Repository navigation
feat(attention): cuDNN FROST attention backend for head_dim in (256, 512], with context parallelism - #3527
feat(attention): cuDNN FROST attention backend for head_dim in (256, 512], with context parallelism#3527nvegesna-netizen wants to merge 108 commits into
Conversation
|
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>
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>
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>
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>
| if fused_attention_backend == FusedAttnBackend["FROST"]: | ||
| # Imported here so a process that never selects FROST never imports cuDNN Frontend. | ||
| # pylint: disable-next=import-outside-toplevel | ||
| from ..attention.dot_product_attention import frost_attention |
There was a problem hiding this comment.
Plugging frost fused_attn_fwd/bwd here might be safe for now, but it may not be the best architecture design to avoid circular imports. Need to think about this a bit more.
There was a problem hiding this comment.
Both directions are already function-local for exactly this reason: frost_attention.py:298 imports from here inside a function, and :359 / :653 here import frost inside a function. That is the established pattern elsewhere too, suppressed at seventeen sites across grouped_mlp.py, dynamo/custom_op.py, utils.py and ops/_common.py.
There is a clean fix available. Everything frost imports from here is pure data defined above the functions: TORCH_DType, QKVFormat, QKVLayout, AttnBiasType, AttnMaskType, SoftmaxType and FusedAttnBackend, all depending only on torch, tex and ..constants. Moving them into a types module and re-exporting from here would break the cycle rather than defer it, and none of the thirteen places importing these names would change.
Leaving this open for you to decide whether that is worth doing, and whether here or separately.
| :Type: ``int`` (0 or 1) | ||
| :Default: ``1`` | ||
| :Description: Enable or disable the FROST sub-backend of FusedAttention for DotProductAttention. FROST wraps the cuDNN FROST CuTe-DSL SDPA kernels through the cuDNN Frontend python API, rather than the C++ fused-attention path the other sub-backends use. **It is experimental and subject to change**, as the underlying cuDNN FROST engines are. When set to ``0``, FROST will not be used. It is selected only where the cuDNN sub-backends decline and is the only released backend serving symmetric ``head_dim`` in (256, 512] together with context parallelism. It is limited to SM100/SM103 with BF16/FP16 inputs, a ``head_dim`` that is a multiple of 8, and ``nvidia-cudnn-frontend>=1.29.0`` and ``nvidia-cutlass-dsl>=4.7.0`` installed. It declines FP8, ``thd`` layouts, dropout, attention bias, KV caching, ``max_logit``, CUDA graph capture, and deterministic execution, the last because cuDNN offers no deterministic backward for these kernels. | ||
|
|
There was a problem hiding this comment.
In general, we/TE wants to avoid adding more environment variables. How necessary is this?
I also discovered that cudnn-frontend has an opt-in mechanism for some kernels, whose accessibility is controlled by CUDNN_FRONTEND_ENABLE_FROST_ENGINES. But for some non-opt-in kernels, they'll be discoverable anyway - not sure how that changes our strategy here. We could possibly use NVTE_FROST_ATTN to control whether we run is_frost_attention_supported at all (see utils.py).
All this text can go into the header docstring in frost_attention.py. (What I suggested in the last iteration wasn't set in stone by the way; please feel free to edit - but please make it succinct. Thanks!)
There was a problem hiding this comment.
The closest analogue argues your way: FusedAttention's other sub-backend, FP8, has no on/off env variable at all. NVTE_FLASH_ATTN_V2/V3/V4 are versions of a different backend rather than sub-backends of one, so they are not really comparable.
What I think still justifies it is narrower. FP8 is opted into through a recipe, so a user who does not want it simply does not ask for it. FROST is selected automatically whenever the C++ sub-backends decline, so without this variable the only ways to rule it out are NVTE_FUSED_ATTN=0, which disables the C++ paths as well, and CUDNN_FRONTEND_ENABLE_FROST_ENGINES=0, which is cuDNN's rather than ours.
On the opt-in mechanism: for SDPA it is all or nothing. Every FROST SDPA slot is opt_in=True, nine forward and four backward (engines/manifest.py:210-233), and offered_ids() filters on enabled or not s.opt_in. The non-opt-in FROST engines belong to the GDN, KDA and GDP families, which the attention path never builds graphs for.
Kept for now, but this is a policy call rather than a technical one, so leaving it to you.
…ing 3 FROST sat at a literal 3, the next free value after NVTE_FP8. If a C++ enumerator ever takes that value, the python member becomes an alias of it. Deriving the value from the pybind enum's current maximum shifts FROST instead, so the two can never occupy the same slot. Measured against the existing import-time sync assert: a C++ addition at 3 is caught in all three orderings (mirrored after FROST, mirrored before it, not mirrored at all), so the clash was never silent. But in one of them FusedAttnBackend.FROST resolves to the new member rather than to FROST, which is only harmless because the assert stops the module loading first. Deriving removes the case entirely. Note that FP8 + 1 would not have helped: it is also 3, before and after. Docstring reworded to describe the enum as what it is, the C++ backends mirrored value-for-value plus one python-only member, rather than claiming it mirrors every member of the C++ enum. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com> Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
The backward test already runs a forward to get out and the LSE, so the forward assertions move there and test_frost_forward_matches_reference goes, along with its case table. 42 cases become 24. One reference evaluation now serves both halves. It is built on tensors that require grad, so its outputs are the exact forward answer and it is also what autograd differentiates for the gradients. Only the lossy evaluation that measures the error floor is extra, so this is two float64 passes per case where the two tests previously cost three between them. Coverage moves rather than only shrinking. Lost: causal_bottom_right and the sq != skv shape, for the unwindowed forward and the LSE; both are still covered for forward output by the sliding-window test, which parametrises causal_bottom_right and a rectangular shape. Gained: the LSE is now checked with a window applied, which the forward test never did. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com> Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
Six sites defaulted to F16_arbitrary_seqlen when the new parameter was None. get_attention_backend has already picked the backend and backends.py passes it on every call that sets use_fused_attention, so the fallback never fired, and silently guessing F16_arbitrary_seqlen would be the wrong backend for a d512 config anyway. Assign it directly instead, and make the contract explicit once at the public entry point rather than six times at the use sites. Without that guard a caller omitting the argument would reach FusedAttnBackend.cast(None) and get "int() argument must be ... not 'NoneType'" instead of a usable message. In-repo behaviour is unchanged: the only two callers are backends.py:1227, which sets neither argument and so cannot reach these lines, and backends.py:2616, which always passes both. Net -6 lines on the file. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com> Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
The fold left no unwindowed causal_bottom_right case checking output and LSE. The folded test parametrises no_mask and causal on square shapes, and the sliding-window test covers causal_bottom_right on a rectangular shape but always with a window and discards the LSE. That gap matters more than the rest of what the fold dropped: top-left and bottom-right coincide when sq == skv, so a swapped anchor is invisible on every other shape, and the LSE is what the context-parallel correction consumes. One targeted case on _SHAPES[2] (sq 256, skv 512) covers it, so the suite is 25 rather than back to 42. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com> Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
…lpers Comment and docstring trims across frost_attention, cudnn_pygraph, envvars.rst and the three test files. FROST is a sub-backend, so it comes out of the top-level backend-ordering paragraph in envvars.rst. Two helpers move into cudnn_pygraph, where flex can reach them: pkg_version, which reads a package version and tolerates absence, and diagonal_band_kwargs, which turns a TE mask type and window into cuDNN's diagonal alignment plus band. flex_attention's _BACKEND_NAME becomes FlexAttention, matching FrostAttention and TE's own class names. Both test guards drop their CUDA and compute-capability pre-checks, which is_frost_attention_available already makes; it imports with the FROST engines off, so probing has no process-wide side effect to avoid. test_frost_serves_an_unambiguous_window_on_a_non_causal_mask folds into test_frost_declines_unsupported_configs, which already carries a positive assertion. test_flex_attention's switch test reuses _score_mod_post_scale_bias rather than defining an identical score_mod inline. utils.py labels the two halves of a combined reject message, so it is clear which selector produced which. 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
… backward return
Two shape/arity mismatches with cpp_extensions.fused_attn, both of which only
worked because every consumer happened to be tolerant.
The forward squeezed the LSE to [b, h, s]. The C++ fused_attn_fwd returns
[b, h, s, 1], and context_parallel.py:2152 squeezes it itself, with the comment
"[b, h, sq, 1] -> [b, h, sq]". Against our pre-squeezed tensor that squeeze was
a no-op, so both paths happened to end at [b, h, s]. At seqlen 1 they do not:
[b, h, 1] loses the sequence dimension and becomes [b, h]. Return the LSE
unsqueezed and let CP do what it already does.
The backward returned four values. The C++ one returns five,
{dQ, dK, dV, dBias, dSoftmaxOffset} (csrc/extensions/attention.cpp:646), which
is also the contract backends.py:2091 documents. Four worked only because every
caller either stars the tail or stops at dbias. Return five.
FROST declines both bias and non-vanilla softmax, so the two extra slots are
None.
Test updates follow from the shapes: the float64 reference LSE is [b, h, s], so
the comparisons squeeze before subtracting rather than broadcasting to
[b, h, s, s], and the _bwd helper stars the tail.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Nitin Vegesna <nvegesna@nvidia.com>
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.FP8 forward with a FROST backward is not supported, and cannot arise: FP8 serves
head_dimup to 256 and FROST only above it, so their ranges are disjoint.
_fused_attn_setup_ctxthereforehands the backward whatever the forward recorded, which is FP8 for an FP8 forward and FROST for a
FROST one. Supporting a mixed pair would need
get_attention_backendto return a backend perdirection, which is out of scope 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.