IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).

Python function

flash_attention_ragged

flash_attention_ragged()​

max.experimental.nn.common_layers.functional_kernels.flash_attention_ragged(kv_params, input, input_row_offsets, kv_collection, layer_idx, mask_variant, scale, local_window_size=-1, sink_weights=None, rel_logits=None, output_dtype=None)

source

Computes flash (self) attention provided the !mo.opaque KV Cache.

Notably, this materializes the attention mask (dependent on MHAMaskVariant) within the kernel. input and input_row_offsets are used together to implement the ragged tensor. input_row_offsets indicates where each batch starts and ends in input.

Note that this is self attention and the KV sequence length is assumed to be equal to the Q sequence length. For KV sequence length != Q sequence length, use cross_attention_ragged().

When rel_logits is set, the kernel gathers an additive relative-position bias by rel_dist = q_pos - k_pos from a (total_q_tokens, heads, extent) table and adds it on every visible position, selecting mo.mha.ragged.paged.rel_logits. This path only supports mask_variant values CAUSAL_MASK (with local_window_size == -1) and SLIDING_WINDOW_CAUSAL_MASK (with a positive local_window_size); sink_weights and rel_logits are mutually exclusive.

Parameters:

  • kv_params (KVCacheParams) – KVCacheParams object containing key-value cache parameters.
  • input (TensorValue) – TensorValue representing the input tensor with shape [total_seq_len, num_heads, head_dim].
  • input_row_offsets (TensorValue) – TensorValue indicating the start and end of each batch in the input tensor with shape [batch_size + 1].
  • kv_collection (KVCacheInputsPerDevice[TensorValue, BufferValue]) – PagedCacheValues object for managing key-value cache.
  • layer_idx (TensorValue) – TensorValue representing the layer index, expected to have dtype uint32.
  • mask_variant (MHAMaskVariant) – MHAMaskVariant specifying the type of attention mask to use. With rel_logits, only CAUSAL_MASK and SLIDING_WINDOW_CAUSAL_MASK are supported.
  • scale (float) – float value used to scale the attention scores.
  • local_window_size (int) – int specifying the size of the local attention window, default is -1 for no local window.
  • sink_weights (TensorValue | None) – Optional tensor of shape [num_heads] containing learnable sink weights for each attention head.
  • rel_logits (TensorValue | None) – Optional relative-position bias table with shape [total_seq_len, num_heads, extent]; row r matches input’s own ragged-flat row convention.
  • output_dtype (DType | None) – Dtype for the attention output. Defaults to input.dtype.

Return type:

TensorValue