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_gpu
flash_attention_ragged_gpu()
max.experimental.nn.common_layers.functional_kernels.flash_attention_ragged_gpu(q, k, v, input_row_offsets, max_seq_len, mask_variant, scale, local_window_size=-1)
Computes flash attention for ragged inputs using a GPU-optimized kernel, without a KV cache.
-
Parameters:
-
- q (Tensor) – The query tensor, of shape
[total_seq_len, num_heads, head_dim](ragged). - k (Tensor) – The key tensor, of shape
[total_seq_len, num_heads, head_dim](ragged). - v (Tensor) – The value tensor, of shape
[total_seq_len, num_heads, head_dim](ragged). - input_row_offsets (Tensor) – The buffer of shape
[batch_size + 1]with dtype uint32. Indicates where each sequence starts and ends in the ragged tensors. The values should be a prefix sum (cumulative sum) of sequence lengths. - max_seq_len (Tensor) – The maximum sequence length across the batch, as a
rank-1
uint32tensor on CPU. - mask_variant (MHAMaskVariant) – The mask variant to use for attention.
- scale (float) – The scaling factor for attention scores.
- local_window_size (int) – The local window size for sliding window attention.
- q (Tensor) – The query tensor, of shape
-
Returns:
-
The output tensor, of shape
[total_seq_len, num_heads, head_dim]. -
Return type: