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_gpu
flash_attention_gpu()
max.experimental.nn.common_layers.functional_kernels.flash_attention_gpu(q, k, v, mask_variant, scale, local_window_size=-1, valid_length=None)
Computes flash attention using a GPU-optimized kernel.
-
Parameters:
-
- q (TensorValue) – The query tensor, of shape
[batch, seq_len, num_heads, head_dim]. - k (TensorValue) – The key tensor, of shape
[batch, seq_len, num_heads, head_dim]. - v (TensorValue) – The value tensor, of shape
[batch, seq_len, num_heads, head_dim]. - 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.
- valid_length (TensorValue | None) – The optional tensor of shape
[batch]with dtype uint32. When provided, uses the padded kernel variant that respects the valid sequence lengths for each batch element.
- q (TensorValue) – The query tensor, of shape
-
Returns:
-
The output tensor, of shape
[batch, seq_len, num_heads, head_dim]. -
Return type: