> For clean Markdown of any page, append .md to the page URL.
> For a complete documentation index, see https://docs.nvidia.com/cudnn/llms.txt.
> For AI client integration (Claude Code, Cursor, etc.), connect to the MCP server at https://docs.nvidia.com/cudnn/_mcp/server.

# 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:

```text
O = softmax(scale * Q K^T + interval_mask) V
```

Rather than materializing a dense Boolean mask, every query row provides an
odd-length endpoint sequence whose visible keys are the interval union

```text
[0, F0) U [F1, F2) U [F3, F4) U ...
```

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:

![Supported static attention mask shapes](https://fdr-prod-docs-files-public.s3.us-east-1.amazonaws.com/nvidia-cudnn.docs.buildwithfern.com/24b36df50bb9bbd650476c941cac8ea89e0f9246b7760e48824a2a65d59d2ee6/_dot_dot_/fe-oss-apis/attention/assets/static_mask_shapes.webp?X-Amz-Algorithm=AWS4-HMAC-SHA256&X-Amz-Content-Sha256=UNSIGNED-PAYLOAD&X-Amz-Credential=AKIA6KXJSKKNFOCF7G4B%2F20261006%2Fus-east-1%2Fs3%2Faws4_request&X-Amz-Date=20261006T223242Z&X-Amz-Expires=604800&X-Amz-Signature=3fd45e9f6e2e4a6ca333dbe4f62cbe74635e0f3e4c934a2eba06249c7717c3a8&X-Amz-SignedHeaders=host&x-amz-checksum-mode=ENABLED&x-id=GetObject)

`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:

```bash
pip install -e .
pip install --group torch
```

For a published package, install `nvidia-cudnn-frontend` and a
compatible CUDA-enabled PyTorch separately.

## Fixed-length BSHD API

```python
import torch
from cudnn import create_mask_plan, flex_attn_func

B, S, H, D = 1, 4096, 8, 128
q = torch.randn(B, S, H, D, device="cuda", dtype=torch.bfloat16, requires_grad=True)
k = torch.randn_like(q)
v = torch.randn_like(q)

# Causal row q sees local key indices [0, q + 1).
mask_func = torch.arange(1, S + 1, device="cuda", dtype=torch.int32).view(1, 1, S)
plan = create_mask_plan(mask_func, q, k, v)

out, lse = flex_attn_func(q, k, v, mask_plan=plan, return_lse=True)
out.backward(torch.randn_like(out))
```

Fixed-length tensor shapes are:

| Tensor | Shape | Dtype |
|---|---|---|
| `q` | `[B, Sq, Hq, Dqk]` | FP16 or BF16 |
| `k` | `[B, Sk, Hkv, Dqk]` | same as `q` |
| `v` | `[B, Sk, Hkv, Dv]` | same as `q` |
| `mask_func` | `[Hmask, nfunc, B * Sq]` | CUDA INT32 |
| `out` | `[B, Sq, Hq, Dv]` | same as `q` |
| `lse` | `[B, Hq, Sq]` | FP32 |
| `max_logit` (optional) | `[Hq]` | FP32 |

`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:

| `return_lse` | `return_max_logit` | Return value |
|---|---|---|
| `False` | `False` | `out` |
| `True` | `False` | `(out, lse)` |
| `False` | `True` | `(out, max_logit)` |
| `True` | `True` | `(out, lse, max_logit)` |

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.

```python
out, lse, max_logit = flex_attn_func(
    q, k, v, mask_plan=plan, return_lse=True, return_max_logit=True
)
```

## 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:

```python
import torch
from cudnn import create_mask_plan, flex_attn_func

q_lengths = (192, 128)
k_lengths = (160, 96)
cu_q = torch.tensor((0, 192, 320), device="cuda", dtype=torch.int32)
cu_k = torch.tensor((0, 160, 256), device="cuda", dtype=torch.int32)
Hq, Hkv, D = 8, 2, 128

q = torch.randn(320, Hq, D, device="cuda", dtype=torch.float16, requires_grad=True)
k = torch.randn(256, Hkv, D, device="cuda", dtype=torch.float16, requires_grad=True)
v = torch.randn(256, Hkv, D, device="cuda", dtype=torch.float16, requires_grad=True)

# Each query sees its whole sample-local K sequence.
endpoints = torch.cat(
    [torch.full((q_length,), k_length, device="cuda", dtype=torch.int32)
     for q_length, k_length in zip(q_lengths, k_lengths)]
)
mask_func = endpoints.view(1, 1, -1)

plan = create_mask_plan(
    mask_func,
    q,
    k,
    v,
    cu_seqlens_q=cu_q,
    cu_seqlens_k=cu_k,
    max_seqlen_q=max(q_lengths),
    max_seqlen_k=max(k_lengths),
)
out, lse = flex_attn_func(q, k, v, mask_plan=plan, return_lse=True)
```

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 leave `None` for
  architecture/configuration-based selection.
- `build_backward`: build backward topology explicitly. With `None`, 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:

```python
plan = create_mask_plan(mask_func, q_sample, k_sample, v_sample, build_backward=True)

for q, k, v in inputs_with_the_same_geometry:
    out = flex_attn_func(q, k, v, mask_plan=plan)
    out.backward(torch.randn_like(out))
```

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()`:

```python
import torch
from cudnn import FlexAttentionBwd, FlexAttentionFwd

plan = create_mask_plan(mask_func, q, k, v, build_backward=True)
out = torch.empty((*q.shape[:-1], v.shape[-1]), dtype=q.dtype, device=q.device)
lse = torch.empty((q.shape[0], q.shape[2], q.shape[1]), dtype=torch.float32, device=q.device)
fwd = FlexAttentionFwd(q, k, v, out, plan, lse)
fwd.check_support()
fwd.compile()

fwd_workspace = torch.empty(fwd.workspace_size, device=q.device, dtype=torch.uint8)
fwd.execute(q, k, v, out, plan, lse, workspace=fwd_workspace)

do = torch.randn_like(out)
dq, dk, dv = torch.empty_like(q), torch.empty_like(k), torch.empty_like(v)
bwd = FlexAttentionBwd(q, k, v, out, do, lse, dq, dk, dv, plan)
bwd.check_support()
bwd.compile()

bwd_workspace = torch.empty(bwd.workspace_size, device=q.device, dtype=torch.uint8)
bwd.execute(q, k, v, out, do, lse, dq, dk, dv, plan, workspace=bwd_workspace)
```

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`](https://github.com/NVIDIA/cudnn-frontend/blob/main/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.

![Static-mask attention performance on NVIDIA GB300, Dqk=Dv=128](https://fdr-prod-docs-files-public.s3.us-east-1.amazonaws.com/nvidia-cudnn.docs.buildwithfern.com/b2bbfbc9b8a24127b8bac3ce8bf879224bbccc525ae8cbb834c25c97b4ab0c41/_dot_dot_/fe-oss-apis/attention/assets/static_mask_benchmark.webp?X-Amz-Algorithm=AWS4-HMAC-SHA256&X-Amz-Content-Sha256=UNSIGNED-PAYLOAD&X-Amz-Credential=AKIA6KXJSKKNFOCF7G4B%2F20261006%2Fus-east-1%2Fs3%2Faws4_request&X-Amz-Date=20261006T223242Z&X-Amz-Expires=604800&X-Amz-Signature=72bf5ebe7201f8e5bc7eab71b155f66fd0954f0c1dd1db02f118220001898e63&X-Amz-SignedHeaders=host&x-amz-checksum-mode=ENABLED&x-id=GetObject)

![Static-mask attention performance on NVIDIA GB300, Dqk=192 and Dv=128](https://fdr-prod-docs-files-public.s3.us-east-1.amazonaws.com/nvidia-cudnn.docs.buildwithfern.com/9dbc5178e0107228251e37d271d5698832182942a86fff4e2da46f5ec8d1f159/_dot_dot_/fe-oss-apis/attention/assets/static_mask_benchmark_d192.webp?X-Amz-Algorithm=AWS4-HMAC-SHA256&X-Amz-Content-Sha256=UNSIGNED-PAYLOAD&X-Amz-Credential=AKIA6KXJSKKNFOCF7G4B%2F20261006%2Fus-east-1%2Fs3%2Faws4_request&X-Amz-Date=20261006T223242Z&X-Amz-Expires=604800&X-Amz-Signature=5927934e2fe79846f31e9bba764befbd711d607902ef39ab3de3b2b29eed09f0&X-Amz-SignedHeaders=host&x-amz-checksum-mode=ENABLED&x-id=GetObject)

![Static-mask attention performance on NVIDIA GB300, Dqk=Dv=256](https://fdr-prod-docs-files-public.s3.us-east-1.amazonaws.com/nvidia-cudnn.docs.buildwithfern.com/347111fffb56e2d637af274b46cd94cbac1e4023f77bb8aa9774213cbd265b9e/_dot_dot_/fe-oss-apis/attention/assets/static_mask_benchmark_d256.webp?X-Amz-Algorithm=AWS4-HMAC-SHA256&X-Amz-Content-Sha256=UNSIGNED-PAYLOAD&X-Amz-Credential=AKIA6KXJSKKNFOCF7G4B%2F20261006%2Fus-east-1%2Fs3%2Faws4_request&X-Amz-Date=20261006T223242Z&X-Amz-Expires=604800&X-Amz-Signature=f2d2ba25838ee556015e46bd953e7e366aee630d0dddaa954432003ce9732405&X-Amz-SignedHeaders=host&x-amz-checksum-mode=ENABLED&x-id=GetObject)

## Design documentation

[Flex Attention Mask Plan Design](/cudnn/fe-oss-apis/attention/flex_attention_design) describes interval
coordinates, planner stages, partial/full CSR topology, consumer-specific
packed predicates, forward/backward consumption, and plan compatibility and
ownership.