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)
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 (PagedCacheValues) – 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, onlyCAUSAL_MASKandSLIDING_WINDOW_CAUSAL_MASKare 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]; rowrmatchesinput’s own ragged-flat row convention. - output_dtype (DType | None) – Dtype for the attention output. Defaults to
input.dtype.
-
Return type: