Flex Attention
This is an experimental API and subject to change.
Overview
Flex Attention implements scaled dot-product attention with a reusable sparse mask plan:
Rather than materializing a dense Boolean mask, every query row provides an odd-length endpoint sequence whose visible keys are the interval union
The endpoint representation supports a range of sparse attention patterns. Rows in the following figure are queries, columns are keys, blue cells are visible, and light-gray cells are masked:

create_mask_plan compiles these endpoints into architecture-native packed
forward and, when requested, backward metadata. The resulting MaskPlan can
be reused with new Q/K/V values that have the same geometry, dtype, and CUDA
device. A single flex_attn_func entry point handles both layouts: supplying
cu_seqlens_q and cu_seqlens_k when the plan is created selects the THD
variable-length path; omitting both selects fixed-length BSHD. Execution uses
the FlexAttentionFwd and FlexAttentionBwd APIBase adapters. Their
check_support(), compile(), and execute() lifecycle supports precompiled,
allocation-free launches with caller-provided outputs and workspace;
flex_attn_func composes the allocating wrappers behind PyTorch custom
autograd.
Requirements
- SM90, SM100, or SM103 NVIDIA GPU
- CUDA-enabled PyTorch
- cuDNN Frontend (its CuTeDSL dependencies are required and install with it)
From a source checkout:
For a published package, install nvidia-cudnn-frontend and a
compatible CUDA-enabled PyTorch separately.
Fixed-length BSHD API
Fixed-length tensor shapes are:
Hq must be divisible by Hkv, supporting MHA, GQA, and MQA. Hmask must
be 1 for a mask shared by all query heads or Hq for per-head masks.
nfunc must be positive and odd. The last mask dimension is flattened in
batch-major order, but every endpoint remains a sample-local key coordinate
in [0, Sk]. Endpoints must be nondecreasing within each query row.
With both output flags disabled (the default), flex_attn_func returns out.
The optional statistics are returned in this order:
LSE is still maintained internally when gradients require it. max_logit[h]
is the maximum of softmax_scale * dot(q[b, i, h], k[b, j, h // (Hq/Hkv)])
over all batches and mask-visible query/key pairs for query head h.
This is a signed maximum, not an absolute maximum. A head with no visible
pair returns -inf. The same [Hq] reduction applies to variable-length
THD inputs across all sequences. return_max_logit=True requires a
non-negative scale; at zero scale a head with visible pairs returns zero.
max_logit is non-differentiable; output and LSE retain their gradient support.
Variable-length THD API
Variable-length Q/K/V are flattened across samples. Sequence geometry is provided when the plan is built and is then owned by the plan:
In this mode Q has shape [total_q, Hq, Dqk], K and V have leading extent
total_k, out has shape [total_q, Hq, Dv], and LSE has shape
[Hq, total_q]. Both cumulative-length tensors are contiguous CUDA INT32
vectors of shape [B + 1], start at zero, end at the corresponding total,
and describe lengths no greater than the supplied maxima. The plan clones
them, so later mutations of the caller’s tensors do not change plan geometry.
Plan and execution options
create_mask_plan(..., pack_gqa=None, build_backward=None) accepts:
pack_gqa: select packed-GQA planning explicitly, or leaveNonefor architecture/configuration-based selection.build_backward: build backward topology explicitly. WithNone, it is enabled when autograd is active and any Q/K/V sample tensor requires a gradient. A plan built without backward payloads cannot later be used for an autograd-enabled call requiring gradients.
The execution function accepts softmax_scale (default
1 / sqrt(Dqk)), deterministic, return_lse, and return_max_logit. Plan reuse requires the
same fixed/variable mode, sequence and head geometry, dtype, and device used
at construction.
Keep MaskPlan outside the execution loop. Q/K/V values and tensor addresses
may change between calls, while the mask and plan geometry stay fixed:
The plan owns immutable forward metadata and, with build_backward=True, the
backward metadata. Mutable scheduler counters, semaphores, and accumulation
buffers are per-execution workspace rather than plan state, so one plan can be
used by independent calls.
Explicit compile and execute APIs
To enable the maximum statistic in the explicit API, provide a contiguous CUDA
FP32 [Hq] sample buffer as FlexAttentionFwd(..., sample_max_logit=max_logit)
and supply max_logit_tensor=max_logit on every execute() call. Its presence,
shape, dtype, stride, and device must match the compiled descriptor. Each call
resets this buffer to -inf on the launch stream before the attention kernel
reduces into it, including during CUDA graph replay. execute() requires no
additional allocation for this output. The statistic uses the kernel’s row
maxima and CTA-local/global atomic reductions; enabling it specializes the
kernel and may affect performance.
Most PyTorch users should use flex_attn_func. Framework integrations can use
FlexAttentionFwd and FlexAttentionBwd directly. Constructors capture sample
tensor descriptors and a sample plan; execute() accepts new tensors and any
compatible plan. Outputs and the workspace_size-byte CUDA uint8 workspace
must be allocated before execute():
The same plan, compiled API objects, output buffers, and workspaces can be
reused for sequential calls with new tensor values that keep the compiled
shape, stride, dtype, and device. Concurrent calls may share the immutable
plan and compiled API objects, but each in-flight call needs separate outputs
and workspace.
flex_attn_func uses internal allocating helpers to manage outputs, workspaces,
compiled-object caches, and autograd state. Explicit backward requires the LSE
produced by forward and a plan built with backward metadata.
Supported configurations and current limits
- FP16 and BF16 inputs; FP32 LSE
- fixed-length BSHD and true variable-length THD layouts
(Dqk, Dv)with each dimension in{8, 16, ..., 128}, plus(192, 128)and the dedicated(256, 256)path- forward and backward on SM90, SM100, and SM103
Paged KV cache, SplitKV, MLA, FP8, SM80, and SM120 are not implemented. Q/K/V must be contiguous in their last dimension. The plan builder may impose additional architecture-specific shared-memory constraints and reports them as validation errors.
Benchmark
The Flex-only static-mask benchmark and its protocol are documented in
benchmark/flex_attention/README.md.
Reference results from the original implementation
The following figures are carried over from the original FlexAttention
repository (docs/assets/static_mask_benchmark*.png, source checkout revision
9c6cbf7). They report measurements on NVIDIA GB300 with BF16 inputs,
B=1, S=128K, and Hq=Hkv=4, for the eight static mask patterns shown above.
The panels report forward and backward active TFLOP/s, counting only visible
query/key pairs. The head dimensions are indicated in each figure.
The figures retain their original comparison labels: PyTorch FlexAttention,
FA4 (native causal / mask_mod), Magi backend, and FlexAttention (ours).
Here, “ours” refers to the original implementation migrated into
cudnn.flex_attention. These are pre-migration reference results; the migrated
cuDNN Frontend wrappers have not been remeasured for these figures. The benchmark
linked above measures the current Flex Attention implementation only.



Design documentation
Flex Attention Mask Plan Design describes interval coordinates, planner stages, partial/full CSR topology, consumer-specific packed predicates, forward/backward consumption, and plan compatibility and ownership.