Repository navigation
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
Open
feat(attention): cuDNN FROST attention backend for head_dim in (256, 512], with context parallelism#3527nvegesna-netizen wants to merge 101 commits into
nvegesna-netizen wants to merge 101 commits into
Conversation
nvegesna-netizen
marked this pull request as ready for review
September 16, 2026 04:35
Contributor
|
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>
for more information, see https://pre-commit.ci
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
force-pushed
the
nvegesna/te-frost-d512-cp
branch
from
September 16, 2026 06:22
7d453a4 to
4069bfa
Compare
…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>
…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>
…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>
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>
cyanguwa
reviewed
Oct 8, 2026
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>
for more information, see https://pre-commit.ci
_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>
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
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Problem
get_attention_backendselects no backend at all for symmetrichead_dim=512with contextparallelism, so that configuration raises rather than running:
head_dim256UnfusedDotProductAttentionserves 512 but is disabled under CPHybrid models such as Gemma 4, which interleave sliding-window layers at
head_dim256 withglobal 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-levelbackend:
frost_attention.pyholds the kernel wrapper, its plan cache, andfused_attn_fwd/bwdbehindthe
cpp_extensions.fused_attnsignaturescpp_extensions/fused_attn.pyroutes to those two functions when the sub-backend is FROST,before any C++ call
utils.pyconsultsis_frost_attention_supportedinside_get_fused_attn_backend, where theC++ selector returns
No_Backend, and checks availability once at the end ofget_attention_backend, the way flash-attn's version is checkedSo
FusedAttnFuncandattn_forward_func_with_cpreach these kernels without knowing whichsub-backend they got, and
dot_product_attention.pyneeds 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_modandsoftcapfilters;and FP8, since
FusedAttnFuncforces the FP8 sub-backend on that path. The all-gather path'scausal -> causal_bottom_rightrewrite applies to FROST for free, which is what gives it thebottom-right band it needs on a trimmed KV range.
Verification
Full suite on B200 against this head, in a container with
nvidia-cudnn-frontend1.29.0 andnvidia-cutlass-dsl4.8.0:flex_attention.pyregressiontest_attention.pysuiteThe broad suite is identical to the pre-change baseline and
flex_attention.py's own suite isclean, 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 sothe 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:p2pall_gathera2aPlus 2 nodes x 2 ranks for every comm type including
a2a+p2p, with aFusedAttentioncontrolpassing throughout.
a2a+p2pneeds the multi-node arm specifically: atworld_size4 the a2alevel 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.pyis non-distributed and checks the forward, the LSE convention and thebackward 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.pycompares a CP run against a non-CP run ofthe 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-py3shipsnvidia-cudnn-frontend1.26.0 and
nvidia-cutlass-dsl4.6.2, both below what FROST needs, sois_frost_attention_available()declines and the tests skip with that reason rather than failing.qa/L3_pytorch_FA_versions_test/test.shtargets sm100+ but pinsnvidia-cutlass-dsl[cu13]==4.4.2for the FA4 path, so it would skip even on a B200; and
nvidia-cutlass-dslis not a declared TEdependency at all, arriving transitively via
flash-attn-4. Happy to wire this into whicheverlane you consider the right home.
Performance
Against
UnfusedDotProductAttentionat the same shape (b2 hq8 hkv4 d512 causal bf16, CP=1,single layer), the only other backend serving this head dim from released components:
The unfused path materialises the full
s x sscore matrix, so it grows O(s^2). Note that inforward+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 thesame root cause: the selected sub-backend was being discarded and re-derived as
F16_arbitrary_seqlenin seven places, so those places now keep it instead.context_parallel.py(+29/-7): each of the three autograd classes recomputed the sub-backendlocally rather than receiving it.
attn_forward_func_with_cpgains one argument, threadedthrough, with the existing derivation kept as the fallback. One further line: the p2p forward
step returned
max_logitwith a starred tail, which is only unambiguous while no backend has astatically 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 functionactually returns.
backends.py(+5/-5): the non-FP8 backward hardcodedF16_arbitrary_seqlen, discarding thevalue the forward had already saved, so it now simply reads
ctx_attrs["fused_attention_backend"]. The CP assert admitted only that sub-backend.FusedAttnBackend.FROSThas no pybind counterpart. It is3, past the C++ enum, and theimport-time sync assert now exempts python-only members. Nothing in
transformer_engine/pytorchconstructs
NVTE_Fused_Attn_Backend(int), and the dispatch happens before any C++ call, so thevalue never reaches pybind. Every post-selection check in
utils.pyis guarded on== FP8or== F16_arbitrary_seqlen, so a FROST value falls through them.is_frost_attention_supporteddeliberately does not probe availability. Probing imports cuDNNFrontend 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_dim512. Availability is checked once at the end, beside the flash-attnversion checks.
Plan selection is deliberately strict.
heur_mode.Awith an explicit name-checkedselect_plan, rather thanA|FALLBACKwithHEURISTICS_CHOICE. Without a pin,build_planswalks 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_supportfatal rather than advisory, which is why it follows the selection.Version constraint worth knowing. FROST enforces
nvidia-cutlass-dsl >= 4.7.0at plan-buildtime while
cudnn-frontendonly declares>= 4.6.2. Below that floor every FROST engine silentlydeclines 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-44.0.0b11, which CI currently pins.Scope
SM100/SM103 only (the cuDNN d512 backward is Blackwell-only), bf16/fp16, symmetric
head_dimin(256, 512],
bshdandsbhd, mask typesno_mask/causal/causal_bottom_right.Symmetric specifically, and the constraint is the backward engine's.
vcarries its own graphnode 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, leavesdqk_ge_dvunset,which cuDNN reads as requiring
d_qk == d_vat every head dim rather than only above 256. Theforward is the opposite:
sdpa_fwd_prefill_sm100lists(192, 128)among its natived_shapesand has no
dqk_ge_dvconcept at all. So the decline covers both directions rather than trainingalone, since
is_trainingismodule.trainingandeval()does not disable autograd, which makesa 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}andd_pad_multiple=8, which is_MIN_HEAD_DIM257,_MAX_HEAD_DIM512 and_HEAD_DIM_MULTIPLE8 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 graphcapture, and deterministic execution.
thdis 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 totop-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_gatherora2a. It is declined withp2panda2a+p2p, whose ring shards KV across steps so a bound measured against the full sequence doesnot survive the per-step tiles, which is the same rule
FusedAttentioncarries. That matters forthe 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 d512causal bf16, single GPU, peak allocated for forward + backward: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 to74 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_sizeThe 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:
bshdstrided viewbhsdcontiguousThe 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.pyextraction is done.frost_attention.pyandflex_attention.pynowdrive 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.pycomes 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.pycarries its own copy, shares no code witheither, and cannot consume a torch-dependent module.
The graph builders are shared as
cudnn_pygraph.build_fwdandbuild_bwd. Each backend passestensor descriptors, any auxiliary tensors its callback reads, and its own
sdpaarguments, so amask, 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_extensionsshims.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 onB200: a FROST plan ranks at index 0 and an unpinned build selects it at
head_dim64, 128, 256 and512, 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_plansnow bars them unless the caller pins a plan by name, which FROST does and flex doesnot. 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_enginesmarks rather than removes,measured on B200 at
head_dim64, 128 and 256: the plan count is unchanged and index 0 still namesthe barred engine after
build_plans, so a post-build assertion on the selected plan would firefalsely. What verifies the bar is
test_frost_switch_does_not_change_what_flex_computes, which runsthe score_mod path with the engines on and off and compares.
NVTE_FLEX_TEST_REQUIREDexists so alane 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 inits docstring rather than becoming permanent by default.
One handle is shared across overlapped CP streams. The
wait_streamserialization incontext_parallel.pyexists for FA3/FA4's internal per-call workspace, and its comment saysFusedAttention 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 whilebeing deliberately left overlapped, so I have not added FROST to that guard.
Dependency floors are intentionally not enforced by the build. Bumping
nvidia-cudnn-frontendto 1.29.0 would forcenvidia-cutlass-dslandapache-tvm-ffiinto everyTE-PyTorch install, because 1.29.0 drops the
cutedslextra marker, and it still would notguarantee 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_REQUIREDandNVTE_FLEX_TEST_REQUIREDmirrorNVTE_GDN_TEST_REQUIRED, so a lanemeant 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_ATTNis indocs/envvars.rst, described as a FusedAttentionsub-backend. The backend tables in
docs/examples/attention/attention.ipynbare not updated yet.