Skip to content

feat(attention): cuDNN FROST attention backend for head_dim in (256, 512], with context parallelism - #3527

Open
nvegesna-netizen wants to merge 101 commits into
NVIDIA:mainfrom
nvegesna-netizen:nvegesna/te-frost-d512-cp
Open

nvegesna-netizen wants to merge 101 commits into
NVIDIA:mainfrom
nvegesna-netizen:nvegesna/te-frost-d512-cp

Conversation

@nvegesna-netizen

@nvegesna-netizen nvegesna-netizen commented Sep 16, 2026 •

Copy link
Copy Markdown
Contributor

Problem

get_attention_backend selects no backend at all for symmetric head_dim=512 with context
parallelism, so that configuration raises rather than running:

  • FlashAttention 2 and 3 instantiate kernels up to head_dim 256
  • FA4 is disabled at symmetric 512 on SM100/SM110 (TMEM budget)
  • the C++ cuDNN fused path caps at 256
  • UnfusedDotProductAttention serves 512 but is disabled under CP

Hybrid models such as Gemma 4, which interleave sliding-window layers at head_dim 256 with
global layers at 512, therefore cannot use context parallelism at all.

To be precise about the claim: this makes context parallelism available at symmetric 512 from
released components. It is not a claim to unbounded context. Dao-AILab/flash-attention#2877 adds
symmetric D512 kernels to FA4 and, with #3532, that path also works and scales further, but it is
unmerged and unreviewed. Measured capacity here, and the workspace behaviour that bounds it, are
in Known limitations.

Approach

cuDNN Frontend >= 1.29.0 ships CuTe-DSL ("FROST") SDPA kernels that serve symmetric 512 forward
and backward on SM100/SM103. They are reachable only through the cuDNN graph Python API, since
the FROST engines register at Python import time and need nvidia-cutlass-dsl, while TE's C++
builds against cuDNN Frontend headers only.

FROST is a FusedAttention sub-backend, FusedAttnBackend["FROST"], not a fourth top-level
backend:

  • frost_attention.py holds the kernel wrapper, its plan cache, and fused_attn_fwd/bwd behind
    the cpp_extensions.fused_attn signatures
  • cpp_extensions/fused_attn.py routes to those two functions when the sub-backend is FROST,
    before any C++ call
  • utils.py consults is_frost_attention_supported inside _get_fused_attn_backend, where the
    C++ selector returns No_Backend, and checks availability once at the end of
    get_attention_backend, the way flash-attn's version is checked

So FusedAttnFunc and attn_forward_func_with_cp reach these kernels without knowing which
sub-backend they got, and dot_product_attention.py needs no change at all.

Several declines now come from the existing fused filters rather than their own copies: the
context-parallel mask, window and a2a restrictions; the score_mod and softcap filters;
and FP8, since FusedAttnFunc forces the FP8 sub-backend on that path. The all-gather path's
causal -> causal_bottom_right rewrite applies to FROST for free, which is what gives it the
bottom-right band it needs on a trimmed KV range.

Verification

Full suite on B200 against this head, in a container with nvidia-cudnn-frontend 1.29.0 and
nvidia-cutlass-dsl 4.8.0:

check result
FROST is reached as a FusedAttention sub-backend 7 gates pass
flex_attention.py regression 27 passed
non-distributed numerics vs a float64 reference 60 passed
context parallelism 22 passed
broader test_attention.py suite 2652 passed, 2605 skipped

The broad suite is identical to the pre-change baseline and flex_attention.py's own suite is
clean, which is what says the shared-module extraction below preserved behaviour rather than merely
compiling.

Correctness uses the criterion FlashAttention applies to itself,
err(kernel, fp64) <= 2 * err(inputs rounded to dtype, fp64), with that floor measured per case so
the bar tracks shape and dtype instead of encoding a number that rots. Coverage is square and
rectangular, causal / bottom-right / non-causal, GQA and MHA, forward and backward in bf16 and
fp16, plus fp16 arms under CP.

Context parallelism, each compared against the non-CP reference via run_attention_with_cp.py:

bshd sbhd
p2p CP=2, CP=4 CP=2, CP=4
all_gather CP=2, CP=4 CP=2
a2a CP=2, CP=4 CP=2

Plus 2 nodes x 2 ranks for every comm type including a2a+p2p, with a FusedAttention control
passing throughout. a2a+p2p needs the multi-node arm specifically: at world_size 4 the a2a
level takes consecutive ranks (0,1),(2,3) and the p2p level takes strided ranks (0,2),(1,3),
so block assignment across two nodes puts the all-to-all inside each node and the ring across the
fabric. On one node both levels sit on NVLink and the arrangement is never exercised. The run
establishes the rank-to-host mapping first and declines to report unless it is block ordered,
since round-robin would invert the two levels and pass while testing the opposite topology.

test_frost_attention.py is non-distributed and checks the forward, the LSE convention and the
backward gradients against an independent float64 reference, across both causal alignments, bf16
and fp16, GQA and MHA, and a rectangular shape where top-left and bottom-right masking differ.
The CP tests cannot do this: run_attention_with_cp.py compares a CP run against a non-CP run of
the same backend, so a systematic kernel error appears on both sides and cancels. The reference
must be float64 rather than float32, because torch computes fp32 matmuls in TF32 on Ampere and
newer and TF32's significand is 11 bits, the same as fp16, so an fp32 reference is no more
accurate than the kernel under test.

End to end, a Gemma 4 dense model (sliding layers on FlashAttention, global layers on FROST)
matched its CP=1 result to 9.7e-05 on the loss.

These tests skip in CI as it stands, and the blocker is the dependency floors rather than the
Blackwell requirement. Stock nvcr.io/nvidia/pytorch:26.08-py3 ships nvidia-cudnn-frontend
1.26.0 and nvidia-cutlass-dsl 4.6.2, both below what FROST needs, so
is_frost_attention_available() declines and the tests skip with that reason rather than failing.
qa/L3_pytorch_FA_versions_test/test.sh targets sm100+ but pins nvidia-cutlass-dsl[cu13]==4.4.2
for the FA4 path, so it would skip even on a B200; and nvidia-cutlass-dsl is not a declared TE
dependency at all, arriving transitively via flash-attn-4. Happy to wire this into whichever
lane you consider the right home.

Performance

Against UnfusedDotProductAttention at the same shape (b2 hq8 hkv4 d512 causal bf16, CP=1,
single layer), the only other backend serving this head dim from released components:

seqlen fwd speedup fwd+bwd speedup unfused peak (fwd) FROST peak (fwd)
4096 5.8x 2.4x 1520 MiB 224 MiB
16384 15.2x 6.4x 51232 MiB 1057 MiB
32768 n/a (OOM) n/a (OOM) OOM on 178 GiB 2850 MiB

The unfused path materialises the full s x s score matrix, so it grows O(s^2). Note that in
forward+backward FROST uses more memory than unfused below roughly 8k, where the saved tensors
and workspace exceed the small score matrix; the crossover is between 8192 and 16384.

Notes for reviewers

Two places the "no changes outside frost_attention.py" goal does not quite hold, both the
same root cause: the selected sub-backend was being discarded and re-derived as
F16_arbitrary_seqlen in seven places, so those places now keep it instead.

  • context_parallel.py (+29/-7): each of the three autograd classes recomputed the sub-backend
    locally rather than receiving it. attn_forward_func_with_cp gains one argument, threaded
    through, with the existing derivation kept as the fallback. One further line: the p2p forward
    step returned max_logit with a starred tail, which is only unambiguous while no backend has a
    statically known return length. The FROST branch gives pylint one, so the tail resolved to empty
    and the five-label unpacks tripped unbalanced-tuple-unpacking. Indexing says what the function
    actually returns.
  • backends.py (+5/-5): the non-FP8 backward hardcoded F16_arbitrary_seqlen, discarding the
    value the forward had already saved, so it now simply reads
    ctx_attrs["fused_attention_backend"]. The CP assert admitted only that sub-backend.

FusedAttnBackend.FROST has no pybind counterpart. It is 3, past the C++ enum, and the
import-time sync assert now exempts python-only members. Nothing in transformer_engine/pytorch
constructs NVTE_Fused_Attn_Backend(int), and the dispatch happens before any C++ call, so the
value never reaches pybind. Every post-selection check in utils.py is guarded on == FP8 or
== F16_arbitrary_seqlen, so a FROST value falls through them.

is_frost_attention_supported deliberately does not probe availability. Probing imports cuDNN
Frontend with the FROST engines enabled, which is process-wide and reorders plan selection for
every other cuDNN consumer. That must not happen for the overwhelming majority of configs, which
are nowhere near head_dim 512. Availability is checked once at the end, beside the flash-attn
version checks.

Plan selection is deliberately strict. heur_mode.A with an explicit name-checked
select_plan, rather than A|FALLBACK with HEURISTICS_CHOICE. Without a pin, build_plans
walks the ranked list and finalizes the first plan that builds, logging declines at INFO. At d512
that matters in the forward, where an ordinary engine may build and compute a different function;
the backward is self-limiting, since no non-FROST d512 backward exists. Pinning is also what makes
check_support fatal rather than advisory, which is why it follows the selection.

Version constraint worth knowing. FROST enforces nvidia-cutlass-dsl >= 4.7.0 at plan-build
time while cudnn-frontend only declares >= 4.6.2. Below that floor every FROST engine silently
declines and ordinary backend plans are returned with no error, which is why the code checks the
selected plan by name rather than assuming the engine was used. Stock 26.08 ships exactly 4.6.2.
Separately, 4.7.1 is incompatible with flash-attn-4 4.0.0b11, which CI currently pins.

Scope

SM100/SM103 only (the cuDNN d512 backward is Blackwell-only), bf16/fp16, symmetric head_dim in
(256, 512], bshd and sbhd, mask types no_mask / causal / causal_bottom_right.

Symmetric specifically, and the constraint is the backward engine's. v carries its own graph
node and its own cache-key entry, so an asymmetric pair is expressible here. It is declined because
sdpa_bwd_sm100, the only f16 FROST backward reachable on SM100/SM103, leaves dqk_ge_dv unset,
which cuDNN reads as requiring d_qk == d_v at every head dim rather than only above 256. The
forward is the opposite: sdpa_fwd_prefill_sm100 lists (192, 128) among its native d_shapes
and has no dqk_ge_dv concept at all. So the decline covers both directions rather than training
alone, since is_training is module.training and eval() does not disable autograd, which makes
a served forward no guarantee that no backward follows.

The head-dim range is that same engine's declared envelope: d_envelope_floor=256 (exclusive),
d={512} and d_pad_multiple=8, which is _MIN_HEAD_DIM 257, _MAX_HEAD_DIM 512 and
_HEAD_DIM_MULTIPLE 8 here.

Declined by the selector, each an explicit decline rather than a silent fallback: FP8, attention
bias, dropout, softcap, non-vanilla softmax, thd, max_logit, KV caching, CUDA graph
capture, and deterministic execution. thd is expressible with these kernels but is not implemented here.
Determinism is declined because cuDNN has no deterministic backward for them at all; measured,
not assumed.

Also declined: a right-bounded window on a non-causal mask when max_seqlen_q != max_seqlen_kv.
That combination takes its alignment only from bottom_right_diagonal, which defaults to
top-left, while the all-gather ring measures its window against the bottom-right diagonal. The two
differ exactly when the lengths do, so FROST declines rather than guessing.

Sliding window is supported, with all_gather or a2a. It is declined with p2p and
a2a+p2p, whose ring shards KV across steps so a bound measured against the full sequence does
not survive the per-step tiles, which is the same rule FusedAttention carries. That matters for
the motivating model: Gemma 4's sliding layers are the other half of it.

Known limitations and open decisions

Backward memory grows super-linearly with sequence length. Measured on B200,
b2 hq8 hkv4 d512 causal bf16, single GPU, peak allocated for forward + backward:

seqlen 4,096 8,192 16,384 32,768 65,536
peak 2.0 GiB 8.0 GiB 8.0 GiB 32.0 GiB 112.0 GiB

The forward is not the source: its workspace measures zero at every length above. The growth is in
the backward, where the cuDNN graph's get_workspace_size() request rises from 1.6 GiB at 4k to
74 GiB at 64k, and this module allocates what the graph asks for.

The consequence is a ceiling on the per-rank shard rather than on the model: shards beyond roughly
32k tokens are impractical on a 180 GB device, so a long-context configuration needs a
correspondingly higher cp_size. Largest global sequence that fits:

cp_size 1 2 4 8
longest context 65,536 131,072 131,072 262,144

The request is also layout-dependent, which matters if you try to reproduce these numbers.
Same logical problem, same plan, differing only in the strides the graph is built against:

seqlen 4,096 8,192 16,384 32,768 65,536
bshd strided view 1.13 GiB 4.25 GiB 4.50 GiB 17.00 GiB 66.00 GiB
bhsd contiguous 1.63 GiB 5.25 GiB 6.50 GiB 21.00 GiB 74.00 GiB

The difference is exactly linear, 128 KiB per query token at both ends of the sweep, while the
super-linear term is identical in both layouts. Note the direction: the contiguous layout is the
more expensive one. This module keys its plan cache on each tensor's actual strides, so which
figure applies follows from the caller's layout. Both columns are reproducible without
TransformerEngine by building the same two graphs through the cuDNN graph API and reading
get_workspace_size(), which is how they were measured: nothing executed, no workspace allocated.

The cudnn_pygraph.py extraction is done. frost_attention.py and flex_attention.py now
drive cuDNN through one shared module: the frontend import and its process-wide engine switch, the
per-device handle with its per-call set_stream, the dtype mapping, the BHSD tensor description,
the SDPA forward and backward graph builders, the plan cache, plan finalization and graph
execution. flex_attention.py comes out at +101/-206 against main.

The reason to share these is ownership rather than line count, and the module docstring says so: the
state behind them is process-global, and giving it two owners is how this code has produced bugs
before. flex also picks up a fix on the way, since it was building graphs under whatever device
happened to be current while only the handle named the right one.

Worth stating plainly: this takes the duplication from three pygraph sites to two, not to one.
transformer_engine/jax/cpp_extensions/flex_attention.py carries its own copy, shares no code with
either, and cannot consume a torch-dependent module.

The graph builders are shared as cudnn_pygraph.build_fwd and build_bwd. Each backend passes
tensor descriptors, any auxiliary tensors its callback reads, and its own sdpa arguments, so a
mask, a score_mod and the choice of which engine to pin or bar stay with the backend that wants
them.

What stayed behind is the part that is not cuDNN mechanics: the selector, the TE mask vocabulary,
the runtime guards and the cpp_extensions shims.

flex is barred from the FROST engines, by default rather than on request. Those engines accept
a score_mod graph, pass check_support, build, and then compute without the callback. Measured on
B200: a FROST plan ranks at index 0 and an unpinned build selects it at head_dim 64, 128, 256 and
512, with the output tracking a float64 reference computed without the modifier. The switch that
offers them is process-wide, so declining to ask for them is not enough.

finalize_plans now bars them unless the caller pins a plan by name, which FROST does and flex does
not. The default direction is deliberate: forgetting to exclude gives silently wrong numbers, while
excluding wrongly gives a slower plan or a loud decline. The engine names live in one place, so the
pin and the bar cannot name different engines; previously they were separate string literals in two
files, and cuDNN has renamed that family once already.

It cannot be checked any cheaper than end to end. deselect_engines marks rather than removes,
measured on B200 at head_dim 64, 128 and 256: the plan count is unchanged and index 0 still names
the barred engine after build_plans, so a post-build assertion on the selected plan would fire
falsely. What verifies the bar is test_frost_switch_does_not_change_what_flex_computes, which runs
the score_mod path with the engines on and off and compares. NVTE_FLEX_TEST_REQUIRED exists so a
lane can make that test fail rather than skip.

The underlying cuDNN defect is filed upstream: those engines do have a score_mod capability gate,
but it reads a key the sdpa() path never writes. The bar carries that as a removal condition in
its docstring rather than becoming permanent by default.

One handle is shared across overlapped CP streams. The wait_stream serialization in
context_parallel.py exists for FA3/FA4's internal per-call workspace, and its comment says
FusedAttention keeps the per-step overlap. FROST's workspace is caller-owned and per-call, and
FusedAttention uses the identical one-handle-per-device plus cudnnSetStream-per-call model while
being deliberately left overlapped, so I have not added FROST to that guard.

Dependency floors are intentionally not enforced by the build. Bumping
nvidia-cudnn-frontend to 1.29.0 would force nvidia-cutlass-dsl and apache-tvm-ffi into every
TE-PyTorch install, because 1.29.0 drops the cutedsl extra marker, and it still would not
guarantee the 4.7.0 floor, since the transitive requirement is 4.6.2. Every other optional backend
here (flash-attn 2/3/4, GDN, quack) is undeclared and runtime-probed, which is what this does.
NVTE_FROST_TEST_REQUIRED and NVTE_FLEX_TEST_REQUIRED mirror NVTE_GDN_TEST_REQUIRED, so a lane
meant to cover either can fail loudly rather than skipping. Neither is set in qa: there is no
Blackwell L0 lane for the first, and I cannot verify the L0 container carries the frontend package
for the second. The flex one matters more than it used to, since that file's single numerical test
is now what says this refactor preserved its behaviour.

Documentation. NVTE_FROST_ATTN is in docs/envvars.rst, described as a FusedAttention
sub-backend. The backend tables in docs/examples/attention/attention.ipynb are not updated yet.

@github-actions github-actions Bot added the community-contribution PRs from external contributor outside the core maintainers, representing community-driven work. label Sep 16, 2026
@nvegesna-netizen
nvegesna-netizen marked this pull request as ready for review September 16, 2026 04:35
@greptile-apps

greptile-apps Bot commented Sep 16, 2026 •

Copy link
Copy Markdown
Contributor

RetriggerConfidence Score: 5/5

[High impact] The PR appears safe to merge; no outstanding previous findings or new actionable defects remain.

Summary

Adds a cuDNN FROST attention sub-backend for symmetric head dimensions in (256, 512] on SM100/SM103, enabling context-parallel attention where existing released backends are unavailable.

  • Routes FROST through the existing FusedAttention interfaces, including forward, backward, and context-parallel communication paths.
  • Introduces shared cuDNN Python graph infrastructure used by FROST and flex attention.
  • Adds dependency, architecture, layout, masking, determinism, and unsupported-feature eligibility checks.
  • Adds numerical reference tests and context-parallel coverage for supported communication modes, layouts, and dtypes.
  • Documents the experimental backend and its runtime dependency requirements.

Diagram

%%{init: {'theme': 'neutral'}}%%
flowchart TD
  A[DotProductAttention request] --> B[get_attention_backend]
  B --> C{Existing cuDNN fused backend available?}
  C -->|Yes| D[FusedAttention C++ sub-backend]
  C -->|No| E{FROST configuration supported and available?}
  E -->|No| F[Try another eligible top-level backend or decline]
  E -->|Yes| G[FusedAttnBackend.FROST]
  G --> H[Existing FusedAttnFunc interface]
  H --> I{Context parallelism?}
  I -->|No| J[FROST cuDNN Python graph]
  I -->|Yes| K[CP p2p / all-gather / a2a path]
  K --> J
  J --> L[Pinned FROST execution plan]
  L --> M[Forward / backward execution]
Loading

Reviews (82) · Last reviewed commit: "docs(attention): shorten the cudnn_pygra..." · Reviewed by Greptile

Comment thread transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py Outdated
Comment thread transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py Outdated
Comment thread tests/pytorch/attention/test_attention_with_cp.py
nvegesna-netizen and others added 8 commits September 15, 2026 23:21
Adds frost_attention.py, a thin PyTorch wrapper around the CuTe-DSL ("FROST")
SDPA kernels in cuDNN Frontend >= 1.29.0. These are the only kernels that serve
symmetric head_dim 512 forward and backward on SM100/SM103, a range no other
backend covers: FlashAttention 2 and 3 cap at 256, FA4 is disabled at symmetric
512, and the C++ cuDNN fused path caps at 256.

Measured on B200 before writing the integration, and each result shaped the code:

- Correctness against the criterion FlashAttention applies to itself,
  err(kernel, fp64) <= 2 * err(naive_bf16, fp64): 0.21x to 0.94x across square
  and rectangular, causal and non-causal, GQA and MHA shapes. An absolute error
  is uninterpretable without that floor.
- The forward LSE is natural-log logsumexp in fp32, matching an fp64 reference
  to 1.8e-06. This is what makes a context-parallel ring merge valid at all.
- Outputs are bitwise reproducible across runs, ruling out a racing split-KV or
  atomic reduction.
- Plan building costs ~1972 ms cold and ~12 ms once cuDNN caches the JIT,
  against a ~0.129 ms execute. Hence _PLAN_CACHE: at ~15000x an execute, caching
  is required rather than an optimisation.

Design notes:

- The cache holds compiled plans only, never output buffers. Buffers are
  allocated per call so a reused plan cannot make one call overwrite another's
  result, and with torch.empty_strided rather than empty_like, which does not
  preserve an arbitrary permuted stride.
- Graphs are built from each tensor's ACTUAL strides, so bshd and sbhd are both
  served without a transpose. sbhd matters because Megatron uses it internally
  and copying every tensor per call would be a real cost.
- _MASK_MODES lists only spellings verified behaviourally. cudnn sdpa() takes
  **kwargs and silently ignores names it does not recognise, so a typo would
  apply no mask and still build and run; inspect.signature is no help either,
  reporting no mask parameters at all. Both top-left and bottom-right causal are
  needed: the p2p ring produces square diagonal tiles where the two coincide,
  while all_gather trims KV and relies on bottom-right, where they differ by
  three orders of magnitude.
- Unsupported configurations are refused rather than approximated, because the
  failure mode of guessing is silent numerical corruption, not an exception.

Scope: SM100/SM103 only (the cuDNN d512 backward is Blackwell-only), bf16/fp16,
symmetric head_dim in (256, 512], bshd and sbhd. thd needs varlen support that
is feasible but not implemented here.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
…tention

Makes the FROST kernels reachable. Before this, get_attention_backend selected
NO backend for symmetric head_dim 512 with context parallelism: FlashAttention
and FusedAttention decline the head dim, and UnfusedDotProductAttention is
disabled under CP. That combination raised rather than running, which is the
gap this series closes.

- get_attention_backend admits head_dim in (256, 512] on SM100/SM103 and
  returns a new use_frost_attention flag. It is consulted only where the
  established backends cannot run the shape, so it never displaces a faster
  path, and it is preferred over the unfused path, which covers the same shapes
  but cannot do CP.
- FrostAttention and FrostAttnFunc in backends.py. TE selects a module class per
  backend, and there was none for cuDNN's Python kernels, so
  attn_forward_func_with_cp was unreachable for them.
- FrostAttention is deliberately narrow: no FP8, bias, dropout, softmax offset
  or paging. Threading a flag through FusedAttention instead would have pulled
  FROST into all of that machinery; the selector declines those configurations
  first, so anything reaching the module is already supported.
- FrostAttnFunc covers the non-CP path only. The CP path does not go through it
  because the ring must interleave per-step kernel calls with KV exchange and
  LSE correction rather than treating attention as one opaque autograd node.

use_frost_attention is a separate flag rather than a FusedAttnBackend value:
that enum mirrors NVTE_Fused_Attn_Backend value-for-value and is consumed by
fused_attn_fwd, which dispatches into C++ that caps at 256, so routing FROST
through it would feed a value into a path that cannot honour it.

Two contracts worth noting, both of which produce runtime errors rather than
type errors when missed: TE attention modules return heads flattened into the
last dimension ([b, s, h*d]), and both return paths need it; and the
"no backend is available" guard must count the new flag, or selecting FROST
alone raises the very error this change removes.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
… and a2a

Dispatches the FROST kernels per step in all three CP comm types, which is what
makes head_dim 512 usable for long-context training rather than only at CP=1.

All three are needed for a real model. Gemma-4 dense is hybrid: sliding-window
layers at head_dim 256 alongside global layers at 512. TE asserts that
sliding-window attention requires a2a or all_gather, never p2p, so a p2p-only
backend passes every attention test and still cannot run the target model.
(MCore accepts cp_comm_type as a per-layer list, so a mixed configuration also
works: sliding layers on all_gather, global layers on the cheaper p2p ring.)

- p2p adds cp_p2p_{fwd,bwd}_frost_attn beside the existing fused and flash
  helpers. A ring step is only ever given causal, no_mask or a padding variant,
  so the dense case needs just causal on or off.
- all_gather is simpler: KV is already gathered and trimmed, so each step is one
  call with no LSE correction. It does require BOTTOM-RIGHT causal, because
  get_kv_seq_info_after_all_gather trims KV and returns a window that is causal
  relative to the trimmed range. Top-left and bottom-right coincide only when
  SQ == SKV, which all_gather never produces, so the wrong choice here would be
  silent corruption rather than an error.
- a2a is simplest of all: after the all-to-all each rank holds the full sequence
  for a subset of heads, so there is no ring and no correction.

Ring tiles are slices and do not carry the strides the cuDNN graphs are built
for, so every to_frost_layout call site passes contiguous tiles. This can copy;
correctness first, worth revisiting if it shows up in a profile.

Also relaxes an sbhd guard that inferred "not fused means flash". FROST is
neither, and its graphs are built from actual strides, so sbhd is served
directly; this matters because Megatron uses sbhd internally. The remaining
instances of that inference are safe by construction (sliding windows and thd
are already declined by the selector).

Note for future work here: the three CP autograd classes are similar enough to
invite generalisation and different enough to punish it. They do not carry
identical ctx state, their aux_ctx_tensors differ in shape, and the tensor saved
for backward is not always the value the branch returns. Adding a branch means
checking what the enclosing function initialises and later consumes, not what
the neighbouring class does. Each class also requires its backward to return one
gradient per forward input; adding a parameter without the matching None breaks
every existing user of that comm type, not just the new path.

Verified on B200: {p2p, all_gather, a2a} x {bshd, sbhd} x {CP=2, CP=4} against
the non-CP reference, plus 2 nodes x 2 ranks, with a FusedAttention regression
control passing throughout.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Adds model_configs_frost_attn with Gemma-4 global-layer shapes (head_dim 512,
GQA and MHA, causal and no_mask) and a FrostAttention kernel_backend in the CP
runner. TE's existing CP matrix stops at head_dim 192, so nothing covered the
range this backend exists for.

The runner leaves NVTE_FLASH_ATTN and NVTE_FUSED_ATTN at 0 and sets
NVTE_FROST_ATTN=1 rather than relying on fallthrough. FROST is the only backend
serving head_dim > 256, so the selector would pick it either way, but making it
an explicit kernel_backend keeps the test honest about what it exercises.

Also updates the three call sites that unpack get_attention_backend for the
added return value.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
model_configs_frost_attn was defined but no pytest function parametrized over
it, so the configs were only reachable by invoking run_attention_with_cp.py
directly and CI would never have executed them.

test_cp_with_frost_attention covers p2p / all_gather / a2a across bshd and sbhd.
It skips rather than fails where the backend cannot run, reporting the reason
from is_frost_attention_available(). That matters for the less obvious
dependency: cuDNN Frontend declares nvidia-cutlass-dsl >= 4.6.2 but FROST
enforces >= 4.7.0 at plan-build time, and below that floor every FROST engine
silently declines and ordinary backend plans come back with no error. An
environment without the dependencies, or without an SM100/SM103 GPU, therefore
reports a skip with a reason rather than a failure that looks like a bug.

thd and a2a+p2p are excluded because the backend declines them.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
…s by device

Two defects found in review.

The availability check verified nvidia-cutlass-dsl but never the cuDNN Frontend
version, and 1.29.0 is the first release carrying the head_dim=512 backward. The
repo pins nvidia-cudnn-frontend>=1.28.0, and a 1.28.0 wheel does ship the d512
forward along with an importable cudnn.sdpa, so the import guard passed, the
forward plan built, and the first backward raised mid-step. The plan-build error
also named only cutlass-dsl, pointing users at the wrong package; it now reports
both versions and their floors.

The plan cache key omitted the device, so in a single-process multi-GPU run the
same shape on a second device would reuse a graph built under the first while
allocating tensors and workspace on the second. Both of TE's other cuDNN caches
already guard against this: the C++ fused-attn cache keys on device_id to
"distinguish graphs on different GPUs in a single-process run", and
flex_attention keys its Python cache on device.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
…version gate

The device element added to _key() in 63294e8 made the key 12 items while
_build_fwd and _build_bwd still unpacked 11, so the first FROST call of any kind
raised ValueError: too many values to unpack. Every path was affected, forward
and backward, CP and non-CP. Nothing caught it because the only test that reaches
_key is gated on SM100/SM103. The unpack is now starred, so further device
components cannot reintroduce the same break.

The key also lacked device.type, letting a CPU tensor alias cuda:0, and called
torch.cuda.current_device() unguarded for a normalisation that a materialized CUDA
tensor never needs -- which would raise a confusing CUDA-init error on a CPU-only
host. It now keys on (type, index) directly, following _score_mod_device_key in
flex_attention.py.

Version handling moves to packaging.Version over distribution metadata, matching
_cudnn_frontend_version_supported in fused_mla_q_uproj.py. The hand-rolled parser
accepted 1.29.0rc1 as 1.29.0, and by replacing the old int() parse it had quietly
relaxed the cutlass floor to admit 4.7.0rc1 as well; both are rejected again. An
undeterminable version no longer reports as "0" and hard-declines a valid source
install -- it defers to _select_frost_plan, which checks the plan by name. That
error message now looks both versions up defensively, since it previously could
raise PackageNotFoundError while formatting the very diagnostic explaining a
failure.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
@nvegesna-netizen
nvegesna-netizen force-pushed the nvegesna/te-frost-d512-cp branch from 7d453a4 to 4069bfa Compare September 16, 2026 06:22
nvegesna-netizen and others added 6 commits September 15, 2026 23:29
…ion probe

Both FROST graphs declare v with k's shape and stride, and the plan-cache key
records only q's and k's, so a v laid out differently from k would hit a plan
built for k's layout and read the wrong elements with no error -- and in the
backward, dv is allocated from v's own stride, disagreeing with the stride the
graph declared. The forward checked shapes but not strides; the backward checked
neither. Both now share one guard. Callers in TE always split k and v from a
single QKV tensor, so this costs nothing and only closes a silent wrong answer.

The version probe now returns the raw string alongside the parsed version, so
"not installed" is distinguishable from "installed but unparseable". Previously
both collapsed to None, which made an odd version string report as not installed
and hard-decline a valid install -- the failure this was meant to remove. Only
absence declines now; an unparseable version defers to the plan-name check, which
is what the accompanying comment already claimed. The module fallback applies to
that case too, and the plan-build error prints the raw string rather than a tuple.

Also corrects a comment in backends.py stating frost_attention raises on any
non-BSHD-contiguous layout. It does not: the graphs are built from each tensor's
actual strides, and .contiguous() is there to keep one plan per shape.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
…n value

get_attention_backend gained a seventh return value, but the mixed-THD mask-policy
path in _get_thd_policy_attention_backend still unpacked six, so every caller of
that path raised ValueError: too many values to unpack. This had nothing to do
with FROST -- it broke existing users of mixed-THD attention. The same function
also rebuilt _attention_backends without use_frost_attention, leaving a stale
value for the read at the scalar forward. Both are fixed, and the fake selector
in test_mixed_thd_attention.py is updated to the same arity so it keeps matching
the real signature rather than masking a mismatch.

The CP runner now asserts that FrostAttention was the backend actually selected,
not merely the one requested. The guarantee was previously emergent -- flash and
fused are env-gated off and CP disables unfused, leaving FROST the only candidate
-- so the assert passes by construction today. It is there so the tests fail
rather than silently exercise another kernel if that ever stops holding.

Docstrings drop the measured timings and error magnitudes. They were accurate,
but no comment in transformer_engine/pytorch or transformer_engine/common cites
figures like these; the fused-attention graph cache states the same constraint
qualitatively. Benchmark numbers from one machine and one shape rot silently, so
the constraints stay and the measurements live in the pull request instead.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
…ong-answer paths

The graphs were built and executed without a cuDNN handle, so cuDNN ran on its
default handle's stream while the tensors and workspace were allocated on
PyTorch's current stream, with nothing ordering the two. That is live on the
path this backend exists for: the p2p ring issues attention inside
`with torch.cuda.stream(cp_stream)`, so on alternating ring steps the kernel and
its buffers sat on different streams. flex_attention.py and the C++ fused path
both bind the stream explicitly; this now does the same, per device, rebinding on
every call because one cached plan is executed from different streams.

Validation now covers what the builders assume. Every node but stats is declared
from q's dtype and execute() binds raw pointers, so a tensor of another dtype had
its bits reinterpreted silently -- dout mattered most, since it arrives from
autograd. k's batch and head_dim, out and dout's shapes, and softmax_lse's dtype
and shape were likewise assumed and unchecked, and the backward additionally
skipped the GQA divisibility check the forward has, which the CP ring can reach
by calling it directly.

Three configurations were selectable but unsupported, each a wrong answer rather
than an error: CP with causal cross-attention or bottom-right masking, which the
ring's square-tile chunking cannot serve and which both other CP backends already
decline; return_max_logit, where this returns a bare tensor while the unfused path
it displaces returns a pair; and load_balancing_strategy, which was dropped on the
way to the ring and silently reverted to DUAL_CHUNK_SWAP. The first two now
decline, the third is threaded through. KV caching declines explicitly -- it was
already unreachable via the padding-mask assert, but only indirectly.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
… names

_handle_for created the handle inside a `with torch.cuda.device(...)` block but
returned before cudnn.pygraph() was called, so graph construction and the
CuTe-DSL plan build ran under whatever device happened to be current, with only
the handle carrying the intended one. A JIT compile path is more likely to read
the ambient CUDA context than the handle, and the guard costs nothing, so the
build now happens under the device the cache key names. Unreachable from TE's
own callers, which always run on the rank's own device, and flex_attention.py
has the same shape -- this is hardening, not a fix for a live bug.

Also corrects the rationale on the new context-parallel mask declines. It read
as though any unequal q/kv length is wrong under CP, which would indict no_mask
too; the restriction is specifically about where the causal diagonal sits, and
no_mask stays allowed when the lengths differ.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
…tself

The CP tests compare a context-parallel run against a non-CP run of the same
backend, so they validate the ring plumbing and nothing about the kernel: a wrong
softmax scale, a causal mask anchored to the wrong corner, or an LSE in the wrong
log base appears identically on both sides and cancels. FrostAttnFunc, the path a
single-GPU head_dim 512 user takes, had no coverage at all.

test_frost_attention.py checks forward output, the LSE convention and the
backward gradients against an fp32 reference computed independently of TE and of
cuDNN, over both causal alignments, both dtypes, GQA and MHA, and a rectangular
shape where top-left and bottom-right masking differ. The bar is the criterion
FlashAttention applies to itself -- error within 2x what the reference itself
incurs from reduced-precision inputs -- measured per case rather than hard-coded,
so it tracks the shape instead of encoding a number that rots. Inputs are
generated in fp32 and cast down, because rounding an already-rounded tensor would
collapse that floor to zero. It also covers the decline paths and the k/v
mismatch guards. Registered in qa/L0 alongside the sibling backends.

Separately, the availability probe now runs after the shape and dtype checks
rather than before. Probing imports cuDNN Frontend and sets
CUDNN_FRONTEND_ENABLE_FROST_ENGINES, which registers engines process-wide and is
therefore visible to flex_attention and the GDN path. That happened for every
attention configuration on any Blackwell machine, at any head dim, including the
overwhelming majority nowhere near 512. It now happens only for a configuration
FROST could actually serve. An explicit CUDNN_FRONTEND_ENABLE_FROST_ENGINES=0
also declines cleanly instead of raising later from plan selection.

Finally, the claim that TE's C++ fused path caps at 256 was wrong: that dispatch
applies no head-dim test and simply asks cuDNN for a graph, so the ceiling is
cuDNN's engine coverage. Stated correctly, along with the actual reason a Python
backend is required -- FROST engines register at Python import time and need
nvidia-cutlass-dsl, while TE's C++ builds against frontend headers only.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
… the backend

cuDNN's SDPA backward is non-deterministic unless the graph asks otherwise --
that is why the C++ fused path calls set_deterministic_algorithm and why
flex_attention passes use_deterministic_algorithm to the same sdpa_backward this
module builds. FROST passed neither, so NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 was
silently not honoured while every other backend either honours it or declines.
The flag is now threaded from DotProductAttention through FrostAttnFunc and the
three context-parallel backward wrappers into the graph, and it is part of the
plan-cache key: the deterministic backward is a different algorithm, so a plan
built one way must not serve a call that asked for the other.

The availability probe gains an NVTE_FROST_TEST_REQUIRED escape hatch, mirroring
NVTE_GDN_TEST_REQUIRED, so a lane intended to cover this backend fails loudly
instead of skipping silently. It is deliberately not set in qa yet, since no
Blackwell L0 lane exists to set it on.

Documents NVTE_FROST_ATTN in docs/envvars.rst, placed by that file's
backend-preference ordering rather than alphabetically, and corrects the stated
preference order, which omitted FrostAttention entirely. FROST sits between
FusedAttention and UnfusedDotProductAttention and is only ever eligible in the
(256, 512] head_dim band that flash and fused do not serve, so it never displaces
a backend that could otherwise have run.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Comment thread docs/envvars.rst Outdated
nvegesna-netizen and others added 5 commits September 16, 2026 00:47
…lism

UnfusedDotProductAttention also serves symmetric head_dim in (256, 512] -- there
is no head-dim filter against it anywhere -- so calling FrostAttention the only
backend for that range was wrong, and would tell a user without context
parallelism that they need a Blackwell-only dependency stack they do not.
FrostAttention is the only backend for that range *with* context parallelism;
without it, unfused covers the same shapes and FROST is merely preferred.

The same paragraph also claimed FrostAttention never displaces a backend that
could otherwise have run, which is wrong in the other direction: it suppresses
UnfusedDotProductAttention when both are eligible. Says so now.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
… p2p wrapper

Threading deterministic into the FROST backward wrappers used a match on the
trailing out_part/dout_part/section parameters, which is not unique to the FROST
one: cp_p2p_bwd_fused_attn ends the same way and already took deterministic
positionally. It therefore gained a second, keyword copy and the module stopped
compiling, taking all of transformer_engine.pytorch down with it.

Not caught before pushing because the syntax check used ast.parse, which parses a
duplicate argument happily -- CPython only rejects it when building the symbol
table in compile(). Verified now with compile() across every file the branch
touches, plus an AST scan for any repeated parameter name.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Measured on B200 with cuDNN Frontend 1.29.0: asking sdpa_backward for a
deterministic algorithm is refused outright -- cudnnGraphNotSupportedError, no
engine proposes a plan for the graph. So unlike the C++ fused path, which opts in
via set_deterministic_algorithm, there is nothing here to opt into, and passing
the flag alone would turn NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 from a silent
violation into a hard failure at plan build.

The selector now declines FROST when determinism is required during training,
which is what the other backends do where they cannot honour it. The graph still
passes use_deterministic_algorithm, so the decline lifts on its own if cuDNN ships
a deterministic d512 backward. The context-parallel tests are unaffected: their
runner sets NVTE_ALLOW_NONDETERMINISTIC_ALGO=1 explicitly.

Also fixes the k/v layout guard test, which could never have failed: it built the
mismatched v as a [b, h, s, d] contiguous tensor, whose strides are exactly those
of a contiguous k, so there was nothing to reject. It is now built as sbhd and
permuted, which keeps the shape and the contiguous head dimension while genuinely
differing in stride order.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
…not a reference

The oracle compared the kernel against an fp32 reference, and on Ampere and newer
torch computes fp32 matmuls in TF32. TF32's significand is 11 bits -- exactly
fp16's -- so for fp16 inputs the "reference" was no more accurate than the kernel
it was judging. Rounding the inputs to fp16 then changed the reference almost not
at all, and the measured error floor collapsed from about 1e-03 to 3e-08,
reducing the bound to the bare absolute slack.

That is how it presented on B200: all seven failures were fp16 with a causal
mask, where the kernel's error is a perfectly normal 1.55e-03 but the bound had
become 1e-03. bf16 was unaffected because its 8-bit significand is far coarser
than TF32, so its floor stayed honest -- which is exactly why the flaw looked
like an fp16-specific kernel problem rather than a broken reference.

The reference is now float64 throughout, immune to TF32 and to whatever the
ambient precision flags are. With it the fp16 causal floor returns to 1.49e-03
and the bound to 3.99e-03, comfortably above the observed error, while a
deliberate 1% scale error is still rejected in every dtype and mask combination.
Tests renamed accordingly, since they no longer compare against fp32.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
…ding window

The backend allowed exactly three mask spellings from a hardcoded table and
declined every sliding window. Neither restriction was necessary: cuDNN's own
engine descriptor for sdpa_bwd_sm100 declares swa and right_band_widening, and
the legacy spellings are not a separate mechanism at all -- pygraph/sdpa.cpp
desugars use_causal_mask to (TOP_LEFT, right_bound=0) and
use_causal_mask_bottom_right to (BOTTOM_RIGHT, right_bound=0), and refuses to
combine either with an explicit right bound.

Masking is therefore built the way the C++ fused path and the in-flight Python
port both build it: a diagonal alignment plus a two-sided band. Causal,
bottom-right and sliding window come from one mechanism instead of three
spellings, the window travels in the plan-cache key, and the all-gather path no
longer raises on a window it can now serve.

The old justification for the allowlist was also wrong. It claimed sdpa()
silently ignores unknown kwargs, so a misspelling would apply no mask and still
run. sdpa is a pybind function with an explicit named-argument list and no kwargs
catch-all; an unknown keyword raises TypeError. The error is deferred to plan
creation rather than raised at validate, which is presumably where the belief
came from, but it is loud, not silent.

Separately, head_dim is now required to be a multiple of 8. The engine pads to
that multiple, so 260 sat inside the advertised (256, 512] range, passed the
gate, and then failed at plan selection complaining about missing engines instead
of declining cleanly.

The oracle test gains sliding-window cases against the float64 reference,
including an assertion that a window changes the output -- a dropped bound would
otherwise still produce finite, plausible numbers.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Comment thread transformer_engine/pytorch/attention/dot_product_attention/utils.py Outdated
nvegesna-netizen and others added 2 commits September 16, 2026 09:43
…for p2p

Allowing sliding window opened two context-parallel paths that could not serve
it. The a2a helpers took no window and so ran plain causal attention with the
left bound silently dropped -- finite, plausible output and wrong gradients. The
p2p ring cannot serve it at all, because a left bound measured against the full
sequence does not survive the per-step KV tiles.

a2a now carries the window, which matches what it can actually do: after the
all-to-all each rank holds the full sequence for a subset of heads, so the user's
window applies unchanged. p2p and a2a+p2p decline, which is the same rule
FusedAttention already carries a few hundred lines above, for the same reason.
all_gather was already correct.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
…y, and per cp_comm_type

The window was exercised only in the forward, yet it changes the backward graph's
dK/dV accumulation rather than just a mask fill, so dq/dk/dv under a window were
entirely unvalidated. The backward test now parametrizes over it.

Adds window=(0,0), the boundary of cuDNN's convention: left_bound counts visible
tokens including the diagonal and has a documented minimum of 1, so this is the
value where an off-by-one stops producing wrong numbers and starts producing an
error instead. Adds the window-validation cases to the decline test, which were
unreachable from the suite even though is_frost_attention_supported accepts and
routes the argument.

Adds a selector test for the rules the previous commit introduced, which shipped
untested: all_gather and a2a may serve a window, p2p and a2a+p2p decline it, and
configurations without a real window must still select FROST under p2p -- the
decline has to key on the window rather than on p2p itself.

Also corrects docs/envvars.rst, which still said the backend declines sliding
window, and which omitted both the multiple-of-8 head_dim constraint and the
determinism decline.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
nvegesna-netizen and others added 20 commits October 6, 2026 22:31
cached_graph entered the requested device for a cache miss but not for an
uncacheable key, so a flex score_mod callback that cannot be keyed would
JIT-compile its plan against whatever device was current. Reachable in a
single-process multi-GPU run where the tensors are not on the current device.

Restructured to one build site rather than one per branch, since a second was
free to forget the scope, which is how this happened.

Reported by Greptile on NVIDIA#3527.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Those engines accept a score_mod graph, pass check_support, build, run, and
then compute without the callback. The switch that offers them is process-wide
and never unset, so flex gets them ranked first once anything else in the
process has enabled them, even though flex never asks.

Barring them is now finalize_plans' default rather than something a caller
opts into. The failure is asymmetric: forgetting to exclude gives silently
wrong numbers, while excluding wrongly gives a slower plan or a loud decline,
so the burden belongs on the caller that wants those engines. A caller pinning
a plan by name has already said which engine it wants, so exclusion is skipped
there rather than contradicting the pin.

The engine names live in cudnn_pygraph now and frost's pin reads them from
there. Previously the pin and the bar were separate string literals in two
files that had to agree; cuDNN has already renamed this family once, and if
they drifted frost would fail loudly while flex failed silently.

Two details the earlier version of this fix did not have:

- deselect_engines is reached through getattr, so a frontend predating it
  degrades instead of raising for every caller.
- when the bar removes the last viable plan, the error names it, so "graph is
  not supported" does not misattribute.

Verified: unpinned and silent bars both tokens, explicit () bars nothing,
explicit list bars that list, pinned bars nothing, and a graph without
deselect_engines still returns its workspace.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Measured against the neighbours rather than guessed. flex_attention.py, the
closest sibling, runs a median docstring of 1 line with a maximum of 8;
backends.py is 1 and 24. The two new files were at a median of 5 and 6 with
maxima of 29 and 12, which is a different convention, not a stricter one.

Trimmed by one rule: state each hazard once where it belongs and use one or two
lines everywhere else. So the stream rebinding, the import-memo ordering, the
plan pin, the engine bar, the diagonal-band off-by-one and the score_mod
exclusion all keep their explanations, while the small helpers that merely
reorder a tuple lose theirs. The three _check_* guards shrank because
_validate_qkv now carries the rationale they were each restating.

Now a median of 5 and 4 with maxima of 16 and 12, inside backends.py's range.
Inline comments were already in register and are untouched: the longest run in
either file is 4 and 5 lines, against 19 in backends.py.

Verified docstring-only: both files' ASTs are identical to before once
docstrings are stripped.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
cc3f601 accepted an asymmetric pair on the grounds that the restriction was
ours rather than the kernels'. Half right. Measured on B200 with cuDNN Frontend
1.29.0 at d_qk=512, d_v=320: the forward is correct against a float64
reference, and the backward has no plan at all. cuDNN says why:

    Embedding dim per head for q and k is not less than equal to 128 at:
      d_qk > 128 && !is_bprop_d_qk_192_d_v_128_case && !is_d_qk_256_d_v_256_case

so its d_qk > 128 backward covers 192/128 and 256/256 and nothing else, and no
engine proposes a plan for anything in between.

Declined outright rather than for training alone. is_training is
module.training and eval() does not disable autograd, so accepting an
asymmetric pair in eval mode would not stop a backward from being recorded and
then failing at plan build, far from the config that caused it.

What stays is the plumbing: v keeps its own graph node and its own cache-key
entry rather than borrowing k's. That is correct whether or not head dims can
differ, it closes a silent-wrong-answer path if k and v ever diverge, and it is
what made this measurable instead of assumed.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Both latent, both the same shape: a check that exists and does not run on the
path that reaches the kernel.

dqkv_layout was accepted and never read. The backward re-derives the format
from qkv_layout and returns the gradients through that, so a caller passing a
different dqkv_layout got them in a layout it did not ask for, silently. It now
joins the o_format and do_format check beside it. The selector already declines
the mismatch, so this matters for a direct call.

_te_mask_spec unpacked window_size before _mask_spec ever saw it, which left
two of _mask_spec's guards dead on the fused path: a malformed window escaped
backend selection as a TypeError or ValueError rather than declining, and the
selector only catches NotImplementedError. The normalisation is now its own
function that both callers go through, and it rejects a non-pair and a pair
that is not integers, the latter having slipped past the length check before.

Verified both directions over int, short tuple, long tuple, string and float
windows, and that well-formed windows are unchanged.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
It belongs with the rest of the TE-to-cuDNN translation. bhsd_dim_stride is
already there and takes a TE format string, so keeping the band translation out
on the grounds that it speaks TE's vocabulary was a line drawn in the wrong
place: both turn TE's spelling into cuDNN descriptors, and neither decides
anything.

The module docstring claimed no attention semantics, which was never quite
true and is less true now. It states the real boundary instead: this module
knows how to say a thing to cuDNN, not which thing to say. No backend policy,
no knowledge of any backend's cache-key layout.

_mask_options went with it, being a two-line wrapper that unpacked a pair.

flex is untouched. Giving it an optional band is a separate question with a
measured answer: its one entry point is reached only when score_mod is not
None, a band and a score_mod cannot share a cuDNN graph, so the parameter would
be permanently None and would still require a mask field in both of its cache
keys.

Verified the translation is byte-identical across every mask and window the
backend serves, the off-by-one included: cuDNN's left bound counts the diagonal
and TE's window_size does not, so a window of w stays a left bound of w + 1.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
…dnn_pygraph"

This reverts commit 86609ba.

The function itself was verified byte-identical across every mask and window
the backend serves, off-by-one included, but the call site changed and that has
only been simulated. The version in frost has three cluster runs behind it.
Not worth carrying an unverified change on a branch that is otherwise close to
green, for a placement question whose answer is a mild preference either way.

Worth reopening after the next run, since the request to move it was explicit
and anticipated the single-caller objection: "Diagonal_band_kwargs but score
mod can't use it".

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
The last piece of plumbing both backends had: allocate a uint8 workspace of the
size the build reported, resolve the device, and run the graph on a handle
rebound to the current stream. flex had it as a helper, frost did it inline
twice.

The handle is resolved inside execute_graph rather than by the caller, so it is
always rebound immediately before execute. That ordering is what keeps cuDNN on
the same stream as the tensors, and having one place to get it wrong is better
than three.

Verified by simulating both callers against the merged form over cuda:0, cuda:1
and an index-less cuda device: the same allocation, on the same resolved
device, followed by the same handle request and the same execute. The
index-less case needed care, since frost passed the device through unnormalised
and relied on torch and handle_for each resolving it independently; the merged
form normalises once and reaches the same physical device.

_get_cudnn_current_stream_handle in flex had no callers left and is gone.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
…piry

It exists because a capability gate reads a key the sdpa() path never writes.
When that is fixed the default should go rather than outlive its cause, and
nothing in the code said so. Also notes that the bar is not self-verifying,
since it matches engine names by a substring that a rename would silently
break.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Measured on B200: deselect_engines marks rather than removes. The plan count is
unchanged and index 0 still names the barred engine after build_plans, at
head_dim 64, 128 and 256 alike. So a post-build assertion on the selected plan
name would fire falsely and is not worth adding.

The previous note called the bar not self-verifying, which overstated it. What
verifies it is test_frost_switch_does_not_change_what_flex_computes, which runs
flex's score_mod path with the engines on and off and compares: if the bar
failed, a FROST plan would answer and drop the callback. That test runs rather
than skips on hardware, which is visible in the suite going from 25 to 27
passed when it landed, with no skip reported.

What is missing is a cheap in-process check, not verification.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Most of this file runs on CPU with the builder monkeypatched. The two tests
that reach cuDNN skip silently wherever the frontend package is absent, so a
lane meant to cover them can be green having never run them.

That matters more than it used to. test_frost_switch_does_not_change_what_flex_
computes is what verifies the engine bar in cudnn_pygraph: it runs the score_mod
path with the FROST engines on and off and compares, and a FROST plan answering
would drop the callback. Measured on B200, deselect_engines marks rather than
removes, so the plan list still names the barred engine afterwards and the bar
cannot be checked any cheaper way.

NVTE_FLEX_TEST_REQUIRED follows the GDN, GDN2, GDP and FROST pattern already in
this directory. Deliberately not set in qa/L0_pytorch_unittest, which cannot be
assumed to carry the package -- the same reason the FROST line there does not
set its own.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
The GDN and GDP guards carry no comment at all; one line matching their
register is enough. The reasoning lives in the commit that added it.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Same over-writing as the flex one, two lines away. The guard note drops to one
line; the pytestmark note keeps its point, which is non-obvious enough to stop
someone hoisting it to module level and silently skipping the ONNX regression
on the hardware that can still hit its bug.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Four trims, chosen by what they cost rather than by what they cover. The
distributed cases dominate: 27 of them took about as long as the 78 numerics
cases combined.

- Sequence length 4096 to 2048 on two CP models. Attention is O(s^2) and the
  ring does the same work per step whatever the length, so this exercises every
  path at a quarter the cost. 2048 already divides cp_size * 2 for both two and
  four ranks, and cp_hd512_2 has been running at it all along.
- a2a+p2p drops from six cases to one, in its own test. It needs four ranks and
  dispatches to the same AttnFuncWithCPAndKVP2P as plain p2p with an a2a stage
  either side, so one case says the composition works and six pay four-rank
  prices to re-cover p2p.
- fp16 in the forward runs on two shapes rather than all four. What fp16 risks
  that bf16 does not is exponent range, and the backward already runs both
  dtypes on every shape it covers.
- The sliding-window matrix keeps (128, 0) and (0, 0) and drops (256, 0). The
  degenerate diagonal-only window is where a band off-by-one shows; a second
  ordinary width repeats the first.

Also removes test_frost_sliding_window_selection_by_cp_comm_type as asked. It
launched no kernel, so this is coverage given up rather than time saved: it was
the only end-to-end get_attention_backend call in the suite and the only place
FROST, a sliding window and context parallelism met.

27 distributed cases become 22, and 78 numerics become 60.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Removed in a88556b during the CI-cost review, but it launches no kernel: it
drives get_attention_backend and asserts on the chosen sub-backend. Removing it
gave up coverage for no time back, which is the opposite of what that review
was for.

It is the only end-to-end get_attention_backend call in this file -- everything
else calls is_frost_attention_supported directly and skips the filters above it
-- and the only place FROST, a sliding window and context parallelism meet,
since none of the CP model configs carry a window.

Restored byte-identical to the version removed.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
This reverts commit c01f48b.

Reviewer call: the behaviour is obvious enough not to pin. It is also not FROST
code -- frost_attention.py contains no occurrence of cp_comm_type, and the
decline comes from the generic FusedAttention window filter in utils.py, so the
three negative rows re-test logic owned elsewhere.

What goes with it, for the record: the only end-to-end get_attention_backend
call in this file, and the only coverage of FROST with a sliding window under
context parallelism, since no CP model config carries a window.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Measured rather than predicted this time. Halving it on two models and dropping
five cases took the CP suite from 209.09 s / 27 cases to 171.35 s / 22. Per case
that is 7.74 s before and 7.79 s after, and 209.09 * 22/27 predicts 170.4 s, so
the entire saving came from the case count and the sequence length contributed
nothing.

The reasoning that made it look worthwhile was that attention is O(s^2), which
is true and irrelevant here: a distributed case is dominated by pool spawn, plan
JIT and NCCL setup, and the attention disappears into them. So it traded the
4096 that the fused and flash CP configs beside it use for no measurable return.

Comment records the measurement, so the next person reaching for this lever can
see it has been tried.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
The measurement belongs in the commit that made it, not beside the configs.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Back to the comment as it was. The divisibility constraint is the only thing a
reader needs here.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Two earlier changes interacted badly. a2a+p2p went from six cases to one, aimed
at cp_hd512_0, and then cp_hd512_0 went back to 4096. That is the heaviest
four-rank case in the suite and the exact configuration that hit the 90 s pool
timeout once already, now with no other a2a+p2p case to fall back on.

cp_hd512_2 is the same causal d512 shape at half the length, divides correctly
for the a2a subgroup, and ran this comm type in every green run before the trim.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Comment thread transformer_engine/pytorch/attention/dot_product_attention/backends.py Outdated
Comment thread transformer_engine/pytorch/attention/dot_product_attention/backends.py Outdated
Comment thread transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py Outdated
Comment thread transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py Outdated
Comment thread transformer_engine/pytorch/attention/dot_product_attention/utils.py Outdated
Comment thread transformer_engine/pytorch/attention/dot_product_attention/utils.py Outdated
Comment thread transformer_engine/pytorch/attention/dot_product_attention/utils.py Outdated
nvegesna-netizen and others added 3 commits October 8, 2026 14:41
Move the SDPA forward and backward graph construction into cudnn_pygraph as
build_fwd/build_bwd, and call them from both flex_attention and
frost_attention. Each backend passes tensor descriptors, any auxiliary runtime
tensors a callback reads, and its extra sdpa arguments, so a mask, a score_mod
and the engine pin or bar stay with the backend that wants them. The caching
was already shared through cudnn_pygraph.cached_graph; this closes the other
half. Drops _build_cudnn_pygraph, _bhsd_graph_tensor and
_make_cudnn_graph_tensor_dict from flex, which no longer have callers.

Traced the cuDNN call sequence of every builder before and after against a
recording stub: flex identical on all four paths, frost forward identical,
frost backward identical apart from set_stride and set_data_type swapping
order on the three gradients, which are independent setters applied before
graph.validate().

Also from review:

- backends.py: the backward sub-backend is ctx_attrs["fused_attention_backend"]
  with no conditional. The local is bound once from args and reassigned only
  under `if fp8`, and the backend enum has no other members, so the condition
  drew a distinction that does not exist.
- frost_attention.py: check the head_dim range before symmetry, so the symmetry
  rule only applies where cuDNN has no asymmetric backward plan. cuDNN does
  plan asymmetric pairs below that range.
- utils.py: separate the C++ and FROST reject reasons into two sentences
  instead of running them together.
- Rewrite both module docstrings and drop the comments that read as
  review-time rationale.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
_check_layout required rank 4 and a unit head-dim stride. cuDNN requires
exactly the same two things, for exactly the engines this backend pins, so the
check was a duplicate rather than a guard.

In cudnn-frontend 1.29.0 both FROST engine families register a validator:

  engines/manifest.py:222  frost_sdpa_fwd  -> _sdpa_validate.validate_graph
  engines/manifest.py:236  frost_sdpa_bwd  -> _sdpa_validate.validate_graph

validate_graph reaches _check_dim_stride (_sdpa_validate.py:67-76), which
raises ValueError when dim or stride is not rank 4 and cudnnGraphNotSupported
when stride[3] != 1. It runs inside graph.validate() (_pygraph.py:843), which
is the first call finalize_plans makes, so a violating tensor fails loudly
before any plan is proposed and names the offending port and the rule.

Rank was doubly covered: dot_product_attention.py:2432 already asserts 4D for
the formats FROST accepts, and thd is declined at selection. For o and d_o the
call sites ran immediately after .contiguous(), which guarantees a unit last
stride on its own.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Comment thread transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py Outdated
nvegesna-netizen and others added 4 commits October 8, 2026 15:48
Removing _check_layout dropped two guards, and only one of them was covered
elsewhere. The head-dim stride reaches cuDNN's descriptor, so _sdpa_validate
rejects a non-unit one in graph.validate(); that half stays removed. The rank
does not reach it: _check_dim_stride inspects the cuDNN graph tensor, whose
dim list bhsd_dim_stride builds as exactly four elements. cuDNN validates the
descriptor, never the torch tensor.

bhsd_dim_stride reads dims 0-3, so a tensor of rank > 4 is described as 4D
with its trailing dims silently dropped instead of rejected. Upstream only
covers the DPA path (dot_product_attention.py:2432 asserts 4D for sbhd/bshd);
fused_attn_fwd is exported and callable directly.

o and d_o need nothing: their shape is compared for equality against the
4-element o_shape, which rejects any other rank already.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
The comment attributed the symmetric-head_dim rule to a B200 measurement and
said cuDNN plans asymmetric pairs below this range. The second half is wrong
on this arch, and the capability rows say it better than the measurement did.

cudnn/sdpa/bwd/engines.py declares dqk_ge_dv per engine, which cuDNN reads as
"serves rectangular head dims with d_qk >= d_v"; unset means d_qk == d_v is
required. sdpa_bwd_sm100 leaves it unset and is the only f16 FROST backward on
SM100/SM103, so rectangular pairs are unavailable at every head dim here, not
merely inside (256, 512]. The 192/128 case belongs to the sm80 and sm120 rows.

Also noted that the head-dim constants mirror that engine's declared envelope:
d_envelope_floor=256, d={512}, d_pad_multiple=8 against 257 / 512 / 8.

No logic change.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
…ng dims

Same hole the q/k/v rank check closes. The stats node is declared [b, h, s, 1],
but fused_attn_bwd compared only shape[:3] and unsqueezed only at dim() == 3,
so an lse of [b, h, s, 4] passed both, was bound by pointer and read with
strides that do not describe it. Not reachable through TE, where the aux
tensor comes from this module's own forward, but fused_attn_bwd is exported.

Also two comment corrections:

- The symmetry rule belongs to the backward engine, not to the head-dim range.
  sdpa_bwd_sm100 leaves dqk_ge_dv unset, so cuDNN requires d_qk == d_v at every
  head dim, and sdpa_fwd_prefill_sm100 lists (192, 128) among its native
  d_shapes. The forward serves rectangular pairs here; only the backward binds,
  which is why both directions are declined.
- The rank rationale now names the sharp edge: the plan cache key carries
  strides but not rank, so a rank-5 tensor can reuse a valid 4D plan rather
  than merely losing its trailing dims.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Eleven lines down to six, and closer to the register of its neighbours, which
run to a single line each ("cuDNN-backed Flex Attention helpers.", "Attention
Backends.", "Context Parallelism.").

Dropped the two details it carried, since both already live where they apply:
import_cudnn_frontend explains the process-wide engine switch, and handle_for
explains the per-device handle and why it is keyed by backend. Also dropped
the note about where the module's state lives, which described the refactor
rather than the code.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

2.21 attention community-contribution PRs from external contributor outside the core maintainers, representing community-driven work.

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants