From 7765b76555034aec27fb5b8a719240aaabbaf350 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 15 Sep 2026 20:58:54 -0700 Subject: [PATCH 01/97] feat(attention): cuDNN FROST kernel wrapper for head_dim in (256, 512] 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/frost_attention.py | 486 ++++++++++++++++++ 1 file changed, 486 insertions(+) create mode 100644 transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py new file mode 100644 index 00000000000..e7b2e2c0545 --- /dev/null +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -0,0 +1,486 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""cuDNN FROST attention backend for head_dim in (256, 512] on SM100/SM103. + +Why this exists. Gemma-4 global layers use symmetric head_dim=512, and no backend TE can select +today serves both that head dim and context parallelism: FlashAttention 2/3 cap at 256, FA4 is +gated off at symmetric 512, the C++ cuDNN fused path caps at 256, and the unfused path supports +512 but cannot do CP. cuDNN Frontend 1.29.0 ships CuTe-DSL ("FROST") SDPA kernels that do serve +symmetric 512 forward and backward on Blackwell, reachable through the ordinary cuDNN graph API. +This module wraps them so TE, including its CP ring, can dispatch to them. + +Three properties were measured on B200 before this was written, and each one constrains the code: + +1. cuDNN's `use_causal_mask` is TOP-LEFT aligned and `use_causal_mask_bottom_right` is + bottom-right; both were verified against references at SQ=1024/SKV=2048, where the two + disagree by three orders of magnitude (1.6e-03 vs 3.5e+00). They coincide when SQ == SKV, so + the distinction is invisible in square tests and decisive for all_gather, which trims KV. + `_MASK_MODES` lists only spellings checked this way: sdpa() ignores unknown kwargs silently, + so an unverified name would apply no mask at all and still run. + +2. Plan building must be cached. Building a plan costs ~1972 ms the first time and ~12 ms once + cuDNN has cached the JIT, against a ~0.129 ms execute. Even the cached rebuild is ~90x an + execute, so a per-call build would make training build-bound. Hence `_PLAN_CACHE`. + +3. The forward LSE is natural-log logsumexp in fp32, shaped [b, h, s, 1]. Squeezed to [b, h, s] + it is exactly what the CP ring correction in context_parallel.py consumes (max err 1.8e-06 vs + an fp64 reference), which is what makes ring attention over these kernels valid at all. + +Numerics were validated against the criterion FlashAttention applies to itself, namely that the +kernel error must stay within 2x the error bf16 inputs alone produce: observed 0.21x to 0.62x +across square and rectangular, causal and non-causal shapes. +""" + +from __future__ import annotations + +import os +from typing import Optional, Tuple + +import torch + +__all__ = [ + "is_frost_attention_available", + "is_frost_attention_supported", + "frost_attn_fwd", + "frost_attn_bwd", + "to_frost_layout", + "from_frost_layout", +] + + +# FROST engines are opt-in inside cuDNN Frontend, and they additionally require a newer +# nvidia-cutlass-dsl than cudnn-frontend itself declares. cudnn-frontend requires >= 4.6.2 while +# FROST enforces >= 4.7.0 at plan-build time; with 4.6.2 installed every FROST engine silently +# declines and ordinary cuDNN backend plans are returned with no error at all. We therefore check +# the selected plan by NAME rather than trusting that the engine was used. +_FROST_FWD_PLAN_TOKEN = "sdpa_fwd_prefill_sm100" +_FROST_BWD_PLAN_TOKEN = "sdpa_bwd_sm100" +_MIN_CUTLASS_DSL = (4, 7, 0) + +_SUPPORTED_ARCHS = ((10, 0), (10, 3)) +_MAX_HEAD_DIM = 512 +_MIN_HEAD_DIM = 257 # below this the existing cuDNN/flash backends already serve the shape + +_cudnn = None +_availability: Optional[Tuple[bool, str]] = None +_PLAN_CACHE: dict = {} + + +def _import_cudnn(): + """Import cuDNN Frontend with FROST engines enabled, once.""" + global _cudnn + if _cudnn is None: + # Must be set before the import: the engines are registered at import time. + os.environ.setdefault("CUDNN_FRONTEND_ENABLE_FROST_ENGINES", "1") + import cudnn # pylint: disable=import-outside-toplevel + import cudnn.sdpa # noqa: F401 pylint: disable=import-outside-toplevel,unused-import + + _cudnn = cudnn + return _cudnn + + +def is_frost_attention_available() -> Tuple[bool, str]: + """Whether the FROST kernels can be used at all, with a reason when they cannot. + + Cached, because this is consulted on every backend-selection call. + """ + global _availability + if _availability is not None: + return _availability + + def _no(reason): + global _availability + _availability = (False, reason) + return _availability + + if not torch.cuda.is_available(): + return _no("no CUDA device") + if torch.cuda.get_device_capability() not in _SUPPORTED_ARCHS: + return _no( + "cuDNN FROST head_dim>256 kernels are SM100/SM103 only; found sm%d%d" + % torch.cuda.get_device_capability() + ) + try: + _import_cudnn() + except ImportError as exc: + return _no("nvidia-cudnn-frontend not importable: %s" % exc) + + from importlib.metadata import PackageNotFoundError, version + + try: + raw = version("nvidia-cutlass-dsl") + except PackageNotFoundError: + return _no("nvidia-cutlass-dsl not installed (FROST requires >= 4.7.0)") + try: + parsed = tuple(int(p) for p in raw.split(".")[:3]) + except ValueError: + parsed = (0, 0, 0) + if parsed < _MIN_CUTLASS_DSL: + # Worth being loud: this combination fails by silently declining, not by raising. + return _no( + "nvidia-cutlass-dsl %s is below the FROST floor 4.7.0; FROST engines would be" + " silently skipped in favour of ordinary cuDNN backend plans" % raw + ) + + _availability = (True, "") + return _availability + + +# cuDNN sdpa() kwargs per TE mask type. +# +# These exact spellings are behaviourally verified, which matters more than it sounds: sdpa() +# takes **kwargs and SILENTLY IGNORES names it does not recognise, so a typo here would apply no +# mask at all and still build and run. Do not add an entry without checking the output against a +# reference for that alignment. +# +# Both alignments are needed. The p2p ring produces square diagonal tiles (top-left and +# bottom-right coincide there), while all_gather trims KV and relies on bottom-right alignment, +# where the two differ completely. +_MASK_MODES = { + "no_mask": {}, + "causal": {"use_causal_mask": True}, + "causal_bottom_right": {"use_causal_mask_bottom_right": True}, +} + + +def _mask_mode(attn_mask_type: str) -> str: + """Validate a TE mask type and return its key in _MASK_MODES. + + Anything not listed is rejected rather than approximated: the failure mode of guessing wrong + is silent numerical corruption, not an exception. + """ + if attn_mask_type in _MASK_MODES: + return attn_mask_type + raise NotImplementedError( + "FROST attention supports attn_mask_type in %s; got %r. Padding variants need varlen" + " support that is not implemented here." % (sorted(_MASK_MODES), attn_mask_type) + ) + + +def is_frost_attention_supported( + head_dim_qk: int, + head_dim_v: int, + qkv_dtype: torch.dtype, + attn_mask_type: str, + dropout: float = 0.0, + attn_bias_type: str = "no_bias", +) -> Tuple[bool, str]: + """Whether this specific attention configuration should route to FROST.""" + ok, reason = is_frost_attention_available() + if not ok: + return False, reason + if head_dim_qk != head_dim_v: + return False, "FROST path requires symmetric head_dim; got %d/%d" % ( + head_dim_qk, + head_dim_v, + ) + if not _MIN_HEAD_DIM <= head_dim_qk <= _MAX_HEAD_DIM: + return False, "FROST path covers head_dim in (256, 512]; got %d" % head_dim_qk + if qkv_dtype not in (torch.bfloat16, torch.float16): + return False, "FROST path supports bf16/fp16; got %s" % qkv_dtype + if dropout != 0.0: + return False, "FROST path does not support dropout" + if attn_bias_type != "no_bias": + return False, "FROST path does not support attention bias" + try: + _mask_mode(attn_mask_type) + except NotImplementedError as exc: + return False, str(exc) + return True, "" + + +def to_frost_layout(t: torch.Tensor, qkv_format: str) -> torch.Tensor: + """View a tensor in TE's qkv_format as [b, h, s, d]. + + No copy: the cuDNN graphs are built from each tensor's actual strides, so both bshd and + sbhd are served directly. sbhd matters because that is what Megatron uses internally, and + transposing into bshd on every call would copy the whole tensor. + """ + if qkv_format == "bshd": # [b, s, h, d] -> [b, h, s, d] + return t.permute(0, 2, 1, 3) + if qkv_format == "sbhd": # [s, b, h, d] -> [b, h, s, d] + return t.permute(1, 2, 0, 3) + raise NotImplementedError( + "FROST attention supports qkv_format 'bshd' and 'sbhd'; got %r." + " thd needs varlen support that is not implemented here." % qkv_format + ) + + +def from_frost_layout(t: torch.Tensor, qkv_format: str) -> torch.Tensor: + """Inverse of to_frost_layout.""" + if qkv_format == "bshd": # [b, h, s, d] -> [b, s, h, d] + return t.permute(0, 2, 1, 3) + if qkv_format == "sbhd": # [b, h, s, d] -> [s, b, h, d] + return t.permute(2, 0, 1, 3) + raise NotImplementedError( + "FROST attention supports qkv_format 'bshd' and 'sbhd'; got %r." % qkv_format + ) + + +def _cudnn_dtype(dtype: torch.dtype): + cudnn = _import_cudnn() + return { + torch.bfloat16: cudnn.data_type.BFLOAT16, + torch.float16: cudnn.data_type.HALF, + }[dtype] + + +def _check_layout(name: str, t: torch.Tensor) -> None: + """Validate a [b, h, s, d] view. + + The graphs are built from each tensor's ACTUAL strides rather than one fixed layout, so bshd + and sbhd are both served without a transpose. The only hard requirement is that the head + dimension is contiguous, which the kernels assume. + """ + if t.dim() != 4: + raise ValueError("%s must be 4D [b, h, s, d]; got %s" % (name, tuple(t.shape))) + if t.stride(3) != 1: + raise ValueError( + "%s must have a contiguous head dimension; got shape %s stride %s" + % (name, tuple(t.shape), tuple(t.stride())) + ) + + +def _select_frost_plan(graph, token: str, what: str): + """Select a plan whose name proves a FROST engine was chosen. + + Falling back to whatever plan happens to be first would defeat the purpose: at these head + dims the non-FROST plans do not exist, so an unnoticed fallback would either fail obscurely + or quietly serve a different shape. + """ + cudnn = _import_cudnn() + graph.create_execution_plans([cudnn.heur_mode.A]) + names = [graph.get_plan_name_at_index(i) for i in range(graph.get_execution_plan_count())] + hits = [i for i, n in enumerate(names) if token in n] + if not hits: + from importlib.metadata import version + + raise RuntimeError( + "no cuDNN FROST %s engine was offered (looked for %r). Candidate plans: %s." + " nvidia-cutlass-dsl=%s (FROST floor 4.7.0)." + % (what, token, names[:6], version("nvidia-cutlass-dsl")) + ) + graph.select_plan(hits[0]) + graph.check_support() + graph.build_plans() + return names[hits[0]] + + +def _build_fwd(key) -> dict: + """Build (and JIT-compile) a forward graph. Expensive; always reached through the cache.""" + cudnn = _import_cudnn() + b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks = key + io_dt = _cudnn_dtype(dtype) + shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] + + graph = cudnn.pygraph( + io_data_type=io_dt, + intermediate_data_type=cudnn.data_type.FLOAT, + compute_data_type=cudnn.data_type.FLOAT, + ) + tq = graph.tensor(name="q", dim=shq, stride=list(qs)) + tk = graph.tensor(name="k", dim=shkv, stride=list(ks)) + tv = graph.tensor(name="v", dim=shkv, stride=list(ks)) + tout, tlse = graph.sdpa( + name="frost_fwd", + q=tq, + k=tk, + v=tv, + generate_stats=True, # the CP ring needs the LSE, and it is cheap + attn_scale=scale, + **_MASK_MODES[mask], + ) + tout.set_output(True).set_dim(shq).set_stride(list(qs)) # out mirrors q + tlse.set_output(True).set_dim([b, hq, sq, 1]).set_stride([hq * sq, sq, 1, 1]).set_data_type( + cudnn.data_type.FLOAT + ) + graph.validate() + graph.build_operation_graph() + plan = _select_frost_plan(graph, _FROST_FWD_PLAN_TOKEN, "forward") + return { + "graph": graph, + "handles": (tq, tk, tv, tout, tlse), + "workspace": max(graph.get_workspace_size(), 1), + "plan": plan, + } + + +def _build_bwd(key) -> dict: + """Build (and JIT-compile) a backward graph. Expensive; always reached through the cache.""" + cudnn = _import_cudnn() + b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks = key + io_dt = _cudnn_dtype(dtype) + shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] + + graph = cudnn.pygraph( + io_data_type=io_dt, + intermediate_data_type=cudnn.data_type.FLOAT, + compute_data_type=cudnn.data_type.FLOAT, + ) + handles = {} + # o and dO share q's layout; k, v and their grads share k's. + for name, shape, stride in ( + ("q", shq, qs), + ("k", shkv, ks), + ("v", shkv, ks), + ("o", shq, qs), + ("do", shq, qs), + ): + handles[name] = graph.tensor(name=name, dim=shape, stride=list(stride)) + handles["stats"] = graph.tensor( + name="stats", + dim=[b, hq, sq, 1], + stride=[hq * sq, sq, 1, 1], + data_type=cudnn.data_type.FLOAT, + ) + tdq, tdk, tdv = graph.sdpa_backward( + name="frost_bwd", + q=handles["q"], + k=handles["k"], + v=handles["v"], + o=handles["o"], + dO=handles["do"], + stats=handles["stats"], + attn_scale=scale, + **_MASK_MODES[mask], + ) + for tensor, stride in ((tdq, qs), (tdk, ks), (tdv, ks)): + tensor.set_output(True).set_data_type(io_dt).set_stride(list(stride)) + graph.validate() + graph.build_operation_graph() + plan = _select_frost_plan(graph, _FROST_BWD_PLAN_TOKEN, "backward") + handles["dq"], handles["dk"], handles["dv"] = tdq, tdk, tdv + return { + "graph": graph, + "handles": handles, + "workspace": max(graph.get_workspace_size(), 1), + "plan": plan, + } + + +def _cached(kind: str, key): + """Plan cache. See module docstring: a build is ~15000x an execute, so this is required.""" + cache_key = (kind,) + key + entry = _PLAN_CACHE.get(cache_key) + if entry is None: + entry = _build_fwd(key) if kind == "fwd" else _build_bwd(key) + _PLAN_CACHE[cache_key] = entry + return entry + + +def _key(q, k, mask, scale): + return ( + q.shape[0], + q.shape[1], + k.shape[1], + q.shape[2], + k.shape[2], + q.shape[3], + q.dtype, + mask, + float(scale), + # Strides are part of the plan: the graph is built for this exact layout, which is what + # lets bshd and sbhd both run without a transpose. + tuple(q.stride()), + tuple(k.stride()), + ) + + +def frost_attn_fwd( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + attn_scale: Optional[float] = None, + attn_mask_type: str = "causal", +) -> Tuple[torch.Tensor, torch.Tensor]: + """Forward attention via cuDNN FROST. + + q, k, v are [b, h, s, d] views over BSHD-contiguous memory. GQA is supported directly + (h_kv may differ from h_q) and SQ need not equal SKV, which is what lets a CP ring step + use this. Returns (out, softmax_lse) with softmax_lse as [b, h, s] fp32 natural-log + logsumexp, the layout and convention the CP ring correction expects. + """ + for name, tensor in (("q", q), ("k", k), ("v", v)): + _check_layout(name, tensor) + if k.shape != v.shape: + raise ValueError("k and v must have the same shape; got %s and %s" % (k.shape, v.shape)) + if q.shape[1] % k.shape[1] != 0: + raise ValueError( + "num_heads must be divisible by num_gqa_groups; got %d and %d" + % (q.shape[1], k.shape[1]) + ) + + mask = _mask_mode(attn_mask_type) + scale = attn_scale if attn_scale is not None else q.shape[-1] ** -0.5 + entry = _cached("fwd", _key(q, k, mask, scale)) + tq, tk, tv, tout, tlse = entry["handles"] + + b, hq, sq, _ = q.shape + # Allocate per call: the cache holds only the compiled plan, never output buffers, so that + # concurrent or nested uses cannot alias each other. empty_strided rather than empty_like: + # the latter does not preserve an arbitrary permuted stride, and the graph was built for + # q's exact strides. + out = torch.empty_strided(q.shape, q.stride(), device=q.device, dtype=q.dtype) + lse = torch.empty(b, hq, sq, 1, device=q.device, dtype=torch.float32) + workspace = torch.empty(entry["workspace"], device=q.device, dtype=torch.uint8) + entry["graph"].execute({tq: q, tk: k, tv: v, tout: out, tlse: lse}, workspace) + return out, lse.squeeze(-1) + + +def frost_attn_bwd( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + out: torch.Tensor, + softmax_lse: torch.Tensor, + dout: torch.Tensor, + attn_scale: Optional[float] = None, + attn_mask_type: str = "causal", +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Backward attention via cuDNN FROST. `softmax_lse` is [b, h, s] as returned by the forward.""" + for name, tensor in (("q", q), ("k", k), ("v", v), ("out", out), ("dout", dout)): + _check_layout(name, tensor) + + mask = _mask_mode(attn_mask_type) + scale = attn_scale if attn_scale is not None else q.shape[-1] ** -0.5 + entry = _cached("bwd", _key(q, k, mask, scale)) + h = entry["handles"] + + if softmax_lse.dim() == 3: + softmax_lse = softmax_lse.unsqueeze(-1) + softmax_lse = softmax_lse.contiguous() + + # The graph expects o and dO in q's layout. A caller may hand us either with different + # strides (dO in particular comes from autograd), so restride rather than silently reading + # the wrong elements. + def _as(t, ref): + if tuple(t.stride()) == tuple(ref.stride()): + return t + buf = torch.empty_strided(t.shape, ref.stride(), device=t.device, dtype=t.dtype) + buf.copy_(t) + return buf + + out = _as(out, q) + dout = _as(dout, q) + + dq = torch.empty_strided(q.shape, q.stride(), device=q.device, dtype=q.dtype) + dk = torch.empty_strided(k.shape, k.stride(), device=k.device, dtype=k.dtype) + dv = torch.empty_strided(v.shape, v.stride(), device=v.device, dtype=v.dtype) + workspace = torch.empty(entry["workspace"], device=q.device, dtype=torch.uint8) + entry["graph"].execute( + { + h["q"]: q, + h["k"]: k, + h["v"]: v, + h["o"]: out, + h["do"]: dout, + h["stats"]: softmax_lse, + h["dq"]: dq, + h["dk"]: dk, + h["dv"]: dv, + }, + workspace, + ) + return dq, dk, dv From 33198264281e61de893d38eb9979d2c59f6becc3 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 15 Sep 2026 20:59:09 -0700 Subject: [PATCH 02/97] feat(attention): select and dispatch FrostAttention from DotProductAttention 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 Signed-off-by: Nitin Vegesna --- tests/pytorch/test_torch_compile.py | 1 + tests/pytorch/utils.py | 2 + .../dot_product_attention/backends.py | 161 ++++++++++++++++++ .../dot_product_attention.py | 49 +++++- .../attention/dot_product_attention/utils.py | 65 +++++++ 5 files changed, 277 insertions(+), 1 deletion(-) diff --git a/tests/pytorch/test_torch_compile.py b/tests/pytorch/test_torch_compile.py index e5d7169da12..723a2ebd07a 100644 --- a/tests/pytorch/test_torch_compile.py +++ b/tests/pytorch/test_torch_compile.py @@ -1286,6 +1286,7 @@ def fn(x, params): fused_attention_backend, use_unfused_attention, _, + _, ) = dpa_utils.get_attention_backend(params) # Encode the full selection (enabled backends + fused sub-backend) in # the tensor value: without a tensor op dynamo skips the frame entirely diff --git a/tests/pytorch/utils.py b/tests/pytorch/utils.py index 6b66458985d..0571153240d 100644 --- a/tests/pytorch/utils.py +++ b/tests/pytorch/utils.py @@ -452,6 +452,7 @@ def test(): use_fused_attention, fused_attention_backend, use_unfused_attention, + _use_frost_attention, available_backends, ) = get_attention_backend(attention_params) # Check if FA3 is an available backend when num_splits != 1 @@ -465,6 +466,7 @@ def test(): _attention_backends["flash_attention_backend"] = flash_attention_backend _attention_backends["fused_attention_backend"] = fused_attention_backend _attention_backends["use_unfused_attention"] = use_unfused_attention + _attention_backends["use_frost_attention"] = _use_frost_attention _attention_backends["backend_selection_requires_update"] = False return available_backends, flash_attention_backend, fused_attention_backend diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index 9a339233a42..355da80830f 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -2286,6 +2286,167 @@ def backward(ctx, d_out, *_args): return (*_fused_attn_backward_impl(bwd_args), None) +class FrostAttnFunc(torch.autograd.Function): + """Autograd wrapper around the cuDNN FROST kernels, for the non-context-parallel path. + + The CP path does not go through here: context_parallel.py calls frost_attn_fwd/bwd per ring + step itself, because the ring has to interleave those calls with KV exchange and LSE + correction rather than treating attention as one opaque autograd node. + """ + + @staticmethod + def forward(ctx, q, k, v, softmax_scale, attn_mask_type, qkv_format, is_training): + # pylint: disable=missing-function-docstring + from .frost_attention import ( # pylint: disable=import-outside-toplevel + frost_attn_fwd, + from_frost_layout, + to_frost_layout, + ) + + # .contiguous() first: the graphs are built for BSHD-contiguous memory and + # frost_attention raises on anything else rather than computing on wrong strides. + q_f = to_frost_layout(q.contiguous(), qkv_format) + k_f = to_frost_layout(k.contiguous(), qkv_format) + v_f = to_frost_layout(v.contiguous(), qkv_format) + out_f, softmax_lse = frost_attn_fwd( + q_f, k_f, v_f, attn_scale=softmax_scale, attn_mask_type=attn_mask_type + ) + out = from_frost_layout(out_f, qkv_format) + if is_training: + ctx.save_for_backward(q_f, k_f, v_f, out_f, softmax_lse) + ctx.softmax_scale = softmax_scale + ctx.attn_mask_type = attn_mask_type + ctx.qkv_format = qkv_format + ctx.unflattened_shape = out.shape + # TE attention modules return the heads flattened into the last dimension + # ([b, s, h*d] for bshd), matching FlashAttention and FusedAttention. Returning the + # unflattened [b, s, h, d] makes autograd reject the incoming grad on shape mismatch. + return out.reshape(out.shape[0], out.shape[1], -1) + + @staticmethod + def backward(ctx, dout): + # pylint: disable=missing-function-docstring + from .frost_attention import ( # pylint: disable=import-outside-toplevel + frost_attn_bwd, + from_frost_layout, + to_frost_layout, + ) + + q_f, k_f, v_f, out_f, softmax_lse = ctx.saved_tensors + fmt = ctx.qkv_format + # dout arrives flattened, matching what forward returned; restore [b, s, h, d]. + dout = dout.reshape(ctx.unflattened_shape) + dq, dk, dv = frost_attn_bwd( + q_f, + k_f, + v_f, + out_f, + softmax_lse, + to_frost_layout(dout.contiguous(), fmt), + attn_scale=ctx.softmax_scale, + attn_mask_type=ctx.attn_mask_type, + ) + return ( + from_frost_layout(dq, fmt), + from_frost_layout(dk, fmt), + from_frost_layout(dv, fmt), + None, + None, + None, + None, + ) + + +class FrostAttention(torch.nn.Module): + """cuDNN FROST attention for symmetric head_dim in (256, 512] on SM100/SM103. + + This is the only backend that serves that head-dim range together with context parallelism, + which is what Gemma-4 global layers need. Deliberately narrow: no FP8, no bias, no dropout, + no softmax offset, no paging. get_attention_backend declines all of those before selecting + this backend, so anything reaching here should already be supported. + """ + + def __init__( + self, + softmax_scale: float, + attention_type: str = "self", + layer_number: Optional[int] = None, + deterministic: bool = False, + **kwargs, # attention_dropout / attention_dropout_ctx: accepted, must be unused + ) -> None: + super().__init__() + self.softmax_scale = softmax_scale + self.attention_type = attention_type + self.layer_number = 1 if layer_number is None else layer_number + self.deterministic = deterministic + self.attention_dropout = kwargs.get("attention_dropout", 0.0) + + def forward( + self, + query_layer: torch.Tensor, + key_layer: torch.Tensor, + value_layer: torch.Tensor, + qkv_format: str = "bshd", + cu_seqlens_q: Optional[torch.Tensor] = None, + cu_seqlens_kv: Optional[torch.Tensor] = None, + max_seqlen_q: Optional[int] = None, + max_seqlen_kv: Optional[int] = None, + cu_seqlens_q_padded: Optional[torch.Tensor] = None, + cu_seqlens_kv_padded: Optional[torch.Tensor] = None, + attn_mask_type: str = "causal", + window_size: Optional[Tuple[int, int]] = None, + cp_group: Optional[Union[dist_group_type, List[dist_group_type]]] = None, + cp_global_ranks: List[int] = None, + cp_stream: torch.cuda.Stream = None, + cp_comm_type: str = "p2p", + ) -> torch.Tensor: + """Forward pass. Routes through the CP ring when a cp_group is present.""" + assert self.attention_dropout == 0.0, "FrostAttention does not support dropout" + + context_parallel = cp_group is not None and get_distributed_world_size(cp_group) != 1 + if context_parallel: + output = attn_forward_func_with_cp( + self.training, + query_layer, + key_layer, + value_layer, + cu_seqlens_q, + cu_seqlens_kv, + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q_padded, + cu_seqlens_kv_padded, + 0.0, + cp_group, + cp_global_ranks, + cp_stream, + cp_comm_type, + softmax_scale=self.softmax_scale, + qkv_format=qkv_format, + attn_mask_type=attn_mask_type, + attn_bias_type="no_bias", + attn_bias=None, + deterministic=self.deterministic, + use_fused_attention=False, + use_frost_attention=True, + window_size=window_size, + layer_number=self.layer_number, + ) + # Same flattening the other backends apply after the CP call: the ring returns + # [b, s_local, h, d] but TE attention modules return heads in the last dimension. + return output.reshape(output.shape[0], output.shape[1], -1).contiguous() + + return FrostAttnFunc.apply( + query_layer, + key_layer, + value_layer, + self.softmax_scale, + attn_mask_type, + qkv_format, + self.training, + ) + + class FusedAttention(torch.nn.Module): """Dot product attention using `cuDNN attention `_: diff --git a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py index 658dab5d88d..13e9aec1eca 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py @@ -65,6 +65,7 @@ UnfusedDotProductAttention, FusedAttention, FlashAttention, + FrostAttention, ) @@ -79,6 +80,7 @@ "use_fused_attention": None, "fused_attention_backend": None, "use_unfused_attention": None, + "use_frost_attention": None, "backend_selection_requires_update": False, } @@ -996,6 +998,16 @@ def __init__( return_max_logit=self.return_max_logit, ) + # Only selectable for symmetric head_dim in (256, 512] on SM100/SM103, where no other + # backend can run at all. Cheap to construct, so instantiate unconditionally like the rest. + self.frost_attention = FrostAttention( + softmax_scale, + attention_type=attention_type, + layer_number=layer_number, + deterministic=self.deterministic, + **attn_kwargs, + ) + self.unfused_attention = UnfusedDotProductAttention( softmax_scale, attention_type=attention_type, @@ -2858,6 +2870,7 @@ def forward( use_fused_attention, fused_attention_backend, use_unfused_attention, + use_frost_attention, _, ) = dpa_utils.get_attention_backend(attention_params) # Set global _attention_backends var using return value @@ -2867,6 +2880,7 @@ def forward( _attention_backends["use_fused_attention"] = use_fused_attention _attention_backends["fused_attention_backend"] = fused_attention_backend _attention_backends["use_unfused_attention"] = use_unfused_attention + _attention_backends["use_frost_attention"] = use_frost_attention _attention_backends["backend_selection_requires_update"] = False # logging.Logger methods graph-break under torch.compile, so # selection is only logged in eager -- as in @@ -2885,6 +2899,8 @@ def forward( "Running with FusedAttention backend (sub-backend %s)", int(fused_attention_backend), ) + elif use_frost_attention: + logger.info("Running with FrostAttention backend (cuDNN FROST)") elif use_unfused_attention: logger.info("Running with UnfusedDotProductAttention backend") else: @@ -2893,9 +2909,20 @@ def forward( use_fused_attention = _attention_backends["use_fused_attention"] fused_attention_backend = _attention_backends["fused_attention_backend"] use_unfused_attention = _attention_backends["use_unfused_attention"] + use_frost_attention = _attention_backends["use_frost_attention"] # raise exception if no backend is available - if sum([use_flash_attention, use_fused_attention, use_unfused_attention]) == 0: + if ( + sum( + [ + use_flash_attention, + use_fused_attention, + use_unfused_attention, + use_frost_attention, + ] + ) + == 0 + ): raise ValueError( "No dot product attention backend is available for the provided inputs. Please" " run with NVTE_DEBUG=1 NVTE_DEBUG_LEVEL=2 to find out the reasons for" @@ -3058,6 +3085,26 @@ def forward( bf16_backward=bf16_backward, ) + if use_frost_attention: + return self.frost_attention( + query_layer, + key_layer, + value_layer, + qkv_format=qkv_format, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + max_seqlen_q=max_seqlen_q, + max_seqlen_kv=max_seqlen_kv, + cu_seqlens_q_padded=cu_seqlens_q_padded, + cu_seqlens_kv_padded=cu_seqlens_kv_padded, + attn_mask_type=attn_mask_type, + window_size=window_size, + cp_group=self.cp_group, + cp_global_ranks=self.cp_global_ranks, + cp_stream=self.cp_stream, + cp_comm_type=self.cp_comm_type, + ) + if use_unfused_attention: allow_emulation = ( os.getenv("NVTE_UnfusedDPA_Emulate_FP8", "0") == "1" or is_in_onnx_export_mode() diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 23b89287dfd..ded2de74551 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -613,6 +613,7 @@ def get_attention_backend( flash_attention_backend = None use_fused_attention = int(os.environ.get("NVTE_FUSED_ATTN", "1")) use_unfused_attention = int(os.environ.get("NVTE_UNFUSED_ATTN", "1")) + use_frost_attention = int(os.environ.get("NVTE_FROST_ATTN", "1")) if not use_flash_attention_2 and FlashAttentionUtils.is_installed: logger.debug("Disabling FlashAttention 2 due to NVTE_FLASH_ATTN=0 or NVTE_FLASH_ATTN_V2=0") if not use_flash_attention_3 and FlashAttentionUtils.v3_is_installed: @@ -1863,6 +1864,62 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt ), ) FlashAttentionUtils.warning_printed = True + # cuDNN FROST (CuTe-DSL SDPA in cuDNN Frontend >= 1.29.0) is the only backend that serves + # symmetric head_dim in (256, 512] on SM100/SM103. Every other option stops short: FA2/FA3 + # cap at 256, FA4 is disabled at symmetric 512 above, the C++ cuDNN fused path caps at 256, + # and UnfusedDotProductAttention supports 512 but not context parallelism. Without this, + # Gemma-4 global layers with CP > 1 select no backend at all. + if use_frost_attention: + # Local import: frost_attention pulls in cudnn lazily, so this stays cheap and keeps + # TE importable on systems without cudnn-frontend installed. + from .frost_attention import ( # pylint: disable=import-outside-toplevel + is_frost_attention_supported, + ) + + frost_supported, frost_reason = is_frost_attention_supported( + head_dim_qk=head_dim_qk, + head_dim_v=head_dim_v, + qkv_dtype=qkv_dtype, + attn_mask_type=attn_mask_type, + dropout=attention_dropout, + attn_bias_type=core_attention_bias_type, + ) + if not frost_supported: + logger.debug("Disabling FrostAttention: %s", frost_reason) + use_frost_attention = False + # Conservative guards for capabilities that exist in cuDNN but are not validated here yet. + # Each is a silent-wrong-answer risk rather than an error, so default to declining. + if use_frost_attention and softmax_type != "vanilla": + # CP asserts non-vanilla softmax needs FusedAttention; FROST implements plain softmax. + logger.debug("Disabling FrostAttention for softmax_type = %s", softmax_type) + use_frost_attention = False + if use_frost_attention and fp8: + logger.debug("Disabling FrostAttention for FP8") + use_frost_attention = False + if use_frost_attention and softcap is not None and softcap != 0.0: + logger.debug("Disabling FrostAttention for softcap") + use_frost_attention = False + if use_frost_attention and window_size not in ((-1, -1), (-1, 0)): + logger.debug("Disabling FrostAttention for sliding window %s", str(window_size)) + use_frost_attention = False + if use_frost_attention and "thd" in qkv_layout: + # bshd and sbhd are served directly from their own strides; thd is packed/varlen, which + # needs cu_seqlens plumbing that is neither implemented nor validated here. + logger.debug("Disabling FrostAttention for qkv_layout = %s", qkv_layout) + use_frost_attention = False + if use_frost_attention and context_parallel and cp_comm_type not in ( + "p2p", + "all_gather", + "a2a", + ): + # p2p (ring), all_gather and a2a are wired up in context_parallel.py; a2a+p2p is not. + # Non-p2p types matter for Gemma-4: TE refuses sliding-window attention with p2p, and the + # model has sliding layers, so those layers need all_gather or a2a. + logger.debug( + "Disabling FrostAttention for context parallelism with cp_comm_type = %s", cp_comm_type + ) + use_frost_attention = False + # All available backends if use_flash_attention_2 and not FlashAttentionUtils.is_installed: use_flash_attention_2 = False @@ -1905,13 +1962,20 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt if use_flash_attention: use_fused_attention = False use_unfused_attention = False + use_frost_attention = False elif use_fused_attention: use_unfused_attention = False + use_frost_attention = False + elif use_frost_attention: + # Preferred over the unfused path: same shape coverage, but fused and CP-capable. + use_unfused_attention = False selected_backend = "NoBackend" if use_flash_attention: selected_backend = f"FlashAttention ({str(flash_attention_backend)})" elif use_fused_attention: selected_backend = f"FusedAttention (sub-backend {int(fused_attention_backend)})" + elif use_frost_attention: + selected_backend = "FrostAttention (cuDNN FROST)" elif use_unfused_attention: selected_backend = "UnfusedDotProductAttention" logger.debug("Selected backend = %s.", selected_backend) @@ -1922,6 +1986,7 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt use_fused_attention, fused_attention_backend, use_unfused_attention, + use_frost_attention, available_backends, ) From d60b3fdc8004dcf5c42557c50a322a266141677d Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 15 Sep 2026 20:59:26 -0700 Subject: [PATCH 03/97] feat(attention): context parallelism for FROST across p2p, all_gather 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/context_parallel.py | 390 +++++++++++++++++- 1 file changed, 373 insertions(+), 17 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index 510d14ac638..024a52fab97 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1583,6 +1583,229 @@ def cp_p2p_bwd_flash_attn( return dq, dk, dv +def _frost_mask_for_section(attn_mask_type, section): + """Per-ring-step mask, mirroring cp_p2p_fwd_fused_attn. + + Only the diagonal tile keeps the causal mask; the off-diagonal tiles see a fully visible KV + block. This matches what was validated on B200: causal on the square diagonal, no_mask on the + rectangular off-diagonal tiles. + """ + if section in ("diagonal", "all"): + return attn_mask_type + if section in ("lower-triangle", "upper-triangle"): + return "no_mask" + raise ValueError("unknown CP section %r" % section) + + +def _frost_mask_for_window(window_size): + """Per-step mask for the all_gather path, derived from its adjusted window. + + get_kv_seq_info_after_all_gather trims KV and returns a window that is BOTTOM-RIGHT aligned: + (-1, 0) means causal relative to the trimmed KV, not top-left causal. Using top-left here + would silently compute a different mask, since the two only coincide when SQ == SKV and + all_gather never produces that. + """ + if window_size is None or tuple(window_size) == (-1, -1): + return "no_mask" + if tuple(window_size) == (-1, 0): + return "causal_bottom_right" + raise NotImplementedError( + "FROST all_gather does not support sliding window %s" % str(window_size) + ) + + +def cp_ag_fwd_frost_attn( + softmax_scale, + qkv_format, + window_size, + q_part, + k_part, + v_part, +): + """Per-step forward for CP all_gather with the cuDNN FROST backend. + + Simpler than the p2p ring: KV is already gathered and trimmed, so each step is a single + attention call with no LSE correction. Returns (out, softmax_lse). + """ + from .frost_attention import ( # pylint: disable=import-outside-toplevel + frost_attn_fwd, + from_frost_layout, + to_frost_layout, + ) + + out, softmax_lse = frost_attn_fwd( + to_frost_layout(q_part.contiguous(), qkv_format), + to_frost_layout(k_part.contiguous(), qkv_format), + to_frost_layout(v_part.contiguous(), qkv_format), + attn_scale=softmax_scale, + attn_mask_type=_frost_mask_for_window(window_size), + ) + return from_frost_layout(out, qkv_format), softmax_lse + + +def cp_ag_bwd_frost_attn( + softmax_scale, + qkv_format, + window_size, + softmax_lse, + q_part, + k_part, + v_part, + out_part, + dout_part, +): + """Per-step backward for CP all_gather with the cuDNN FROST backend.""" + from .frost_attention import ( # pylint: disable=import-outside-toplevel + frost_attn_bwd, + from_frost_layout, + to_frost_layout, + ) + + dq, dk, dv = frost_attn_bwd( + to_frost_layout(q_part.contiguous(), qkv_format), + to_frost_layout(k_part.contiguous(), qkv_format), + to_frost_layout(v_part.contiguous(), qkv_format), + to_frost_layout(out_part.contiguous(), qkv_format), + softmax_lse, + to_frost_layout(dout_part.contiguous(), qkv_format), + attn_scale=softmax_scale, + attn_mask_type=_frost_mask_for_window(window_size), + ) + return ( + from_frost_layout(dq, qkv_format), + from_frost_layout(dk, qkv_format), + from_frost_layout(dv, qkv_format), + ) + + +def cp_a2a_fwd_frost_attn(softmax_scale, attn_mask_type, qkv_format, q, k, v): + """Forward for CP a2a with the cuDNN FROST backend. + + The simplest of the three. After the all-to-all each rank holds the FULL sequence for a subset + of heads, so there is no ring, no KV trimming and no LSE correction: one ordinary attention + call with the caller mask type, top-left causal as usual. + """ + from .frost_attention import ( # pylint: disable=import-outside-toplevel + frost_attn_fwd, + from_frost_layout, + to_frost_layout, + ) + + out, softmax_lse = frost_attn_fwd( + to_frost_layout(q.contiguous(), qkv_format), + to_frost_layout(k.contiguous(), qkv_format), + to_frost_layout(v.contiguous(), qkv_format), + attn_scale=softmax_scale, + attn_mask_type=attn_mask_type, + ) + return from_frost_layout(out, qkv_format), softmax_lse + + +def cp_a2a_bwd_frost_attn( + softmax_scale, attn_mask_type, qkv_format, softmax_lse, q, k, v, out, dout +): + """Backward for CP a2a with the cuDNN FROST backend.""" + from .frost_attention import ( # pylint: disable=import-outside-toplevel + frost_attn_bwd, + from_frost_layout, + to_frost_layout, + ) + + dq, dk, dv = frost_attn_bwd( + to_frost_layout(q.contiguous(), qkv_format), + to_frost_layout(k.contiguous(), qkv_format), + to_frost_layout(v.contiguous(), qkv_format), + to_frost_layout(out.contiguous(), qkv_format), + softmax_lse, + to_frost_layout(dout.contiguous(), qkv_format), + attn_scale=softmax_scale, + attn_mask_type=attn_mask_type, + ) + return ( + from_frost_layout(dq, qkv_format), + from_frost_layout(dk, qkv_format), + from_frost_layout(dv, qkv_format), + ) + + +def cp_p2p_fwd_frost_attn( + softmax_scale, + attn_mask_type, + qkv_format, + q_part, + k_part, + v_part, + cu_seqlens_q_per_step, # noqa: ARG001 unused for bshd; matches the fused call convention + cu_seqlens_kv_per_step, # noqa: ARG001 + section, +): + """Per-tile forward call of CP P2P with the cuDNN FROST backend. + + Returns the same 5-tuple shape as cp_p2p_fwd_fused_attn so the ring code can consume it + unchanged. rng_state, attn_bias and max_logit are None: FROST supports neither dropout nor + bias, and the selector declines those configurations before we get here. + + softmax_lse comes back as [b, h, s] natural-log logsumexp in fp32, which is what the ring + correction in this file consumes (measured against an fp64 reference at 1.8e-06). + """ + from .frost_attention import ( # pylint: disable=import-outside-toplevel + frost_attn_fwd, + from_frost_layout, + to_frost_layout, + ) + + out, softmax_lse = frost_attn_fwd( + to_frost_layout(q_part.contiguous(), qkv_format), + to_frost_layout(k_part.contiguous(), qkv_format), + to_frost_layout(v_part.contiguous(), qkv_format), + attn_scale=softmax_scale, + attn_mask_type=_frost_mask_for_section(attn_mask_type, section), + ) + return from_frost_layout(out, qkv_format), softmax_lse, None, None, None + + +def cp_p2p_bwd_frost_attn( + softmax_scale, + attn_mask_type, + qkv_format, + softmax_lse, + softmax_lse_, + q_part, + k_part, + v_part, + out_part, + dout_part, + section, +): + """Per-tile backward call of CP P2P with the cuDNN FROST backend. + + Returns (dq, dk, dv, dbias) to match cp_p2p_bwd_fused_attn; dbias is always None. + """ + from .frost_attention import ( # pylint: disable=import-outside-toplevel + frost_attn_bwd, + from_frost_layout, + to_frost_layout, + ) + + softmax_lse_part = softmax_lse_ if section == "upper-triangle" else softmax_lse + dq, dk, dv = frost_attn_bwd( + to_frost_layout(q_part.contiguous(), qkv_format), + to_frost_layout(k_part.contiguous(), qkv_format), + to_frost_layout(v_part.contiguous(), qkv_format), + to_frost_layout(out_part.contiguous(), qkv_format), + softmax_lse_part, + to_frost_layout(dout_part.contiguous(), qkv_format), + attn_scale=softmax_scale, + attn_mask_type=_frost_mask_for_section(attn_mask_type, section), + ) + return ( + from_frost_layout(dq, qkv_format), + from_frost_layout(dk, qkv_format), + from_frost_layout(dv, qkv_format), + None, + ) + + class AttnFuncWithCPAndKVP2P(torch.autograd.Function): """ Attention implementation with context parallelism. Exchange KV between CP ranks @@ -1629,6 +1852,7 @@ def forward( use_flash_attn_4, fp8_output, layer_number, + use_frost_attention, ): # pylint: disable=missing-function-docstring @@ -1978,7 +2202,9 @@ def forward( i, cp_size, ] - if use_fused_attention: + if use_frost_attention: + frost_attn_inputs = [softmax_scale, attn_mask_type, qkv_format] + elif use_fused_attention: fused_attn_inputs = [ attn_bias, attn_bias_, @@ -2049,7 +2275,17 @@ def forward( cu_seqlens_kv_per_step[i], ) = prepare_outputs q_inputs[i % 2] = q_part - if use_fused_attention: + if use_frost_attention: + ( + out_per_step[i], + softmax_lse_per_step[i], + rng_states[i], + attn_biases[i], + max_logit_per_step[i], + ) = cp_p2p_fwd_frost_attn( + *frost_attn_inputs, *prepare_outputs, section + ) + elif use_fused_attention: ( out_per_step[i], softmax_lse_per_step[i], @@ -2078,7 +2314,17 @@ def forward( cu_seqlens_kv_per_step[i], ) = prepare_outputs q_inputs[i % 2] = q_part - if use_fused_attention: + if use_frost_attention: + ( + out_per_step[i], + softmax_lse_per_step[i], + rng_states[i], + attn_biases[i], + max_logit_per_step[i], + ) = cp_p2p_fwd_frost_attn( + *frost_attn_inputs, *prepare_outputs, section + ) + elif use_fused_attention: ( out_per_step[i], softmax_lse_per_step[i], @@ -2107,7 +2353,17 @@ def forward( cu_seqlens_kv_per_step[i], ) = prepare_outputs q_inputs[i % 2] = q_part - if use_fused_attention: + if use_frost_attention: + ( + out_per_step[i], + softmax_lse_per_step[i], + rng_states[i], + attn_biases[i], + max_logit_per_step[i], + ) = cp_p2p_fwd_frost_attn( + *frost_attn_inputs, *prepare_outputs, section + ) + elif use_fused_attention: ( out_per_step[i], softmax_lse_per_step[i], @@ -2137,7 +2393,15 @@ def forward( cu_seqlens_kv_per_step[i], ) = prepare_outputs q_inputs[i % 2] = q_part - if use_fused_attention: + if use_frost_attention: + ( + out_per_step[i], + softmax_lse_per_step[i], + rng_states[i], + attn_biases[i], + max_logit_per_step[i], + ) = cp_p2p_fwd_frost_attn(*frost_attn_inputs, *prepare_outputs, section) + elif use_fused_attention: ( out_per_step[i], softmax_lse_per_step[i], @@ -2402,6 +2666,7 @@ def forward( ctx.deterministic = deterministic ctx.softcap = softcap ctx.use_fused_attention = use_fused_attention + ctx.use_frost_attention = use_frost_attention ctx.pad_between_seqs = pad_between_seqs ctx.softmax_lse_in_packed_format = softmax_lse_in_packed_format ctx.second_half_lse_seqlen = second_half_lse_seqlen @@ -2763,7 +3028,15 @@ def backward(ctx, dout, *_args): cu_seqlens_q_padded, cu_seqlens_kv_padded, ] - if ctx.use_fused_attention: + if ctx.use_frost_attention: + frost_attn_inputs = [ + ctx.softmax_scale, + ctx.attn_mask_type, + ctx.qkv_format, + softmax_lse, + softmax_lse_, + ] + elif ctx.use_fused_attention: fused_attn_inputs = [ ctx.fp8, ctx.fp8_recipe, @@ -2835,7 +3108,11 @@ def backward(ctx, dout, *_args): if i == (cp_size - 1): section = "diagonal" prepare_outputs = cp_p2p_bwd_prepare_qkv(*prepare_inputs, section) - if ctx.use_fused_attention: + if ctx.use_frost_attention: + dq_, dk_, dv_, dbias_ = cp_p2p_bwd_frost_attn( + *frost_attn_inputs, *prepare_outputs, section + ) + elif ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( *fused_attn_inputs, *prepare_outputs, section ) @@ -2848,7 +3125,11 @@ def backward(ctx, dout, *_args): elif i >= (cp_size - rank - 1): section = "lower-triangle" prepare_outputs = cp_p2p_bwd_prepare_qkv(*prepare_inputs, section) - if ctx.use_fused_attention: + if ctx.use_frost_attention: + dq_, dk_, dv_, dbias_ = cp_p2p_bwd_frost_attn( + *frost_attn_inputs, *prepare_outputs, section + ) + elif ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( *fused_attn_inputs, *prepare_outputs, section ) @@ -2861,7 +3142,11 @@ def backward(ctx, dout, *_args): else: section = "upper-triangle" prepare_outputs = cp_p2p_bwd_prepare_qkv(*prepare_inputs, section) - if ctx.use_fused_attention: + if ctx.use_frost_attention: + dq_, dk_, dv_, dbias_ = cp_p2p_bwd_frost_attn( + *frost_attn_inputs, *prepare_outputs, section + ) + elif ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( *fused_attn_inputs, *prepare_outputs, section ) @@ -2874,7 +3159,11 @@ def backward(ctx, dout, *_args): else: section = "all" prepare_outputs = cp_p2p_bwd_prepare_qkv(*prepare_inputs, section) - if ctx.use_fused_attention: + if ctx.use_frost_attention: + dq_, dk_, dv_, dbias_ = cp_p2p_bwd_frost_attn( + *frost_attn_inputs, *prepare_outputs, section + ) + elif ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( *fused_attn_inputs, *prepare_outputs, section ) @@ -3215,6 +3504,7 @@ def backward(ctx, dout, *_args): None, None, None, + None, # use_frost_attention ) @@ -3304,6 +3594,7 @@ def forward( quantizers, fp8_output, load_balancing_strategy, + use_frost_attention, ): # pylint: disable=missing-function-docstring nvtx_range_push("transformer_engine.AttnFuncWithCPAndKVAllGather.forward") @@ -3710,7 +4001,17 @@ def forward( Float8Tensor.make_like(x, data=y, dtype=fwd_nominal_dtype) for x, y in zip([q_fp8, k_fp8, v_fp8], [q_part, k_part, v_part]) ] - if use_fused_attention: + if use_frost_attention: + out_per_step[i], softmax_lse_per_step[i] = cp_ag_fwd_frost_attn( + softmax_scale, + qkv_format, + window_size_per_step[i], + q_part, + k_part, + v_part, + ) + rng_states[i] = None # FROST has no dropout, so no RNG state + elif use_fused_attention: # Set per-step parameters for THD vs bshd/sbhd if qkv_format == "thd": cu_seqlens_q_ = thd_cu_seqlens_q_per_step[i] @@ -3980,6 +4281,7 @@ def forward( ctx.deterministic = deterministic ctx.softcap = softcap ctx.use_fused_attention = use_fused_attention + ctx.use_frost_attention = use_frost_attention ctx.use_flash_attn_3 = use_flash_attn_3 ctx.use_flash_attn_4 = use_flash_attn_4 ctx.pad_between_seqs = pad_between_seqs @@ -4250,7 +4552,23 @@ def backward(ctx, dout, *_args): out_part = out.select(seq_dim_o, i).contiguous() dout_part = dout.select(seq_dim_o, i).contiguous() - if ctx.use_fused_attention: + if ctx.use_frost_attention: + ( + dq_per_step[i], + dk_per_step[i], + dv_per_step[i], + ) = cp_ag_bwd_frost_attn( + ctx.softmax_scale, + ctx.qkv_format, + window_size_per_step[i], + softmax_lse_per_step[i], + q_part, + k_part, + v_part, + out_part, + dout_part, + ) + elif ctx.use_fused_attention: # Set per-step parameters for THD if ctx.qkv_format == "thd": cu_seqlens_q_ = thd_cu_seqlens_q_per_step[i] @@ -4577,6 +4895,7 @@ def backward(ctx, dout, *_args): None, None, None, + None, # use_frost_attention ) @@ -4621,6 +4940,7 @@ def forward( softmax_type, softmax_offset, fp8_output, + use_frost_attention, ): # pylint: disable=missing-function-docstring nvtx_range_push("transformer_engine.AttnFuncWithCPAndQKVOA2A.forward") @@ -4806,7 +5126,19 @@ def forward( ) ) qkv_scale_inv_format = None - if use_fused_attention: + if use_frost_attention: + out_, softmax_lse = cp_a2a_fwd_frost_attn( + softmax_scale, attn_mask_type, qkv_format, q, k, v + ) + # Only the LSE: FROST has no dropout, so there is no RNG state to carry, and a + # None in this list would have to survive the save/restore machinery. + aux_ctx_tensors = [softmax_lse] + # out_part is what gets saved for backward (f16_tensors below). Leaving it at its + # None initialisation makes `out` arrive as None in backward, which is not obvious + # from this branch alone: the fused path sets it inside its fp8 bookkeeping. + out_part = out_ + out_f16 = out_ + elif use_fused_attention: if fp8: if fp8_recipe.mxfp8(): q_fp8, k_fp8, v_fp8, qkv_layout, qkv_scale_inv_format = combine_and_quantize( @@ -5038,6 +5370,10 @@ def forward( ctx.softcap = softcap ctx.window_size = window_size ctx.use_fused_attention = use_fused_attention + ctx.use_frost_attention = use_frost_attention + # The a2a class never needed qkv_format in backward before: the fused and flash paths + # take a qkv_layout instead. FROST builds its graphs from the tensor layout, so it does. + ctx.qkv_format = qkv_format ctx.fp8_meta = fp8_meta ctx.is_input_fp8 = is_input_fp8 ctx.is_output_fp8 = is_output_fp8 @@ -5189,7 +5525,19 @@ def backward(ctx, dout, *_args): fa_backward_kwargs["softcap"] = ctx.softcap dq_fp8, dk_fp8, dv_fp8 = None, None, None - if ctx.use_fused_attention: + if ctx.use_frost_attention: + dq, dk, dv = cp_a2a_bwd_frost_attn( + ctx.softmax_scale, + ctx.attn_mask_type, + ctx.qkv_format, + aux_ctx_tensors[0], + q, + k, + v, + out, + dout, + ) + elif ctx.use_fused_attention: do_format = ctx.o_format do_scale_inv_format = None q_part, k_part, v_part, out_part, dout_part = q, k, v, out, dout @@ -5417,6 +5765,7 @@ def backward(ctx, dout, *_args): None, d_softmax_offset, None, + None, # use_frost_attention ) @@ -5556,6 +5905,7 @@ def attn_forward_func_with_cp( attn_bias=None, deterministic=False, use_fused_attention=False, + use_frost_attention=False, window_size=None, softcap=0.0, fp8=False, @@ -5693,9 +6043,12 @@ def attn_forward_func_with_cp( assert cu_seqlens_q is cu_seqlens_kv and ( cu_seqlens_q_padded is cu_seqlens_kv_padded ), "No-load-balance THD self-attention requires shared Q/KV sequence metadata tensors." - assert ( - qkv_format != "sbhd" or use_fused_attention - ), "Context parallelism does not support FlashAttention backend with qkv_format = 'sbhd'!" + # The restriction is FlashAttention-specific; the condition infers "not fused means flash", + # which predates FROST. FROST builds its cuDNN graphs from each tensor's actual strides, so + # sbhd is served directly. This matters because Megatron uses sbhd internally. + assert qkv_format != "sbhd" or use_fused_attention or use_frost_attention, ( + "Context parallelism does not support FlashAttention backend with qkv_format = 'sbhd'!" + ) assert attn_bias is None or (use_fused_attention and "padding" not in attn_mask_type), ( "Context parallelism only supports attention bias with FusedAttention backend and" " non-padding mask types!" @@ -5760,6 +6113,7 @@ def attn_forward_func_with_cp( use_flash_attn_4, fp8_output, layer_number, + use_frost_attention, ] out = AttnFuncWithCPAndKVP2P.apply(*args) elif cp_comm_type == "all_gather": @@ -5775,6 +6129,7 @@ def attn_forward_func_with_cp( quantizers, fp8_output, load_balancing_strategy, + use_frost_attention, ] out = AttnFuncWithCPAndKVAllGather.apply(*args) elif cp_comm_type == "a2a": @@ -5791,6 +6146,7 @@ def attn_forward_func_with_cp( softmax_type, softmax_offset, fp8_output, + use_frost_attention, ] out = AttnFuncWithCPAndQKVOA2A.apply(*args) else: From 62ffe3936a16b29c162c5d34eba8d7fc99eb0ea1 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 15 Sep 2026 20:59:35 -0700 Subject: [PATCH 04/97] test(attention): CP coverage for FrostAttention at head_dim 512 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 Signed-off-by: Nitin Vegesna --- tests/pytorch/attention/run_attention_with_cp.py | 9 +++++++++ tests/pytorch/attention/test_attention_with_cp.py | 10 ++++++++++ 2 files changed, 19 insertions(+) diff --git a/tests/pytorch/attention/run_attention_with_cp.py b/tests/pytorch/attention/run_attention_with_cp.py index 8d1b870d3a0..e820f212e98 100644 --- a/tests/pytorch/attention/run_attention_with_cp.py +++ b/tests/pytorch/attention/run_attention_with_cp.py @@ -17,6 +17,7 @@ from transformer_engine.pytorch import DType from test_attention_with_cp import ( model_configs_flash_attn, + model_configs_frost_attn, model_configs_fused_attn, ) from transformer_engine.pytorch import ( @@ -273,6 +274,14 @@ def run_dpa_with_cp( config = copy.deepcopy(model_configs_fused_attn[model]) else: assert False, f"{model=} is not a known FusedAttention CP config!" + if kernel_backend == "FrostAttention": + # Leave NVTE_FLASH_ATTN and NVTE_FUSED_ATTN at 0: FROST is the only backend that serves + # head_dim > 256, so get_attention_backend selects it on its own. + os.environ["NVTE_FROST_ATTN"] = "1" + if model in model_configs_frost_attn: + config = copy.deepcopy(model_configs_frost_attn[model]) + else: + assert False, f"{model=} is not a known FrostAttention CP config!" assert config.attn_mask_type in [ "causal", "no_mask", diff --git a/tests/pytorch/attention/test_attention_with_cp.py b/tests/pytorch/attention/test_attention_with_cp.py index 8b85c300577..9542157930e 100644 --- a/tests/pytorch/attention/test_attention_with_cp.py +++ b/tests/pytorch/attention/test_attention_with_cp.py @@ -459,6 +459,16 @@ def test_cp_with_flash_attention_softcap(cp_pool, cp_comm_type): ) +# cuDNN FROST: symmetric head_dim in (256, 512] on SM100/SM103, the range no other backend +# serves together with context parallelism. Shapes are Gemma-4 global layers, which is what +# motivated the backend. seqlen must stay divisible by cp_size * 2 for causal load balancing. +model_configs_frost_attn = { + # test: ModelConfig(b, sq, hq, dqk) + "cp_hd512_0": ModelConfig(2, 4096, 8, 512, num_gqa_groups=4, attn_mask_type="causal"), + "cp_hd512_1": ModelConfig(2, 4096, 8, 512, num_gqa_groups=4, attn_mask_type="no_mask"), + "cp_hd512_2": ModelConfig(2, 2048, 8, 512, num_gqa_groups=8, attn_mask_type="causal"), +} + model_configs_fused_attn = { # test: ModelConfig(b, sq, hq, dqk) "cp_1_0": ModelConfig(2, 4096, 12, 128, attn_mask_type="causal", return_max_logit=True), # MHA From fe72e4ae413b076464a3900b8c755f605779ecc3 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Wed, 16 Sep 2026 04:33:27 +0000 Subject: [PATCH 05/97] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- .../dot_product_attention/context_parallel.py | 6 +++--- .../attention/dot_product_attention/utils.py | 13 +++++++++---- 2 files changed, 12 insertions(+), 7 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index 024a52fab97..f3478e611ca 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -6046,9 +6046,9 @@ def attn_forward_func_with_cp( # The restriction is FlashAttention-specific; the condition infers "not fused means flash", # which predates FROST. FROST builds its cuDNN graphs from each tensor's actual strides, so # sbhd is served directly. This matters because Megatron uses sbhd internally. - assert qkv_format != "sbhd" or use_fused_attention or use_frost_attention, ( - "Context parallelism does not support FlashAttention backend with qkv_format = 'sbhd'!" - ) + assert ( + qkv_format != "sbhd" or use_fused_attention or use_frost_attention + ), "Context parallelism does not support FlashAttention backend with qkv_format = 'sbhd'!" assert attn_bias is None or (use_fused_attention and "padding" not in attn_mask_type), ( "Context parallelism only supports attention bias with FusedAttention backend and" " non-padding mask types!" diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index ded2de74551..b8670926c95 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -1907,10 +1907,15 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt # needs cu_seqlens plumbing that is neither implemented nor validated here. logger.debug("Disabling FrostAttention for qkv_layout = %s", qkv_layout) use_frost_attention = False - if use_frost_attention and context_parallel and cp_comm_type not in ( - "p2p", - "all_gather", - "a2a", + if ( + use_frost_attention + and context_parallel + and cp_comm_type + not in ( + "p2p", + "all_gather", + "a2a", + ) ): # p2p (ring), all_gather and a2a are wired up in context_parallel.py; a2a+p2p is not. # Non-p2p types matter for Gemma-4: TE refuses sliding-window attention with p2p, and the From a8957aeee1df9c59fe621c695df683aa495201ab Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 15 Sep 2026 22:15:53 -0700 Subject: [PATCH 06/97] test(attention): run the FrostAttention CP configs from pytest 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 Signed-off-by: Nitin Vegesna --- .../attention/test_attention_with_cp.py | 49 +++++++++++++++++++ 1 file changed, 49 insertions(+) diff --git a/tests/pytorch/attention/test_attention_with_cp.py b/tests/pytorch/attention/test_attention_with_cp.py index 9542157930e..d1dd818721f 100644 --- a/tests/pytorch/attention/test_attention_with_cp.py +++ b/tests/pytorch/attention/test_attention_with_cp.py @@ -758,6 +758,55 @@ def test_cp_with_fused_attention( ) +def _frost_availability(): + """Why FrostAttention cannot run here, or None if it can. + + The backend needs cuDNN Frontend >= 1.29.0 and, less obviously, + nvidia-cutlass-dsl >= 4.7.0: cudnn-frontend only declares >= 4.6.2, and below the FROST floor + every FROST engine silently declines and ordinary backend plans are returned with no error. + Reporting the reason as a skip keeps that distinguishable from a real failure. + """ + if get_device_compute_capability() not in ((10, 0), (10, 3)): + return "FrostAttention requires SM100/SM103 (the cuDNN d512 backward is Blackwell-only)." + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + is_frost_attention_available, + ) + + ok, reason = is_frost_attention_available() + return None if ok else reason + + +@pytest.mark.parametrize("model", model_configs_frost_attn.keys()) +@pytest.mark.parametrize("qkv_format", ["bshd", "sbhd"]) +@pytest.mark.parametrize("cp_comm_type", ["p2p", "all_gather", "a2a"]) +def test_cp_with_frost_attention(cp_pool, model, qkv_format, cp_comm_type): + """Context parallelism at head_dim 512, which no other backend serves. + + thd and a2a+p2p are excluded because the backend declines them: thd needs varlen support that + is not implemented, and a2a+p2p is not wired up. + """ + reason = _frost_availability() + if reason is not None: + pytest.skip(reason) + + config = model_configs_frost_attn[model] + config.context_parallel = True + config.cp_comm_type = cp_comm_type + + pool = cp_pool(2) + + _submit( + pool, + dtype="bf16", + model=model, + qkv_format=qkv_format, + kernel_backend="FrostAttention", + cp_comm_type=cp_comm_type, + is_training=True, + log_level=pytest_logging_level, + ) + + @pytest.mark.skipif(get_cudnn_version() < (8, 9, 7), reason="cuDNN 8.9.7+ is required.") @pytest.mark.skipif( get_device_compute_capability() < (9, 0), reason="FusedAttention THD requires sm90+." From bdbb36aad5304a21d32b26b28a01150910261a42 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 15 Sep 2026 23:02:45 -0700 Subject: [PATCH 07/97] fix(attention): gate FROST on the cuDNN Frontend version and key plans 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/frost_attention.py | 53 ++++++++++++++++--- 1 file changed, 47 insertions(+), 6 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index e7b2e2c0545..37db3c9f5fe 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -59,6 +59,11 @@ _FROST_BWD_PLAN_TOKEN = "sdpa_bwd_sm100" _MIN_CUTLASS_DSL = (4, 7, 0) +# 1.29.0 is the first release carrying the head_dim=512 BACKWARD (bprop_d512_f16_sm100). 1.28.0 +# ships the forward only, and the repo's own pin allows it, so without this check training would +# build a forward plan and then raise on the first backward. +_MIN_CUDNN_FRONTEND = (1, 29, 0) + _SUPPORTED_ARCHS = ((10, 0), (10, 3)) _MAX_HEAD_DIM = 512 _MIN_HEAD_DIM = 257 # below this the existing cuDNN/flash backends already serve the shape @@ -81,6 +86,22 @@ def _import_cudnn(): return _cudnn +def _parse_version(raw: str) -> Tuple[int, ...]: + """Leading numeric components of a version, ignoring any suffix. Unparseable sorts lowest.""" + parts = [] + for piece in str(raw).split(".")[:3]: + digits = "" + for ch in piece: + if not ch.isdigit(): + break + digits += ch + if not digits: + break + parts.append(int(digits)) + # Pad, or "1.29" would compare below (1, 29, 0) and be rejected as too old. + return tuple(parts + [0] * (3 - len(parts))) if parts else (0, 0, 0) + + def is_frost_attention_available() -> Tuple[bool, str]: """Whether the FROST kernels can be used at all, with a reason when they cannot. @@ -109,14 +130,24 @@ def _no(reason): from importlib.metadata import PackageNotFoundError, version + fe_raw = getattr(_cudnn, "__version__", None) + if fe_raw is None: + try: + fe_raw = version("nvidia-cudnn-frontend") + except PackageNotFoundError: + fe_raw = "0" + if _parse_version(fe_raw) < _MIN_CUDNN_FRONTEND: + return _no( + "nvidia-cudnn-frontend %s does not carry the head_dim>256 backward; >= 1.29.0 is" + " required (1.28.0 ships the forward only, so this would raise on the first" + " backward rather than here)" % fe_raw + ) + try: raw = version("nvidia-cutlass-dsl") except PackageNotFoundError: return _no("nvidia-cutlass-dsl not installed (FROST requires >= 4.7.0)") - try: - parsed = tuple(int(p) for p in raw.split(".")[:3]) - except ValueError: - parsed = (0, 0, 0) + parsed = _parse_version(raw) if parsed < _MIN_CUTLASS_DSL: # Worth being loud: this combination fails by silently declining, not by raising. return _no( @@ -259,8 +290,14 @@ def _select_frost_plan(graph, token: str, what: str): raise RuntimeError( "no cuDNN FROST %s engine was offered (looked for %r). Candidate plans: %s." - " nvidia-cutlass-dsl=%s (FROST floor 4.7.0)." - % (what, token, names[:6], version("nvidia-cutlass-dsl")) + " nvidia-cudnn-frontend=%s (floor 1.29.0), nvidia-cutlass-dsl=%s (floor 4.7.0)." + % ( + what, + token, + names[:6], + getattr(_cudnn, "__version__", None) or version("nvidia-cudnn-frontend"), + version("nvidia-cutlass-dsl"), + ) ) graph.select_plan(hits[0]) graph.check_support() @@ -372,6 +409,10 @@ def _cached(kind: str, key): def _key(q, k, mask, scale): return ( + # The graph is built under whichever device was current, so it must not be reused on + # another one. Matches the C++ fused-attn cache, which keys on device_id for the same + # reason. Normalise None, or "cuda" and "cuda:0" would build two plans for one device. + q.device.index if q.device.index is not None else torch.cuda.current_device(), q.shape[0], q.shape[1], k.shape[1], From e30ea42b484ee001d61096a3481367aa8d4fe6f8 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 15 Sep 2026 23:20:00 -0700 Subject: [PATCH 08/97] fix(attention): repair the FROST plan-cache key arity and harden the version gate The device element added to _key() in 63294e88 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/frost_attention.py | 92 ++++++++++--------- 1 file changed, 47 insertions(+), 45 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 37db3c9f5fe..c15e114deaf 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -36,9 +36,11 @@ from __future__ import annotations import os +from importlib.metadata import PackageNotFoundError, version as get_pkg_version from typing import Optional, Tuple import torch +from packaging.version import InvalidVersion, Version as PkgVersion __all__ = [ "is_frost_attention_available", @@ -57,12 +59,12 @@ # the selected plan by NAME rather than trusting that the engine was used. _FROST_FWD_PLAN_TOKEN = "sdpa_fwd_prefill_sm100" _FROST_BWD_PLAN_TOKEN = "sdpa_bwd_sm100" -_MIN_CUTLASS_DSL = (4, 7, 0) +_MIN_CUTLASS_DSL = PkgVersion("4.7.0") # 1.29.0 is the first release carrying the head_dim=512 BACKWARD (bprop_d512_f16_sm100). 1.28.0 # ships the forward only, and the repo's own pin allows it, so without this check training would # build a forward plan and then raise on the first backward. -_MIN_CUDNN_FRONTEND = (1, 29, 0) +_MIN_CUDNN_FRONTEND = PkgVersion("1.29.0") _SUPPORTED_ARCHS = ((10, 0), (10, 3)) _MAX_HEAD_DIM = 512 @@ -86,20 +88,23 @@ def _import_cudnn(): return _cudnn -def _parse_version(raw: str) -> Tuple[int, ...]: - """Leading numeric components of a version, ignoring any suffix. Unparseable sorts lowest.""" - parts = [] - for piece in str(raw).split(".")[:3]: - digits = "" - for ch in piece: - if not ch.isdigit(): - break - digits += ch - if not digits: - break - parts.append(int(digits)) - # Pad, or "1.29" would compare below (1, 29, 0) and be rejected as too old. - return tuple(parts + [0] * (3 - len(parts))) if parts else (0, 0, 0) +def _pkg_version(name: str, module=None) -> Optional[PkgVersion]: + """Installed version of a package, or None if it cannot be determined. + + Distribution metadata first, matching the sibling check in fused_mla_q_uproj.py, with the + module attribute as a fallback so a source or vendored install is not misreported as old. + """ + raw = None + try: + raw = get_pkg_version(name) + except PackageNotFoundError: + raw = getattr(module, "__version__", None) + if not isinstance(raw, str): + return None + try: + return PkgVersion(raw) + except InvalidVersion: + return None def is_frost_attention_available() -> Tuple[bool, str]: @@ -128,31 +133,25 @@ def _no(reason): except ImportError as exc: return _no("nvidia-cudnn-frontend not importable: %s" % exc) - from importlib.metadata import PackageNotFoundError, version - - fe_raw = getattr(_cudnn, "__version__", None) - if fe_raw is None: - try: - fe_raw = version("nvidia-cudnn-frontend") - except PackageNotFoundError: - fe_raw = "0" - if _parse_version(fe_raw) < _MIN_CUDNN_FRONTEND: + # Decline only on positive evidence of a too-old install. An undeterminable version is left + # to _select_frost_plan, which checks the plan by name and fails loudly with both versions. + frontend = _pkg_version("nvidia-cudnn-frontend", _cudnn) + if frontend is not None and frontend < _MIN_CUDNN_FRONTEND: return _no( - "nvidia-cudnn-frontend %s does not carry the head_dim>256 backward; >= 1.29.0 is" - " required (1.28.0 ships the forward only, so this would raise on the first" - " backward rather than here)" % fe_raw + "nvidia-cudnn-frontend %s registers no sm100 backward engine; >= %s is required" + " (1.28.0 ships the d512 forward only, so this would otherwise raise on the first" + " backward rather than here)" % (frontend, _MIN_CUDNN_FRONTEND) ) - try: - raw = version("nvidia-cutlass-dsl") - except PackageNotFoundError: - return _no("nvidia-cutlass-dsl not installed (FROST requires >= 4.7.0)") - parsed = _parse_version(raw) - if parsed < _MIN_CUTLASS_DSL: + cutlass = _pkg_version("nvidia-cutlass-dsl") + if cutlass is None: + return _no("nvidia-cutlass-dsl not installed (FROST requires >= %s)" % _MIN_CUTLASS_DSL) + if cutlass < _MIN_CUTLASS_DSL: # Worth being loud: this combination fails by silently declining, not by raising. return _no( - "nvidia-cutlass-dsl %s is below the FROST floor 4.7.0; FROST engines would be" - " silently skipped in favour of ordinary cuDNN backend plans" % raw + "nvidia-cutlass-dsl %s is below the FROST floor %s; FROST engines would be" + " silently skipped in favour of ordinary cuDNN backend plans" + % (cutlass, _MIN_CUTLASS_DSL) ) _availability = (True, "") @@ -286,17 +285,19 @@ def _select_frost_plan(graph, token: str, what: str): names = [graph.get_plan_name_at_index(i) for i in range(graph.get_execution_plan_count())] hits = [i for i, n in enumerate(names) if token in n] if not hits: - from importlib.metadata import version - + # Both versions, because either floor can cause this and blaming one misdirects. Looked + # up defensively: this is the message explaining a failure, so it must not raise itself. raise RuntimeError( "no cuDNN FROST %s engine was offered (looked for %r). Candidate plans: %s." - " nvidia-cudnn-frontend=%s (floor 1.29.0), nvidia-cutlass-dsl=%s (floor 4.7.0)." + " nvidia-cudnn-frontend=%s (floor %s), nvidia-cutlass-dsl=%s (floor %s)." % ( what, token, names[:6], - getattr(_cudnn, "__version__", None) or version("nvidia-cudnn-frontend"), - version("nvidia-cutlass-dsl"), + _pkg_version("nvidia-cudnn-frontend", _cudnn) or "unknown", + _MIN_CUDNN_FRONTEND, + _pkg_version("nvidia-cutlass-dsl") or "unknown", + _MIN_CUTLASS_DSL, ) ) graph.select_plan(hits[0]) @@ -308,7 +309,7 @@ def _select_frost_plan(graph, token: str, what: str): def _build_fwd(key) -> dict: """Build (and JIT-compile) a forward graph. Expensive; always reached through the cache.""" cudnn = _import_cudnn() - b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks = key + *_device, b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks = key io_dt = _cudnn_dtype(dtype) shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] @@ -347,7 +348,7 @@ def _build_fwd(key) -> dict: def _build_bwd(key) -> dict: """Build (and JIT-compile) a backward graph. Expensive; always reached through the cache.""" cudnn = _import_cudnn() - b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks = key + *_device, b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks = key io_dt = _cudnn_dtype(dtype) shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] @@ -411,8 +412,9 @@ def _key(q, k, mask, scale): return ( # The graph is built under whichever device was current, so it must not be reused on # another one. Matches the C++ fused-attn cache, which keys on device_id for the same - # reason. Normalise None, or "cuda" and "cuda:0" would build two plans for one device. - q.device.index if q.device.index is not None else torch.cuda.current_device(), + # reason. Type is included too, so a CPU tensor cannot alias cuda:0. + q.device.type, + q.device.index, q.shape[0], q.shape[1], k.shape[1], From 0957f712150a81d3bd6c85dbd049d1e1339a35be Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 15 Sep 2026 23:29:01 -0700 Subject: [PATCH 09/97] fix(attention): reject a v that does not match k, and refine the version 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/backends.py | 5 +- .../dot_product_attention/frost_attention.py | 70 +++++++++++++------ 2 files changed, 50 insertions(+), 25 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index 355da80830f..8ce1a336b4b 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -2303,8 +2303,9 @@ def forward(ctx, q, k, v, softmax_scale, attn_mask_type, qkv_format, is_training to_frost_layout, ) - # .contiguous() first: the graphs are built for BSHD-contiguous memory and - # frost_attention raises on anything else rather than computing on wrong strides. + # .contiguous() first: the graphs are built from each tensor's actual strides, so an + # arbitrary incoming layout would key a separate plan per layout and require k and v to + # agree. Normalising here keeps one plan per shape. q_f = to_frost_layout(q.contiguous(), qkv_format) k_f = to_frost_layout(k.contiguous(), qkv_format) v_f = to_frost_layout(v.contiguous(), qkv_format) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index c15e114deaf..9eb1f6c5d97 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -88,23 +88,29 @@ def _import_cudnn(): return _cudnn -def _pkg_version(name: str, module=None) -> Optional[PkgVersion]: - """Installed version of a package, or None if it cannot be determined. +def _pkg_version(name: str, module=None) -> Tuple[Optional[PkgVersion], Optional[str]]: + """(parsed version, raw string) for a package. Either element is None if undeterminable. Distribution metadata first, matching the sibling check in fused_mla_q_uproj.py, with the - module attribute as a fallback so a source or vendored install is not misreported as old. + module attribute as a fallback so a source or vendored install is not misreported as absent. + The raw string is returned separately so callers can tell "not installed" from "installed but + unparseable"; those warrant different answers, and conflating them declines valid installs. """ raw = None + for candidate in (lambda: get_pkg_version(name), lambda: getattr(module, "__version__", None)): + try: + raw = candidate() + except PackageNotFoundError: + raw = None + if isinstance(raw, str): + break + raw = None + if raw is None: + return None, None try: - raw = get_pkg_version(name) - except PackageNotFoundError: - raw = getattr(module, "__version__", None) - if not isinstance(raw, str): - return None - try: - return PkgVersion(raw) + return PkgVersion(raw), raw except InvalidVersion: - return None + return None, raw def is_frost_attention_available() -> Tuple[bool, str]: @@ -133,25 +139,26 @@ def _no(reason): except ImportError as exc: return _no("nvidia-cudnn-frontend not importable: %s" % exc) - # Decline only on positive evidence of a too-old install. An undeterminable version is left - # to _select_frost_plan, which checks the plan by name and fails loudly with both versions. - frontend = _pkg_version("nvidia-cudnn-frontend", _cudnn) + # Decline on positive evidence that FROST cannot work: a version below a floor, or a package + # that is absent outright. A version that is present but unparseable is NOT evidence, so it + # defers to _select_frost_plan, which checks the plan by name and reports both versions. + frontend, frontend_raw = _pkg_version("nvidia-cudnn-frontend", _cudnn) if frontend is not None and frontend < _MIN_CUDNN_FRONTEND: return _no( "nvidia-cudnn-frontend %s registers no sm100 backward engine; >= %s is required" " (1.28.0 ships the d512 forward only, so this would otherwise raise on the first" - " backward rather than here)" % (frontend, _MIN_CUDNN_FRONTEND) + " backward rather than here)" % (frontend_raw, _MIN_CUDNN_FRONTEND) ) - cutlass = _pkg_version("nvidia-cutlass-dsl") - if cutlass is None: + cutlass, cutlass_raw = _pkg_version("nvidia-cutlass-dsl") + if cutlass_raw is None: return _no("nvidia-cutlass-dsl not installed (FROST requires >= %s)" % _MIN_CUTLASS_DSL) - if cutlass < _MIN_CUTLASS_DSL: + if cutlass is not None and cutlass < _MIN_CUTLASS_DSL: # Worth being loud: this combination fails by silently declining, not by raising. return _no( "nvidia-cutlass-dsl %s is below the FROST floor %s; FROST engines would be" " silently skipped in favour of ordinary cuDNN backend plans" - % (cutlass, _MIN_CUTLASS_DSL) + % (cutlass_raw, _MIN_CUTLASS_DSL) ) _availability = (True, "") @@ -273,6 +280,23 @@ def _check_layout(name: str, t: torch.Tensor) -> None: ) +def _check_kv_match(k: torch.Tensor, v: torch.Tensor) -> None: + """Require v to match k in both shape and layout. + + Both graphs declare v with k's shape and stride, and _key records only q's and k's, so a v + that differs would hit a cached plan built for k's layout and read the wrong elements with no + error at all. Callers in TE always split k and v from one QKV tensor, so this costs nothing + and is purely a guard against a silent wrong answer. + """ + if k.shape != v.shape: + raise ValueError("k and v must have the same shape; got %s and %s" % (k.shape, v.shape)) + if k.stride() != v.stride(): + raise ValueError( + "k and v must have the same layout; got strides %s and %s" + % (tuple(k.stride()), tuple(v.stride())) + ) + + def _select_frost_plan(graph, token: str, what: str): """Select a plan whose name proves a FROST engine was chosen. @@ -294,9 +318,9 @@ def _select_frost_plan(graph, token: str, what: str): what, token, names[:6], - _pkg_version("nvidia-cudnn-frontend", _cudnn) or "unknown", + _pkg_version("nvidia-cudnn-frontend", _cudnn)[1] or "unknown", _MIN_CUDNN_FRONTEND, - _pkg_version("nvidia-cutlass-dsl") or "unknown", + _pkg_version("nvidia-cutlass-dsl")[1] or "unknown", _MIN_CUTLASS_DSL, ) ) @@ -447,8 +471,7 @@ def frost_attn_fwd( """ for name, tensor in (("q", q), ("k", k), ("v", v)): _check_layout(name, tensor) - if k.shape != v.shape: - raise ValueError("k and v must have the same shape; got %s and %s" % (k.shape, v.shape)) + _check_kv_match(k, v) if q.shape[1] % k.shape[1] != 0: raise ValueError( "num_heads must be divisible by num_gqa_groups; got %d and %d" @@ -485,6 +508,7 @@ def frost_attn_bwd( """Backward attention via cuDNN FROST. `softmax_lse` is [b, h, s] as returned by the forward.""" for name, tensor in (("q", q), ("k", k), ("v", v), ("out", out), ("dout", dout)): _check_layout(name, tensor) + _check_kv_match(k, v) mask = _mask_mode(attn_mask_type) scale = attn_scale if attn_scale is not None else q.shape[-1] ** -0.5 From 438b4da5d1c602b67d7597a05fb65e5d0f687c89 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 15 Sep 2026 23:40:16 -0700 Subject: [PATCH 10/97] fix(attention): update the mixed-THD backend unpack for the new return 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 Signed-off-by: Nitin Vegesna --- .../attention/run_attention_with_cp.py | 13 +++++++++ .../attention/test_mixed_thd_attention.py | 2 +- .../dot_product_attention/context_parallel.py | 2 +- .../dot_product_attention.py | 2 ++ .../dot_product_attention/frost_attention.py | 29 ++++++++++--------- 5 files changed, 32 insertions(+), 16 deletions(-) diff --git a/tests/pytorch/attention/run_attention_with_cp.py b/tests/pytorch/attention/run_attention_with_cp.py index e820f212e98..842be0400d8 100644 --- a/tests/pytorch/attention/run_attention_with_cp.py +++ b/tests/pytorch/attention/run_attention_with_cp.py @@ -602,6 +602,19 @@ def run_dpa_with_cp( pad_between_seqs=pad_between_seqs, fp8_output=fp8_mha, ) + if kernel_backend == "FrostAttention": + # Assert the backend actually used, not just the one requested. FROST is currently + # the only selectable backend for these configs -- flash and fused are env-gated off + # and CP disables unfused -- so a silent substitution is impossible today and this + # would pass by construction. It is here so it stops passing if that stops being + # true, rather than quietly testing some other kernel. + from transformer_engine.pytorch.attention.dot_product_attention.dot_product_attention import ( # pylint: disable=import-outside-toplevel + _attention_backends, + ) + + assert _attention_backends[ + "use_frost_attention" + ], "expected FrostAttention to be selected, got %s" % (_attention_backends,) if config.return_max_logit: out_, max_logit_ = out_ if is_training: diff --git a/tests/pytorch/attention/test_mixed_thd_attention.py b/tests/pytorch/attention/test_mixed_thd_attention.py index d665df4ceff..d4618db126d 100644 --- a/tests/pytorch/attention/test_mixed_thd_attention.py +++ b/tests/pytorch/attention/test_mixed_thd_attention.py @@ -453,7 +453,7 @@ def test_thd_mask_type_runtime_dispatch_uses_backend_selection(monkeypatch): def fake_get_attention_backend(attention_params): observed_params.append(attention_params) available_backends = [False, attention_params.attn_mask_type == "padding", False] - return False, None, available_backends[1], None, False, available_backends + return False, None, available_backends[1], None, False, False, available_backends monkeypatch.setattr(dpa_module.dpa_utils, "get_attention_backend", fake_get_attention_backend) padded_policies, grouped_policies = DotProductAttention._partition_thd_mask_policies( diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index f3478e611ca..6cfc1e11ea5 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1746,7 +1746,7 @@ def cp_p2p_fwd_frost_attn( bias, and the selector declines those configurations before we get here. softmax_lse comes back as [b, h, s] natural-log logsumexp in fp32, which is what the ring - correction in this file consumes (measured against an fp64 reference at 1.8e-06). + correction in this file consumes. """ from .frost_attention import ( # pylint: disable=import-outside-toplevel frost_attn_fwd, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py index 13e9aec1eca..dc16cd967e9 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py @@ -158,6 +158,7 @@ def _get_thd_policy_attention_backend( use_fused_attention, fused_attention_backend, use_unfused_attention, + use_frost_attention, _, ) = selection _attention_backends.update( @@ -168,6 +169,7 @@ def _get_thd_policy_attention_backend( "use_fused_attention": use_fused_attention, "fused_attention_backend": fused_attention_backend, "use_unfused_attention": use_unfused_attention, + "use_frost_attention": use_frost_attention, "backend_selection_requires_update": False, } ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 9eb1f6c5d97..38c63fc5984 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -11,26 +11,26 @@ symmetric 512 forward and backward on Blackwell, reachable through the ordinary cuDNN graph API. This module wraps them so TE, including its CP ring, can dispatch to them. -Three properties were measured on B200 before this was written, and each one constrains the code: +Three properties of these kernels were verified on Blackwell before this was written, and each +one constrains the code: 1. cuDNN's `use_causal_mask` is TOP-LEFT aligned and `use_causal_mask_bottom_right` is - bottom-right; both were verified against references at SQ=1024/SKV=2048, where the two - disagree by three orders of magnitude (1.6e-03 vs 3.5e+00). They coincide when SQ == SKV, so - the distinction is invisible in square tests and decisive for all_gather, which trims KV. - `_MASK_MODES` lists only spellings checked this way: sdpa() ignores unknown kwargs silently, - so an unverified name would apply no mask at all and still run. + bottom-right. They coincide when SQ == SKV, so the distinction is invisible in square tests + and decisive for all_gather, which trims KV. `_MASK_MODES` lists only spellings checked + against a reference for their alignment: sdpa() ignores unknown kwargs silently, so an + unverified name would apply no mask at all and still run. -2. Plan building must be cached. Building a plan costs ~1972 ms the first time and ~12 ms once - cuDNN has cached the JIT, against a ~0.129 ms execute. Even the cached rebuild is ~90x an - execute, so a per-call build would make training build-bound. Hence `_PLAN_CACHE`. +2. Plan building must be cached. Building a plan is by far the most expensive cuDNN frontend + call here, and dominates an execute even after cuDNN has cached the JIT and made rebuilds + cheap, so a per-call build would leave training build-bound. Hence `_PLAN_CACHE`. 3. The forward LSE is natural-log logsumexp in fp32, shaped [b, h, s, 1]. Squeezed to [b, h, s] - it is exactly what the CP ring correction in context_parallel.py consumes (max err 1.8e-06 vs - an fp64 reference), which is what makes ring attention over these kernels valid at all. + it is exactly what the CP ring correction in context_parallel.py consumes, which is what + makes ring attention over these kernels valid at all. Numerics were validated against the criterion FlashAttention applies to itself, namely that the -kernel error must stay within 2x the error bf16 inputs alone produce: observed 0.21x to 0.62x -across square and rectangular, causal and non-causal shapes. +kernel error must stay within 2x the error bf16 inputs alone produce, across square and +rectangular, causal and non-causal shapes. """ from __future__ import annotations @@ -423,7 +423,8 @@ def _build_bwd(key) -> dict: def _cached(kind: str, key): - """Plan cache. See module docstring: a build is ~15000x an execute, so this is required.""" + """Plan cache. See module docstring: building dominates executing even once the JIT is + cached, so this is required rather than an optimisation.""" cache_key = (kind,) + key entry = _PLAN_CACHE.get(cache_key) if entry is None: From c0d071307efff92a0933d97f37172433152a7aec Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 15 Sep 2026 23:46:09 -0700 Subject: [PATCH 11/97] fix(attention): bind a cuDNN stream and close the remaining silent-wrong-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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/backends.py | 4 + .../dot_product_attention.py | 1 + .../dot_product_attention/frost_attention.py | 89 +++++++++++++++++-- .../attention/dot_product_attention/utils.py | 25 ++++++ 4 files changed, 114 insertions(+), 5 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index 8ce1a336b4b..58ef53e5685 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -2400,6 +2400,9 @@ def forward( cp_global_ranks: List[int] = None, cp_stream: torch.cuda.Stream = None, cp_comm_type: str = "p2p", + load_balancing_strategy: CPLoadBalancingStrategy = ( + CPLoadBalancingStrategy.DUAL_CHUNK_SWAP + ), ) -> torch.Tensor: """Forward pass. Routes through the CP ring when a cp_group is present.""" assert self.attention_dropout == 0.0, "FrostAttention does not support dropout" @@ -2432,6 +2435,7 @@ def forward( use_frost_attention=True, window_size=window_size, layer_number=self.layer_number, + load_balancing_strategy=load_balancing_strategy, ) # Same flattening the other backends apply after the CP call: the ring returns # [b, s_local, h, d] but TE attention modules return heads in the last dimension. diff --git a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py index dc16cd967e9..b4ecff5e836 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py @@ -3105,6 +3105,7 @@ def forward( cp_global_ranks=self.cp_global_ranks, cp_stream=self.cp_stream, cp_comm_type=self.cp_comm_type, + load_balancing_strategy=self.load_balancing_strategy, ) if use_unfused_attention: diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 38c63fc5984..da0a0a43012 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -73,6 +73,7 @@ _cudnn = None _availability: Optional[Tuple[bool, str]] = None _PLAN_CACHE: dict = {} +_HANDLES: dict = {} def _import_cudnn(): @@ -88,6 +89,36 @@ def _import_cudnn(): return _cudnn +def _handle_for(device: torch.device): + """A cuDNN handle for `device`, bound to PyTorch's current stream on it. + + Without this, cuDNN runs on its default handle's stream while the tensors and workspace are + allocated on PyTorch's current stream, and nothing orders the two. That is not hypothetical + here: the p2p CP ring issues attention inside `with torch.cuda.stream(cp_stream)`, so on + alternating ring steps the kernel and its buffers would be on different streams. Re-binding + on every call is what flex_attention.py does, and is required because the same cached plan is + executed from different streams across ring steps. + """ + if device.type != "cuda": + raise ValueError("FrostAttention requires CUDA tensors; got device %s" % device) + cudnn = _import_cudnn() + if device.index is None: + device = torch.device("cuda", torch.cuda.current_device()) + with torch.cuda.device(device): + handle = _HANDLES.get(device) + if handle is None: + handle = cudnn.create_handle() + _HANDLES[device] = handle + cudnn.set_stream(handle=handle, stream=torch.cuda.current_stream(device).cuda_stream) + return handle + + +def _device_from_key(device_key) -> torch.device: + """Rebuild the torch.device that _key recorded, for building under the right device.""" + kind, index = device_key + return torch.device(kind) if index is None else torch.device(kind, index) + + def _pkg_version(name: str, module=None) -> Tuple[Optional[PkgVersion], Optional[str]]: """(parsed version, raw string) for a package. Either element is None if undeterminable. @@ -280,6 +311,17 @@ def _check_layout(name: str, t: torch.Tensor) -> None: ) +def _check_dtype(name: str, t: torch.Tensor, expected: torch.dtype) -> None: + """Require a tensor to carry the dtype its graph node was declared with. + + Every node but `stats` is declared from q's dtype, and execute() binds raw pointers, so a + tensor of another dtype would have its bits reinterpreted with no error at all. `dout` + matters most: it arrives from autograd and is not this module's to control. + """ + if t.dtype != expected: + raise ValueError("%s must be %s to match q; got %s" % (name, expected, t.dtype)) + + def _check_kv_match(k: torch.Tensor, v: torch.Tensor) -> None: """Require v to match k in both shape and layout. @@ -341,6 +383,7 @@ def _build_fwd(key) -> dict: io_data_type=io_dt, intermediate_data_type=cudnn.data_type.FLOAT, compute_data_type=cudnn.data_type.FLOAT, + handle=_handle_for(_device_from_key(_device)), ) tq = graph.tensor(name="q", dim=shq, stride=list(qs)) tk = graph.tensor(name="k", dim=shkv, stride=list(ks)) @@ -380,6 +423,7 @@ def _build_bwd(key) -> dict: io_data_type=io_dt, intermediate_data_type=cudnn.data_type.FLOAT, compute_data_type=cudnn.data_type.FLOAT, + handle=_handle_for(_device_from_key(_device)), ) handles = {} # o and dO share q's layout; k, v and their grads share k's. @@ -465,14 +509,22 @@ def frost_attn_fwd( ) -> Tuple[torch.Tensor, torch.Tensor]: """Forward attention via cuDNN FROST. - q, k, v are [b, h, s, d] views over BSHD-contiguous memory. GQA is supported directly - (h_kv may differ from h_q) and SQ need not equal SKV, which is what lets a CP ring step - use this. Returns (out, softmax_lse) with softmax_lse as [b, h, s] fp32 natural-log - logsumexp, the layout and convention the CP ring correction expects. + q, k, v are [b, h, s, d] views; bshd and sbhd are both served, since the graph is built from + each tensor's actual strides. GQA is supported directly (h_kv may differ from h_q) and SQ + need not equal SKV, which is what lets a CP ring step use this. Returns (out, softmax_lse) + with softmax_lse as [b, h, s] fp32 natural-log logsumexp, the layout and convention the CP + ring correction expects. """ for name, tensor in (("q", q), ("k", k), ("v", v)): _check_layout(name, tensor) + _check_dtype(name, tensor, q.dtype) _check_kv_match(k, v) + if k.shape[0] != q.shape[0] or k.shape[3] != q.shape[3]: + # The graph declares k and v with q's batch and head_dim, so a mismatch would bind a + # differently shaped buffer to that node and read the wrong elements silently. + raise ValueError( + "k must match q in batch and head_dim; got q %s and k %s" % (q.shape, k.shape) + ) if q.shape[1] % k.shape[1] != 0: raise ValueError( "num_heads must be divisible by num_gqa_groups; got %d and %d" @@ -492,7 +544,9 @@ def frost_attn_fwd( out = torch.empty_strided(q.shape, q.stride(), device=q.device, dtype=q.dtype) lse = torch.empty(b, hq, sq, 1, device=q.device, dtype=torch.float32) workspace = torch.empty(entry["workspace"], device=q.device, dtype=torch.uint8) - entry["graph"].execute({tq: q, tk: k, tv: v, tout: out, tlse: lse}, workspace) + entry["graph"].execute( + {tq: q, tk: k, tv: v, tout: out, tlse: lse}, workspace, handle=_handle_for(q.device) + ) return out, lse.squeeze(-1) @@ -509,7 +563,31 @@ def frost_attn_bwd( """Backward attention via cuDNN FROST. `softmax_lse` is [b, h, s] as returned by the forward.""" for name, tensor in (("q", q), ("k", k), ("v", v), ("out", out), ("dout", dout)): _check_layout(name, tensor) + _check_dtype(name, tensor, q.dtype) _check_kv_match(k, v) + # The same shape assumptions the forward makes, plus o/dO, which the graph declares with q's + # shape. The forward runs first in autograd, but the CP ring calls this directly. + if k.shape[0] != q.shape[0] or k.shape[3] != q.shape[3]: + raise ValueError( + "k must match q in batch and head_dim; got q %s and k %s" % (q.shape, k.shape) + ) + if q.shape[1] % k.shape[1] != 0: + raise ValueError( + "num_heads must be divisible by num_gqa_groups; got %d and %d" + % (q.shape[1], k.shape[1]) + ) + for name, tensor in (("out", out), ("dout", dout)): + if tensor.shape != q.shape: + raise ValueError( + "%s must have q's shape; got %s and %s" % (name, tensor.shape, q.shape) + ) + if softmax_lse.dtype != torch.float32: + raise ValueError("softmax_lse must be fp32; got %s" % softmax_lse.dtype) + if tuple(softmax_lse.shape[:3]) != tuple(q.shape[:3]): + raise ValueError( + "softmax_lse must be [b, h, s] matching q; got %s and %s" + % (tuple(softmax_lse.shape), tuple(q.shape)) + ) mask = _mask_mode(attn_mask_type) scale = attn_scale if attn_scale is not None else q.shape[-1] ** -0.5 @@ -550,5 +628,6 @@ def _as(t, ref): h["dv"]: dv, }, workspace, + handle=_handle_for(q.device), ) return dq, dk, dv diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index b8670926c95..c8b7b7abaae 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -1907,6 +1907,31 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt # needs cu_seqlens plumbing that is neither implemented nor validated here. logger.debug("Disabling FrostAttention for qkv_layout = %s", qkv_layout) use_frost_attention = False + if use_frost_attention and return_max_logit: + # FrostAttention returns the context layer alone, where UnfusedDotProductAttention returns + # (context, max_logit). Selecting it here would break the caller's unpack. + logger.debug("Disabling FrostAttention for max_logit") + use_frost_attention = False + if use_frost_attention and inference_params is not None: + # Unreachable today, since KV caching asserts a padding mask and FROST declines those. + # Explicit anyway: no page table reaches the backend, so a paged cache would be read raw. + logger.debug("Disabling FrostAttention for KV caching") + use_frost_attention = False + if use_frost_attention and context_parallel: + # Same two restrictions FlashAttention and FusedAttention carry above. The ring chunking + # assumes square tiles, so an unequal q/kv length is a wrong answer rather than an error. + if "bottom_right" in attn_mask_type: + logger.debug( + "Disabling FrostAttention as it does not support context parallelism with" + " causal_bottom_right masking" + ) + use_frost_attention = False + elif "causal" in attn_mask_type and max_seqlen_q != max_seqlen_kv: + logger.debug( + "Disabling FrostAttention as it does not support context parallelism with causal" + " masking for cross-attention" + ) + use_frost_attention = False if ( use_frost_attention and context_parallel From 7504bffbd6fa0d56573bfabe43b14bbcf17775b2 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 00:05:58 -0700 Subject: [PATCH 12/97] fix(attention): JIT-compile FROST plans under the device their handle 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 Signed-off-by: Nitin Vegesna --- .../attention/dot_product_attention/frost_attention.py | 8 +++++++- .../pytorch/attention/dot_product_attention/utils.py | 6 ++++-- 2 files changed, 11 insertions(+), 3 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index da0a0a43012..038f83c746a 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -35,6 +35,7 @@ from __future__ import annotations +import contextlib import os from importlib.metadata import PackageNotFoundError, version as get_pkg_version from typing import Optional, Tuple @@ -472,7 +473,12 @@ def _cached(kind: str, key): cache_key = (kind,) + key entry = _PLAN_CACHE.get(cache_key) if entry is None: - entry = _build_fwd(key) if kind == "fwd" else _build_bwd(key) + # Build under the device the key names, not merely with that device's handle: the plans + # are CuTe-DSL JIT-compiled, and a compile path is far more likely to read the ambient + # CUDA context than the handle. Free to do, and removes the question entirely. + device = _device_from_key(key[:2]) + with torch.cuda.device(device) if device.type == "cuda" else contextlib.nullcontext(): + entry = _build_fwd(key) if kind == "fwd" else _build_bwd(key) _PLAN_CACHE[cache_key] = entry return entry diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index c8b7b7abaae..6674654996f 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -1918,8 +1918,10 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt logger.debug("Disabling FrostAttention for KV caching") use_frost_attention = False if use_frost_attention and context_parallel: - # Same two restrictions FlashAttention and FusedAttention carry above. The ring chunking - # assumes square tiles, so an unequal q/kv length is a wrong answer rather than an error. + # Same two restrictions FlashAttention and FusedAttention carry above. Both are about + # where the causal diagonal sits: the ring shards q and kv independently, so a mask whose + # position depends on the q/kv lengths lands differently per step. no_mask is unaffected + # and stays allowed even when the lengths differ. if "bottom_right" in attn_mask_type: logger.debug( "Disabling FrostAttention as it does not support context parallelism with" From 85a0f512ce0bcead1cc3e67719a4f037276966cf Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 00:15:06 -0700 Subject: [PATCH 13/97] test(attention): anchor FROST numerics to an fp32 reference, not to itself 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 Signed-off-by: Nitin Vegesna --- qa/L0_pytorch_unittest/test.sh | 1 + .../pytorch/attention/test_frost_attention.py | 218 ++++++++++++++++++ .../dot_product_attention/frost_attention.py | 42 +++- .../attention/dot_product_attention/utils.py | 6 +- 4 files changed, 255 insertions(+), 12 deletions(-) create mode 100644 tests/pytorch/attention/test_frost_attention.py diff --git a/qa/L0_pytorch_unittest/test.sh b/qa/L0_pytorch_unittest/test.sh index a78a99d7f95..6584ff83337 100644 --- a/qa/L0_pytorch_unittest/test.sh +++ b/qa/L0_pytorch_unittest/test.sh @@ -63,6 +63,7 @@ python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_hybrid_quantizat python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_identity_quantizer.xml $TE_PATH/tests/pytorch/test_identity_quantizer.py || test_fail "test_identity_quantizer.py" NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_attention.xml $TE_PATH/tests/pytorch/attention/test_attention.py || test_fail "test_attention.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_flex_attention.xml $TE_PATH/tests/pytorch/attention/test_flex_attention.py || test_fail "test_flex_attention.py" +python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_frost_attention.xml $TE_PATH/tests/pytorch/attention/test_frost_attention.py || test_fail "test_frost_attention.py" NVTE_GDN_TEST_REQUIRED=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_gdn_attention.xml $TE_PATH/tests/pytorch/attention/test_gdn_attention.py || test_fail "test_gdn_attention.py" NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE=1 NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_attention_deterministic.xml $TE_PATH/tests/pytorch/attention/test_attention.py || test_fail "NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 test_attention.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_linear_mxfp8_attention.xml $TE_PATH/tests/pytorch/attention/test_linear_mxfp8_attention.py || test_fail "test_linear_mxfp8_attention.py" diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py new file mode 100644 index 00000000000..7068ec835ae --- /dev/null +++ b/tests/pytorch/attention/test_frost_attention.py @@ -0,0 +1,218 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Numerical tests for the cuDNN FROST attention backend. + +These exist because the CP tests cannot catch what this backend is most likely to get wrong. +run_attention_with_cp.py compares a context-parallel run against a non-CP run *of the same +backend*, which validates the ring plumbing and nothing about the kernel: a systematic error -- +a wrong softmax scale, a causal mask anchored to the wrong corner, an LSE in the wrong log base +-- appears identically on both sides and cancels. Everything here is anchored to an independent +fp32 reference instead. + +The pass criterion is the one FlashAttention applies to itself: the kernel's error against an +fp32 reference must stay within 2x the error that comes from feeding the same reference bf16 +inputs. That floor is measured per case rather than hard-coded, so the bar tracks the shape and +dtype instead of encoding a number that silently rots. +""" + +import math + +import pytest +import torch + +from transformer_engine.pytorch import get_device_compute_capability + + +def _frost_availability(): + """Why FrostAttention cannot run here, or None if it can.""" + if not torch.cuda.is_available(): + return "no CUDA device" + if get_device_compute_capability() not in ((10, 0), (10, 3)): + return "FrostAttention requires SM100/SM103 (the cuDNN d512 backward is Blackwell-only)." + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + is_frost_attention_available, + ) + + ok, reason = is_frost_attention_available() + return None if ok else reason + + +_SKIP = _frost_availability() +pytestmark = pytest.mark.skipif(_SKIP is not None, reason=str(_SKIP)) + +# head_dim 512 is the whole point of the backend; 320 checks the interior of the (256, 512] range +# rather than only its endpoint. +_SHAPES = [ + # b, hq, hkv, sq, skv, d + (2, 8, 4, 1024, 1024, 512), # Gemma-4 global layer, GQA + (2, 8, 8, 512, 512, 512), # MHA + (1, 4, 4, 256, 512, 512), # sq != skv, which is where mask alignment matters + (2, 4, 4, 512, 512, 320), # interior head_dim +] + + +def _reference(q, k, v, scale, mask): + """Attention in fp32, computed independently of TE and of cuDNN.""" + qq, kk, vv = q.float(), k.float(), v.float() + rep = qq.shape[1] // kk.shape[1] + kk = kk.repeat_interleave(rep, dim=1) + vv = vv.repeat_interleave(rep, dim=1) + s = (qq @ kk.transpose(-1, -2)) * scale + if mask != "no_mask": + sq, skv = qq.shape[2], kk.shape[2] + # Top-left for "causal", bottom-right for "causal_bottom_right". These coincide only when + # sq == skv, which is exactly why _SHAPES includes a rectangular case. + offset = 0 if mask == "causal" else skv - sq + causal = torch.ones(sq, skv, device=q.device, dtype=torch.bool).triu(offset + 1) + s = s.masked_fill(causal, float("-inf")) + p = s.softmax(-1) + return p @ vv, torch.logsumexp(s, dim=-1) + + +def _floor(q32, k32, v32, scale, mask, dtype): + """The error `dtype` inputs alone cause, and the exact answer to measure the kernel against. + + The inputs must originate in fp32: rounding an already-rounded tensor is a no-op, which would + collapse the floor to zero and turn the criterion below into an impossible bound. + """ + exact, exact_lse = _reference(q32, k32, v32, scale, mask) + lossy, lossy_lse = _reference( + q32.to(dtype).float(), k32.to(dtype).float(), v32.to(dtype).float(), scale, mask + ) + return ( + (exact - lossy).abs().max().item(), + (exact_lse - lossy_lse).abs().max().item(), + exact, + exact_lse, + ) + + +@pytest.mark.parametrize("shape", _SHAPES, ids=lambda s: "b%d_hq%d_hkv%d_sq%d_skv%d_d%d" % s) +@pytest.mark.parametrize("mask", ["no_mask", "causal", "causal_bottom_right"]) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_frost_forward_matches_fp32_reference(shape, mask, dtype): + """Forward output and LSE against an independent fp32 reference.""" + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + frost_attn_fwd, + ) + + b, hq, hkv, sq, skv, d = shape + torch.manual_seed(0) + # Generate in fp32 so there is a true high-precision original to measure against, then cast + # for the kernel. [b, h, s, d] views over bshd-contiguous memory is what the backend consumes. + mk = lambda s_, h_: torch.randn(b, s_, h_, d, device="cuda").permute(0, 2, 1, 3).contiguous() + q32, k32, v32 = mk(sq, hq), mk(skv, hkv), mk(skv, hkv) + q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) + scale = 1.0 / math.sqrt(d) + + out, lse = frost_attn_fwd(q, k, v, attn_scale=scale, attn_mask_type=mask) + + floor_o, floor_l, ref_o, ref_lse = _floor(q32, k32, v32, scale, mask, dtype) + err_o = (out.float() - ref_o).abs().max().item() + err_l = (lse.float() - ref_lse).abs().max().item() + + assert torch.isfinite(out).all(), "forward produced non-finite values" + # A floor of exactly zero would make the ratio meaningless; guard with a small absolute term. + assert err_o <= 2 * floor_o + 1e-3, "out err %.3e exceeds 2x the %s floor %.3e" % ( + err_o, + dtype, + floor_o, + ) + # The LSE convention is what the CP ring correction depends on, so check it explicitly: a + # log2-based or unscaled LSE would still give a plausible-looking output above. + assert err_l <= 2 * floor_l + 1e-3, "lse err %.3e exceeds 2x the %s floor %.3e" % ( + err_l, + dtype, + floor_l, + ) + assert lse.shape == (b, hq, sq), "lse must be [b, h, s]; got %s" % (tuple(lse.shape),) + assert lse.dtype == torch.float32, "lse must be fp32; got %s" % lse.dtype + + +@pytest.mark.parametrize("shape", _SHAPES[:2], ids=lambda s: "b%d_hq%d_hkv%d_sq%d_skv%d_d%d" % s) +@pytest.mark.parametrize("mask", ["no_mask", "causal"]) +def test_frost_backward_matches_fp32_reference(shape, mask): + """dq/dk/dv against autograd on the same independent fp32 reference.""" + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + frost_attn_bwd, + frost_attn_fwd, + ) + + b, hq, hkv, sq, skv, d = shape + dtype = torch.bfloat16 + torch.manual_seed(0) + mk = lambda s_, h_: torch.randn(b, s_, h_, d, device="cuda").permute(0, 2, 1, 3).contiguous() + q32, k32, v32 = mk(sq, hq), mk(skv, hkv), mk(skv, hkv) + q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) + scale = 1.0 / math.sqrt(d) + + out, lse = frost_attn_fwd(q, k, v, attn_scale=scale, attn_mask_type=mask) + dout = torch.randn_like(out) + dq, dk, dv = frost_attn_bwd(q, k, v, out, lse, dout, attn_scale=scale, attn_mask_type=mask) + + qr = q32.detach().clone().requires_grad_(True) + kr = k32.detach().clone().requires_grad_(True) + vr = v32.detach().clone().requires_grad_(True) + ref_o, _ = _reference(qr, kr, vr, scale, mask) + ref_o.backward(dout.float()) + + for name, got, want in (("dq", dq, qr.grad), ("dk", dk, kr.grad), ("dv", dv, vr.grad)): + assert torch.isfinite(got).all(), "%s has non-finite values" % name + assert got.shape == want.shape, "%s shape %s != %s" % (name, got.shape, want.shape) + err = (got.float() - want).abs().max().item() + # Gradients accumulate over the sequence, so scale the bar with skv rather than reusing + # the forward's floor. This is a sanity bound on systematic error, not a tight check. + assert err <= 0.05 * want.abs().max().item() + 1e-2, "%s max|err|=%.3e vs ref max %.3e" % ( + name, + err, + want.abs().max().item(), + ) + + +def test_frost_declines_unsupported_configs(): + """The selector must decline what the kernels do not serve, rather than computing wrongly.""" + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + is_frost_attention_supported, + ) + + base = dict(head_dim_qk=512, head_dim_v=512, qkv_dtype=torch.bfloat16, attn_mask_type="causal") + assert is_frost_attention_supported(**base)[0], "the supported case must be accepted" + + for override, why in ( + (dict(head_dim_qk=256, head_dim_v=256), "head_dim at the exclusive lower bound"), + (dict(head_dim_v=256), "asymmetric head_dim"), + (dict(qkv_dtype=torch.float32), "fp32"), + (dict(dropout=0.1), "dropout"), + (dict(attn_bias_type="post_scale_bias"), "attention bias"), + (dict(attn_mask_type="padding_causal"), "padding mask"), + (dict(attn_mask_type="arbitrary"), "arbitrary mask"), + ): + cfg = dict(base) + cfg.update(override) + ok, reason = is_frost_attention_supported(**cfg) + assert not ok, "%s must be declined" % why + assert reason, "a decline must explain itself" + + +def test_frost_rejects_mismatched_kv(): + """k and v must agree: the graphs declare v with k's shape and stride.""" + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + frost_attn_fwd, + ) + + b, h, s, d = 2, 4, 512, 512 + dtype = torch.bfloat16 + mk = lambda hh: torch.randn(b, s, hh, d, device="cuda", dtype=dtype).permute(0, 2, 1, 3) + q, k = mk(h).contiguous(), mk(h).contiguous() + + with pytest.raises(ValueError, match="same shape"): + frost_attn_fwd(q, k, mk(h * 2).contiguous()) + with pytest.raises(ValueError, match="same layout"): + # Same shape, different stride order: a cache hit would otherwise run a graph built for + # k's layout over v's memory and read the wrong elements silently. + v_odd = torch.randn(b, h, s, d, device="cuda", dtype=dtype) + frost_attn_fwd(q, k, v_odd) + with pytest.raises(ValueError, match="match q"): + frost_attn_fwd(q, k, k.to(torch.float32)) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 038f83c746a..b35d1b6797c 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -6,10 +6,17 @@ Why this exists. Gemma-4 global layers use symmetric head_dim=512, and no backend TE can select today serves both that head dim and context parallelism: FlashAttention 2/3 cap at 256, FA4 is -gated off at symmetric 512, the C++ cuDNN fused path caps at 256, and the unfused path supports -512 but cannot do CP. cuDNN Frontend 1.29.0 ships CuTe-DSL ("FROST") SDPA kernels that do serve -symmetric 512 forward and backward on Blackwell, reachable through the ordinary cuDNN graph API. -This module wraps them so TE, including its CP ring, can dispatch to them. +gated off at symmetric 512, the C++ cuDNN fused path is refused a graph by cuDNN above 256, and +the unfused path supports 512 but cannot do CP. cuDNN Frontend 1.29.0 ships CuTe-DSL ("FROST") +SDPA kernels that do serve symmetric 512 forward and backward on Blackwell. + +Why a separate Python backend rather than teaching the existing C++ fused path. The 256 ceiling +there is not a TE check -- the f16 dispatch applies no head-dim test and simply asks cuDNN to +build a graph -- so the natural question is why the new engines cannot just be picked up. They +cannot: FROST engines are registered at Python import time behind +CUDNN_FRONTEND_ENABLE_FROST_ENGINES and require the nvidia-cutlass-dsl Python package, while +TE's C++ builds against cuDNN Frontend headers only. Reaching them therefore requires a Python +graph, which is what this module is. Three properties of these kernels were verified on Blackwell before this was written, and each one constrains the code: @@ -78,7 +85,12 @@ def _import_cudnn(): - """Import cuDNN Frontend with FROST engines enabled, once.""" + """Import cuDNN Frontend with FROST engines enabled, once. + + Note the ordering hazard: the engines register at import time, so if another module imported + cudnn first without the switch set, setdefault here is too late and no FROST engine exists. + _select_frost_plan catches that by checking the plan name, but only once a plan is built. + """ global _cudnn if _cudnn is None: # Must be set before the import: the engines are registered at import time. @@ -161,6 +173,10 @@ def _no(reason): if not torch.cuda.is_available(): return _no("no CUDA device") + if os.environ.get("CUDNN_FRONTEND_ENABLE_FROST_ENGINES", "1") == "0": + # Explicitly switched off. Declining here is the difference between falling back cleanly + # and raising from _select_frost_plan once a plan is built. + return _no("CUDNN_FRONTEND_ENABLE_FROST_ENGINES=0 disables the FROST engines") if torch.cuda.get_device_capability() not in _SUPPORTED_ARCHS: return _no( "cuDNN FROST head_dim>256 kernels are SM100/SM103 only; found sm%d%d" @@ -236,10 +252,15 @@ def is_frost_attention_supported( dropout: float = 0.0, attn_bias_type: str = "no_bias", ) -> Tuple[bool, str]: - """Whether this specific attention configuration should route to FROST.""" - ok, reason = is_frost_attention_available() - if not ok: - return False, reason + """Whether this specific attention configuration should route to FROST. + + Shape and dtype are checked before availability, and the ordering is deliberate rather than + stylistic. Probing availability imports cuDNN Frontend and sets + CUDNN_FRONTEND_ENABLE_FROST_ENGINES, which registers extra engines process-wide and so is + visible to every other cuDNN consumer in the process. This function runs for every attention + config on the machine, the vast majority of which are nowhere near head_dim 512, and none of + them should pay that cost or have their engine pool changed underneath them. + """ if head_dim_qk != head_dim_v: return False, "FROST path requires symmetric head_dim; got %d/%d" % ( head_dim_qk, @@ -257,6 +278,9 @@ def is_frost_attention_supported( _mask_mode(attn_mask_type) except NotImplementedError as exc: return False, str(exc) + ok, reason = is_frost_attention_available() + if not ok: + return False, reason return True, "" diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 6674654996f..bf6a01d1d58 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -1866,9 +1866,9 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt FlashAttentionUtils.warning_printed = True # cuDNN FROST (CuTe-DSL SDPA in cuDNN Frontend >= 1.29.0) is the only backend that serves # symmetric head_dim in (256, 512] on SM100/SM103. Every other option stops short: FA2/FA3 - # cap at 256, FA4 is disabled at symmetric 512 above, the C++ cuDNN fused path caps at 256, - # and UnfusedDotProductAttention supports 512 but not context parallelism. Without this, - # Gemma-4 global layers with CP > 1 select no backend at all. + # cap at 256, FA4 is disabled at symmetric 512 above, the C++ cuDNN fused path is refused a + # graph by cuDNN above 256, and UnfusedDotProductAttention supports 512 but not context + # parallelism. Without this, Gemma-4 global layers with CP > 1 select no backend at all. if use_frost_attention: # Local import: frost_attention pulls in cudnn lazily, so this stays cheap and keeps # TE importable on systems without cudnn-frontend installed. From c5f9cd2eeda6f2cc2e2281f3264caae1b364cff5 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 00:36:27 -0700 Subject: [PATCH 14/97] feat(attention): honour deterministic on the FROST path, and document 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 Signed-off-by: Nitin Vegesna --- docs/envvars.rst | 15 +++++++-- .../pytorch/attention/test_frost_attention.py | 6 ++++ .../dot_product_attention/backends.py | 10 +++++- .../dot_product_attention/context_parallel.py | 31 ++++++++++++++++--- .../dot_product_attention/frost_attention.py | 15 ++++++--- .../attention/dot_product_attention/utils.py | 8 ++++- 6 files changed, 71 insertions(+), 14 deletions(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index 0fee105fd05..1e04f69790a 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -149,9 +149,12 @@ Then it applies a performance-based preference order among the remaining eligibl In PyTorch, the broad preference order is ``FlashAttention > FusedAttention > UnfusedDotProductAttention`` on supported pre-Hopper GPUs such as Ampere/Ada, and ``FusedAttention > FlashAttention > UnfusedDotProductAttention`` on Hopper and newer GPUs, -including Blackwell. In JAX, Transformer Engine uses cuDNN fused attention when -``NVTE_FUSED_ATTN=1`` and an eligible cuDNN kernel is available; otherwise it falls back to the -JAX-native implementation. See :doc:`examples/attention/attention` for a longer +including Blackwell. On Blackwell SM100/SM103 the order is ``FusedAttention > FlashAttention > +FrostAttention > UnfusedDotProductAttention``; FrostAttention only becomes eligible for +symmetric ``head_dim`` in (256, 512], which flash and fused attention do not serve, so it never +displaces a backend that could otherwise have run. In JAX, Transformer Engine uses cuDNN fused +attention when ``NVTE_FUSED_ATTN=1`` and an eligible cuDNN kernel is available; otherwise it +falls back to the JAX-native implementation. See :doc:`examples/attention/attention` for a longer backend-selection overview. .. envvar:: NVTE_FLASH_ATTN @@ -184,6 +187,12 @@ backend-selection overview. :Default: ``1`` :Description: Enable or disable FusedAttention backend (cuDNN-based) for DotProductAttention. When set to ``0``, FusedAttention will not be used. +.. envvar:: NVTE_FROST_ATTN + + :Type: ``int`` (0 or 1) + :Default: ``1`` + :Description: Enable or disable FrostAttention backend (the cuDNN FROST CuTe-DSL SDPA kernels in cuDNN Frontend) for DotProductAttention. When set to ``0``, FrostAttention will not be used. FrostAttention is the only backend serving symmetric ``head_dim`` in (256, 512], and is limited to SM100/SM103 with BF16/FP16 inputs and ``nvidia-cudnn-frontend>=1.29.0`` and ``nvidia-cutlass-dsl>=4.7.0`` installed. It supports context parallelism with ``cp_comm_type`` of ``p2p``, ``all_gather`` or ``a2a``, and declines FP8, ``thd`` layouts, dropout, attention bias, sliding window, softcap, KV caching and ``max_logit``. + .. envvar:: NVTE_UNFUSED_ATTN :Type: ``int`` (0 or 1) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 7068ec835ae..7d400393d0f 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -18,6 +18,7 @@ """ import math +import os import pytest import torch @@ -40,6 +41,11 @@ def _frost_availability(): _SKIP = _frost_availability() +# Mirrors NVTE_GDN_TEST_REQUIRED in test_gdn_attention.py. These tests skip on any machine that +# cannot reach the backend, which on most CI hardware is every machine; setting this on a lane +# that is supposed to cover FROST turns a silent skip into a loud failure. +if os.getenv("NVTE_FROST_TEST_REQUIRED", "0") == "1" and _SKIP is not None: + raise RuntimeError("NVTE_FROST_TEST_REQUIRED=1, but FrostAttention is unavailable: %s" % _SKIP) pytestmark = pytest.mark.skipif(_SKIP is not None, reason=str(_SKIP)) # head_dim 512 is the whole point of the backend; 320 checks the interior of the (256, 512] range diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index 58ef53e5685..3ca708a7150 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -2295,7 +2295,9 @@ class FrostAttnFunc(torch.autograd.Function): """ @staticmethod - def forward(ctx, q, k, v, softmax_scale, attn_mask_type, qkv_format, is_training): + def forward( + ctx, q, k, v, softmax_scale, attn_mask_type, qkv_format, is_training, deterministic + ): # pylint: disable=missing-function-docstring from .frost_attention import ( # pylint: disable=import-outside-toplevel frost_attn_fwd, @@ -2319,6 +2321,7 @@ def forward(ctx, q, k, v, softmax_scale, attn_mask_type, qkv_format, is_training ctx.attn_mask_type = attn_mask_type ctx.qkv_format = qkv_format ctx.unflattened_shape = out.shape + ctx.deterministic = deterministic # TE attention modules return the heads flattened into the last dimension # ([b, s, h*d] for bshd), matching FlashAttention and FusedAttention. Returning the # unflattened [b, s, h, d] makes autograd reject the incoming grad on shape mismatch. @@ -2346,7 +2349,10 @@ def backward(ctx, dout): to_frost_layout(dout.contiguous(), fmt), attn_scale=ctx.softmax_scale, attn_mask_type=ctx.attn_mask_type, + deterministic=ctx.deterministic, ) + # One None per non-tensor forward argument: softmax_scale, attn_mask_type, qkv_format, + # is_training, deterministic. Must track forward's signature exactly. return ( from_frost_layout(dq, fmt), from_frost_layout(dk, fmt), @@ -2355,6 +2361,7 @@ def backward(ctx, dout): None, None, None, + None, ) @@ -2449,6 +2456,7 @@ def forward( attn_mask_type, qkv_format, self.training, + self.deterministic, ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index 6cfc1e11ea5..43231e38b8e 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1347,6 +1347,7 @@ def cp_p2p_bwd_fused_attn( out_part, dout_part, section, + deterministic=False, ): """Per-tile backward call of CP P2P with FusedAttention backend""" aux_tensors = [softmax_lse, rng_states[cp_size - step - 1]] @@ -1467,6 +1468,7 @@ def cp_p2p_bwd_flash_attn( out_part, dout_part, section, + deterministic=False, ): """Per-tile backward call of CP P2P with FlashAttention backend""" if pad_between_seqs: @@ -1653,6 +1655,7 @@ def cp_ag_bwd_frost_attn( v_part, out_part, dout_part, + deterministic=False, ): """Per-step backward for CP all_gather with the cuDNN FROST backend.""" from .frost_attention import ( # pylint: disable=import-outside-toplevel @@ -1670,6 +1673,7 @@ def cp_ag_bwd_frost_attn( to_frost_layout(dout_part.contiguous(), qkv_format), attn_scale=softmax_scale, attn_mask_type=_frost_mask_for_window(window_size), + deterministic=deterministic, ) return ( from_frost_layout(dq, qkv_format), @@ -1702,7 +1706,7 @@ def cp_a2a_fwd_frost_attn(softmax_scale, attn_mask_type, qkv_format, q, k, v): def cp_a2a_bwd_frost_attn( - softmax_scale, attn_mask_type, qkv_format, softmax_lse, q, k, v, out, dout + softmax_scale, attn_mask_type, qkv_format, softmax_lse, q, k, v, out, dout, deterministic=False ): """Backward for CP a2a with the cuDNN FROST backend.""" from .frost_attention import ( # pylint: disable=import-outside-toplevel @@ -1720,6 +1724,7 @@ def cp_a2a_bwd_frost_attn( to_frost_layout(dout.contiguous(), qkv_format), attn_scale=softmax_scale, attn_mask_type=attn_mask_type, + deterministic=deterministic, ) return ( from_frost_layout(dq, qkv_format), @@ -1776,6 +1781,7 @@ def cp_p2p_bwd_frost_attn( out_part, dout_part, section, + deterministic=False, ): """Per-tile backward call of CP P2P with the cuDNN FROST backend. @@ -1797,6 +1803,7 @@ def cp_p2p_bwd_frost_attn( to_frost_layout(dout_part.contiguous(), qkv_format), attn_scale=softmax_scale, attn_mask_type=_frost_mask_for_section(attn_mask_type, section), + deterministic=deterministic, ) return ( from_frost_layout(dq, qkv_format), @@ -3110,7 +3117,10 @@ def backward(ctx, dout, *_args): prepare_outputs = cp_p2p_bwd_prepare_qkv(*prepare_inputs, section) if ctx.use_frost_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_frost_attn( - *frost_attn_inputs, *prepare_outputs, section + *frost_attn_inputs, + *prepare_outputs, + section, + deterministic=ctx.deterministic, ) elif ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( @@ -3127,7 +3137,10 @@ def backward(ctx, dout, *_args): prepare_outputs = cp_p2p_bwd_prepare_qkv(*prepare_inputs, section) if ctx.use_frost_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_frost_attn( - *frost_attn_inputs, *prepare_outputs, section + *frost_attn_inputs, + *prepare_outputs, + section, + deterministic=ctx.deterministic, ) elif ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( @@ -3144,7 +3157,10 @@ def backward(ctx, dout, *_args): prepare_outputs = cp_p2p_bwd_prepare_qkv(*prepare_inputs, section) if ctx.use_frost_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_frost_attn( - *frost_attn_inputs, *prepare_outputs, section + *frost_attn_inputs, + *prepare_outputs, + section, + deterministic=ctx.deterministic, ) elif ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( @@ -3161,7 +3177,10 @@ def backward(ctx, dout, *_args): prepare_outputs = cp_p2p_bwd_prepare_qkv(*prepare_inputs, section) if ctx.use_frost_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_frost_attn( - *frost_attn_inputs, *prepare_outputs, section + *frost_attn_inputs, + *prepare_outputs, + section, + deterministic=ctx.deterministic, ) elif ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( @@ -4567,6 +4586,7 @@ def backward(ctx, dout, *_args): v_part, out_part, dout_part, + deterministic=ctx.deterministic, ) elif ctx.use_fused_attention: # Set per-step parameters for THD @@ -5536,6 +5556,7 @@ def backward(ctx, dout, *_args): v, out, dout, + deterministic=ctx.deterministic, ) elif ctx.use_fused_attention: do_format = ctx.o_format diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index b35d1b6797c..11754ef7bc1 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -400,7 +400,9 @@ def _select_frost_plan(graph, token: str, what: str): def _build_fwd(key) -> dict: """Build (and JIT-compile) a forward graph. Expensive; always reached through the cache.""" cudnn = _import_cudnn() - *_device, b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks = key + # deterministic is unused here: it selects a backward algorithm. Callers pass False for the + # forward so the two never split the forward cache. + *_device, b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks, _deterministic = key io_dt = _cudnn_dtype(dtype) shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] @@ -440,7 +442,7 @@ def _build_fwd(key) -> dict: def _build_bwd(key) -> dict: """Build (and JIT-compile) a backward graph. Expensive; always reached through the cache.""" cudnn = _import_cudnn() - *_device, b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks = key + *_device, b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks, deterministic = key io_dt = _cudnn_dtype(dtype) shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] @@ -475,6 +477,7 @@ def _build_bwd(key) -> dict: dO=handles["do"], stats=handles["stats"], attn_scale=scale, + use_deterministic_algorithm=deterministic, **_MASK_MODES[mask], ) for tensor, stride in ((tdq, qs), (tdk, ks), (tdv, ks)): @@ -507,7 +510,7 @@ def _cached(kind: str, key): return entry -def _key(q, k, mask, scale): +def _key(q, k, mask, scale, deterministic=False): return ( # The graph is built under whichever device was current, so it must not be reused on # another one. Matches the C++ fused-attn cache, which keys on device_id for the same @@ -527,6 +530,9 @@ def _key(q, k, mask, scale): # lets bshd and sbhd both run without a transpose. tuple(q.stride()), tuple(k.stride()), + # The deterministic backward is a different algorithm, not a flag on the same one, so a + # plan built either way must not be handed to a call that asked for the other. + bool(deterministic), ) @@ -589,6 +595,7 @@ def frost_attn_bwd( dout: torch.Tensor, attn_scale: Optional[float] = None, attn_mask_type: str = "causal", + deterministic: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Backward attention via cuDNN FROST. `softmax_lse` is [b, h, s] as returned by the forward.""" for name, tensor in (("q", q), ("k", k), ("v", v), ("out", out), ("dout", dout)): @@ -621,7 +628,7 @@ def frost_attn_bwd( mask = _mask_mode(attn_mask_type) scale = attn_scale if attn_scale is not None else q.shape[-1] ** -0.5 - entry = _cached("bwd", _key(q, k, mask, scale)) + entry = _cached("bwd", _key(q, k, mask, scale, deterministic)) h = entry["handles"] if softmax_lse.dim() == 3: diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index bf6a01d1d58..11fcccf1b97 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -481,6 +481,8 @@ def get_attention_backend( available_backends : List[bool] All available backends that could support the provided input. A list of Booleans in the form of [use_flash_attention, use_fused_attention, use_unfused_attention]. + FrostAttention is deliberately not a member: the list's length is relied on by + existing three-way unpacks. Use the `use_frost_attention` return value instead. """ # NOTE: As part of refactoring attention.py, populating the _attention_backends cache in attention # is no longer performed at the end of get_attention_backend(), but the responsibility of doing so @@ -1970,7 +1972,7 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt logger.debug( "Available backends = {FlashAttention=%s%s, FusedAttention=%s%s," - " UnfusedDotProductAttention=%s}", + " UnfusedDotProductAttention=%s, FrostAttention=%s}", bool(available_backends[0]), (f" ({str(flash_attention_backend)})" if flash_attention_backend is not None else ""), bool(available_backends[1]), @@ -1980,6 +1982,10 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt else "" ), bool(available_backends[2]), + # Read from the local flag rather than available_backends, which excludes FROST by + # design. Without this the log reports every backend as unavailable and then selects + # FrostAttention a few lines later, which reads as a contradiction. + bool(use_frost_attention), ) # Select FusedAttention for performance From a97496ebcd97b7e42461305bf3d75dcbb1ac4c13 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 00:47:55 -0700 Subject: [PATCH 15/97] docs(attention): scope the FROST exclusivity claim to context parallelism 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 Signed-off-by: Nitin Vegesna --- docs/envvars.rst | 12 ++++++------ .../pytorch/attention/dot_product_attention/utils.py | 9 +++++---- 2 files changed, 11 insertions(+), 10 deletions(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index 1e04f69790a..9fe25fbad85 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -151,11 +151,11 @@ UnfusedDotProductAttention`` on supported pre-Hopper GPUs such as Ampere/Ada, an ``FusedAttention > FlashAttention > UnfusedDotProductAttention`` on Hopper and newer GPUs, including Blackwell. On Blackwell SM100/SM103 the order is ``FusedAttention > FlashAttention > FrostAttention > UnfusedDotProductAttention``; FrostAttention only becomes eligible for -symmetric ``head_dim`` in (256, 512], which flash and fused attention do not serve, so it never -displaces a backend that could otherwise have run. In JAX, Transformer Engine uses cuDNN fused -attention when ``NVTE_FUSED_ATTN=1`` and an eligible cuDNN kernel is available; otherwise it -falls back to the JAX-native implementation. See :doc:`examples/attention/attention` for a longer -backend-selection overview. +symmetric ``head_dim`` in (256, 512], which flash and fused attention do not serve, so the +backend it can displace is UnfusedDotProductAttention. In JAX, Transformer Engine uses cuDNN +fused attention when ``NVTE_FUSED_ATTN=1`` and an eligible cuDNN kernel is available; otherwise +it falls back to the JAX-native implementation. See :doc:`examples/attention/attention` for a +longer backend-selection overview. .. envvar:: NVTE_FLASH_ATTN @@ -191,7 +191,7 @@ backend-selection overview. :Type: ``int`` (0 or 1) :Default: ``1`` - :Description: Enable or disable FrostAttention backend (the cuDNN FROST CuTe-DSL SDPA kernels in cuDNN Frontend) for DotProductAttention. When set to ``0``, FrostAttention will not be used. FrostAttention is the only backend serving symmetric ``head_dim`` in (256, 512], and is limited to SM100/SM103 with BF16/FP16 inputs and ``nvidia-cudnn-frontend>=1.29.0`` and ``nvidia-cutlass-dsl>=4.7.0`` installed. It supports context parallelism with ``cp_comm_type`` of ``p2p``, ``all_gather`` or ``a2a``, and declines FP8, ``thd`` layouts, dropout, attention bias, sliding window, softcap, KV caching and ``max_logit``. + :Description: Enable or disable FrostAttention backend (the cuDNN FROST CuTe-DSL SDPA kernels in cuDNN Frontend) for DotProductAttention. When set to ``0``, FrostAttention will not be used. It is the only backend serving symmetric ``head_dim`` in (256, 512] together with context parallelism; without context parallelism UnfusedDotProductAttention also covers that range, and FrostAttention is preferred over it where both are eligible. It is limited to SM100/SM103 with BF16/FP16 inputs and ``nvidia-cudnn-frontend>=1.29.0`` and ``nvidia-cutlass-dsl>=4.7.0`` installed. It supports context parallelism with ``cp_comm_type`` of ``p2p``, ``all_gather`` or ``a2a``, and declines FP8, ``thd`` layouts, dropout, attention bias, sliding window, softcap, KV caching and ``max_logit``. .. envvar:: NVTE_UNFUSED_ATTN diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 11fcccf1b97..1c7aadbf930 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -1867,10 +1867,11 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt ) FlashAttentionUtils.warning_printed = True # cuDNN FROST (CuTe-DSL SDPA in cuDNN Frontend >= 1.29.0) is the only backend that serves - # symmetric head_dim in (256, 512] on SM100/SM103. Every other option stops short: FA2/FA3 - # cap at 256, FA4 is disabled at symmetric 512 above, the C++ cuDNN fused path is refused a - # graph by cuDNN above 256, and UnfusedDotProductAttention supports 512 but not context - # parallelism. Without this, Gemma-4 global layers with CP > 1 select no backend at all. + # symmetric head_dim in (256, 512] with context parallelism on SM100/SM103. Every other option + # stops short: FA2/FA3 cap at 256, FA4 is disabled at symmetric 512 above, the C++ cuDNN fused + # path is refused a graph by cuDNN above 256, and UnfusedDotProductAttention supports 512 but + # not context parallelism. Without this, Gemma-4 global layers with CP > 1 select no backend at + # all. if use_frost_attention: # Local import: frost_attention pulls in cudnn lazily, so this stays cheap and keeps # TE importable on systems without cudnn-frontend installed. From 5a675c2d2092f8784e79ac0c7b3c02d37ac27a5e Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 08:21:08 -0700 Subject: [PATCH 16/97] fix(attention): drop a duplicate deterministic parameter on the fused 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 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/dot_product_attention/context_parallel.py | 1 - 1 file changed, 1 deletion(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index 43231e38b8e..beb0ee20fc9 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1347,7 +1347,6 @@ def cp_p2p_bwd_fused_attn( out_part, dout_part, section, - deterministic=False, ): """Per-tile backward call of CP P2P with FusedAttention backend""" aux_tensors = [softmax_lse, rng_states[cp_size - step - 1]] From e82981f128b23711ef5ca12801b5e20e45157d96 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 08:34:42 -0700 Subject: [PATCH 17/97] fix(attention): decline FROST when determinism is required 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 Signed-off-by: Nitin Vegesna --- tests/pytorch/attention/test_frost_attention.py | 7 +++++-- .../pytorch/attention/dot_product_attention/utils.py | 9 +++++++++ 2 files changed, 14 insertions(+), 2 deletions(-) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 7d400393d0f..377eaf7f609 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -217,8 +217,11 @@ def test_frost_rejects_mismatched_kv(): frost_attn_fwd(q, k, mk(h * 2).contiguous()) with pytest.raises(ValueError, match="same layout"): # Same shape, different stride order: a cache hit would otherwise run a graph built for - # k's layout over v's memory and read the wrong elements silently. - v_odd = torch.randn(b, h, s, d, device="cuda", dtype=dtype) + # k's layout over v's memory and read the wrong elements silently. Build it as sbhd and + # permute, so the strides genuinely differ -- a [b, h, s, d] contiguous tensor would come + # out with exactly k's strides and prove nothing. + v_odd = torch.randn(s, b, h, d, device="cuda", dtype=dtype).permute(1, 2, 0, 3) + assert v_odd.shape == k.shape and v_odd.stride() != k.stride() frost_attn_fwd(q, k, v_odd) with pytest.raises(ValueError, match="match q"): frost_attn_fwd(q, k, k.to(torch.float32)) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 1c7aadbf930..34774521199 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -1910,6 +1910,15 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt # needs cu_seqlens plumbing that is neither implemented nor validated here. logger.debug("Disabling FrostAttention for qkv_layout = %s", qkv_layout) use_frost_attention = False + if use_frost_attention and deterministic and is_training: + # Measured on B200 with cuDNN Frontend 1.29.0: requesting a deterministic backward is + # refused outright -- cudnnGraphNotSupportedError, no engine proposes a plan -- so unlike + # the C++ fused path there is nothing to opt into. Declining keeps + # NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 an honest guarantee instead of silently running the + # non-deterministic kernel. The graph still passes the flag, so this lifts on its own if + # cuDNN ships a deterministic d512 backward. + logger.debug("Disabling FrostAttention as its backward has no deterministic cuDNN plan") + use_frost_attention = False if use_frost_attention and return_max_logit: # FrostAttention returns the context layer alone, where UnfusedDotProductAttention returns # (context, max_logit). Selecting it here would break the caller's unpack. From 3edba8598970785d20a7f7bb83c9fb357e71cea9 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 09:07:20 -0700 Subject: [PATCH 18/97] test(attention): make the FROST oracle float64, since an fp32 one is 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 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_frost_attention.py | 44 ++++++++++++------- 1 file changed, 28 insertions(+), 16 deletions(-) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 377eaf7f609..34e7bfff211 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -9,12 +9,17 @@ backend*, which validates the ring plumbing and nothing about the kernel: a systematic error -- a wrong softmax scale, a causal mask anchored to the wrong corner, an LSE in the wrong log base -- appears identically on both sides and cancels. Everything here is anchored to an independent -fp32 reference instead. +float64 reference instead. -The pass criterion is the one FlashAttention applies to itself: the kernel's error against an -fp32 reference must stay within 2x the error that comes from feeding the same reference bf16 +The pass criterion is the one FlashAttention applies to itself: the kernel's error against that +reference must stay within 2x the error the reference itself incurs from reduced-precision inputs. That floor is measured per case rather than hard-coded, so the bar tracks the shape and dtype instead of encoding a number that silently rots. + +The reference is float64, not float32. torch uses TF32 for fp32 matmuls on Ampere and newer, and +TF32's significand is 11 bits -- the same as fp16 -- so an fp32 reference is no more accurate +than an fp16 kernel and the floor collapses to nothing. Measured on B200: the fp16 floor came out +at 3e-08 instead of ~1e-03, which turned the bound into the bare absolute slack. """ import math @@ -60,8 +65,15 @@ def _frost_availability(): def _reference(q, k, v, scale, mask): - """Attention in fp32, computed independently of TE and of cuDNN.""" - qq, kk, vv = q.float(), k.float(), v.float() + """Attention in float64, computed independently of TE and of cuDNN. + + float64 rather than float32 on purpose. torch uses TF32 for fp32 matmuls on Ampere and newer, + and TF32 carries an 11-bit significand -- the same as fp16. An fp32 reference is therefore no + more accurate than the fp16 kernel it is meant to judge, which silently collapses the error + floor below and makes the comparison meaningless. float64 is immune to that and to whatever + the ambient TF32 flags happen to be. + """ + qq, kk, vv = q.double(), k.double(), v.double() rep = qq.shape[1] // kk.shape[1] kk = kk.repeat_interleave(rep, dim=1) vv = vv.repeat_interleave(rep, dim=1) @@ -80,12 +92,12 @@ def _reference(q, k, v, scale, mask): def _floor(q32, k32, v32, scale, mask, dtype): """The error `dtype` inputs alone cause, and the exact answer to measure the kernel against. - The inputs must originate in fp32: rounding an already-rounded tensor is a no-op, which would - collapse the floor to zero and turn the criterion below into an impossible bound. + The inputs must originate in higher precision: rounding an already-rounded tensor is a no-op, + which would collapse the floor to zero and turn the criterion below into an impossible bound. """ exact, exact_lse = _reference(q32, k32, v32, scale, mask) lossy, lossy_lse = _reference( - q32.to(dtype).float(), k32.to(dtype).float(), v32.to(dtype).float(), scale, mask + q32.to(dtype).double(), k32.to(dtype).double(), v32.to(dtype).double(), scale, mask ) return ( (exact - lossy).abs().max().item(), @@ -98,8 +110,8 @@ def _floor(q32, k32, v32, scale, mask, dtype): @pytest.mark.parametrize("shape", _SHAPES, ids=lambda s: "b%d_hq%d_hkv%d_sq%d_skv%d_d%d" % s) @pytest.mark.parametrize("mask", ["no_mask", "causal", "causal_bottom_right"]) @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) -def test_frost_forward_matches_fp32_reference(shape, mask, dtype): - """Forward output and LSE against an independent fp32 reference.""" +def test_frost_forward_matches_reference(shape, mask, dtype): + """Forward output and LSE against an independent float64 reference.""" from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( frost_attn_fwd, ) @@ -116,8 +128,8 @@ def test_frost_forward_matches_fp32_reference(shape, mask, dtype): out, lse = frost_attn_fwd(q, k, v, attn_scale=scale, attn_mask_type=mask) floor_o, floor_l, ref_o, ref_lse = _floor(q32, k32, v32, scale, mask, dtype) - err_o = (out.float() - ref_o).abs().max().item() - err_l = (lse.float() - ref_lse).abs().max().item() + err_o = (out.double() - ref_o).abs().max().item() + err_l = (lse.double() - ref_lse).abs().max().item() assert torch.isfinite(out).all(), "forward produced non-finite values" # A floor of exactly zero would make the ratio meaningless; guard with a small absolute term. @@ -139,8 +151,8 @@ def test_frost_forward_matches_fp32_reference(shape, mask, dtype): @pytest.mark.parametrize("shape", _SHAPES[:2], ids=lambda s: "b%d_hq%d_hkv%d_sq%d_skv%d_d%d" % s) @pytest.mark.parametrize("mask", ["no_mask", "causal"]) -def test_frost_backward_matches_fp32_reference(shape, mask): - """dq/dk/dv against autograd on the same independent fp32 reference.""" +def test_frost_backward_matches_reference(shape, mask): + """dq/dk/dv against autograd on the same independent float64 reference.""" from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( frost_attn_bwd, frost_attn_fwd, @@ -162,12 +174,12 @@ def test_frost_backward_matches_fp32_reference(shape, mask): kr = k32.detach().clone().requires_grad_(True) vr = v32.detach().clone().requires_grad_(True) ref_o, _ = _reference(qr, kr, vr, scale, mask) - ref_o.backward(dout.float()) + ref_o.backward(dout.double()) for name, got, want in (("dq", dq, qr.grad), ("dk", dk, kr.grad), ("dv", dv, vr.grad)): assert torch.isfinite(got).all(), "%s has non-finite values" % name assert got.shape == want.shape, "%s shape %s != %s" % (name, got.shape, want.shape) - err = (got.float() - want).abs().max().item() + err = (got.double() - want).abs().max().item() # Gradients accumulate over the sequence, so scale the bar with skv rather than reusing # the forward's floor. This is a sanity bound on systematic error, not a tight check. assert err <= 0.05 * want.abs().max().item() + 1e-2, "%s max|err|=%.3e vs ref max %.3e" % ( From 064396e2f8c2c5335665ca2694f3160204b9cec1 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 09:36:53 -0700 Subject: [PATCH 19/97] feat(attention): express FROST masking as a diagonal band, adding sliding 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 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_frost_attention.py | 57 ++++++++-- .../dot_product_attention/backends.py | 24 ++++- .../dot_product_attention/context_parallel.py | 19 ++-- .../dot_product_attention/frost_attention.py | 100 ++++++++++++------ .../attention/dot_product_attention/utils.py | 4 +- 5 files changed, 151 insertions(+), 53 deletions(-) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 34e7bfff211..62ee83d21ed 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -64,7 +64,7 @@ def _frost_availability(): ] -def _reference(q, k, v, scale, mask): +def _reference(q, k, v, scale, mask, window=None): """Attention in float64, computed independently of TE and of cuDNN. float64 rather than float32 on purpose. torch uses TF32 for fp32 matmuls on Ampere and newer, @@ -83,21 +83,27 @@ def _reference(q, k, v, scale, mask): # Top-left for "causal", bottom-right for "causal_bottom_right". These coincide only when # sq == skv, which is exactly why _SHAPES includes a rectangular case. offset = 0 if mask == "causal" else skv - sq - causal = torch.ones(sq, skv, device=q.device, dtype=torch.bool).triu(offset + 1) - s = s.masked_fill(causal, float("-inf")) + blocked = torch.ones(sq, skv, device=q.device, dtype=torch.bool).triu(offset + 1) + if window is not None and window[0] != -1: + # A left window keeps only the most recent window[0] keys before the diagonal, so + # everything further back is masked as well. + blocked |= torch.ones(sq, skv, device=q.device, dtype=torch.bool).tril( + offset - window[0] - 1 + ) + s = s.masked_fill(blocked, float("-inf")) p = s.softmax(-1) return p @ vv, torch.logsumexp(s, dim=-1) -def _floor(q32, k32, v32, scale, mask, dtype): +def _floor(q32, k32, v32, scale, mask, dtype, window=None): """The error `dtype` inputs alone cause, and the exact answer to measure the kernel against. The inputs must originate in higher precision: rounding an already-rounded tensor is a no-op, which would collapse the floor to zero and turn the criterion below into an impossible bound. """ - exact, exact_lse = _reference(q32, k32, v32, scale, mask) + exact, exact_lse = _reference(q32, k32, v32, scale, mask, window) lossy, lossy_lse = _reference( - q32.to(dtype).double(), k32.to(dtype).double(), v32.to(dtype).double(), scale, mask + q32.to(dtype).double(), k32.to(dtype).double(), v32.to(dtype).double(), scale, mask, window ) return ( (exact - lossy).abs().max().item(), @@ -149,6 +155,45 @@ def test_frost_forward_matches_reference(shape, mask, dtype): assert lse.dtype == torch.float32, "lse must be fp32; got %s" % lse.dtype +@pytest.mark.parametrize("window", [(256, 0), (128, 0)], ids=lambda w: "win%d" % w[0]) +@pytest.mark.parametrize("mask", ["causal", "causal_bottom_right"]) +def test_frost_sliding_window_matches_reference(mask, window): + """Sliding window against the float64 reference. + + The engine advertises swa support, and cuDNN expresses a window as a left bound on the same + diagonal band that gives causal masking, so this shares a code path with the cases above. It + is worth its own test because a left bound that is off by one, or silently dropped, still + produces finite plausible-looking output -- the reference is the only thing that catches it. + """ + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + frost_attn_fwd, + ) + + b, hq, hkv, sq, skv, d = 2, 8, 4, 1024, 1024, 512 + dtype = torch.bfloat16 + torch.manual_seed(0) + mk = lambda s_, h_: torch.randn(b, s_, h_, d, device="cuda").permute(0, 2, 1, 3).contiguous() + q32, k32, v32 = mk(sq, hq), mk(skv, hkv), mk(skv, hkv) + q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) + scale = 1.0 / math.sqrt(d) + + out, _ = frost_attn_fwd(q, k, v, attn_scale=scale, attn_mask_type=mask, window_size=window) + + floor_o, _, ref_o, _ = _floor(q32, k32, v32, scale, mask, dtype, window) + err = (out.double() - ref_o).abs().max().item() + assert torch.isfinite(out).all(), "sliding-window forward produced non-finite values" + assert err <= 2 * floor_o + 1e-3, "out err %.3e exceeds 2x the floor %.3e for window %s" % ( + err, + floor_o, + window, + ) + + # A window must actually change the result; if the bound were dropped this would match the + # unwindowed output and the check above would still pass. + full, _ = frost_attn_fwd(q, k, v, attn_scale=scale, attn_mask_type=mask) + assert not torch.equal(out, full), "window %s produced the same output as no window" % (window,) + + @pytest.mark.parametrize("shape", _SHAPES[:2], ids=lambda s: "b%d_hq%d_hkv%d_sq%d_skv%d_d%d" % s) @pytest.mark.parametrize("mask", ["no_mask", "causal"]) def test_frost_backward_matches_reference(shape, mask): diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index 3ca708a7150..e67d012b2ba 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -2296,7 +2296,16 @@ class FrostAttnFunc(torch.autograd.Function): @staticmethod def forward( - ctx, q, k, v, softmax_scale, attn_mask_type, qkv_format, is_training, deterministic + ctx, + q, + k, + v, + softmax_scale, + attn_mask_type, + qkv_format, + is_training, + deterministic, + window_size, ): # pylint: disable=missing-function-docstring from .frost_attention import ( # pylint: disable=import-outside-toplevel @@ -2312,13 +2321,19 @@ def forward( k_f = to_frost_layout(k.contiguous(), qkv_format) v_f = to_frost_layout(v.contiguous(), qkv_format) out_f, softmax_lse = frost_attn_fwd( - q_f, k_f, v_f, attn_scale=softmax_scale, attn_mask_type=attn_mask_type + q_f, + k_f, + v_f, + attn_scale=softmax_scale, + attn_mask_type=attn_mask_type, + window_size=window_size, ) out = from_frost_layout(out_f, qkv_format) if is_training: ctx.save_for_backward(q_f, k_f, v_f, out_f, softmax_lse) ctx.softmax_scale = softmax_scale ctx.attn_mask_type = attn_mask_type + ctx.window_size = window_size ctx.qkv_format = qkv_format ctx.unflattened_shape = out.shape ctx.deterministic = deterministic @@ -2349,10 +2364,11 @@ def backward(ctx, dout): to_frost_layout(dout.contiguous(), fmt), attn_scale=ctx.softmax_scale, attn_mask_type=ctx.attn_mask_type, + window_size=ctx.window_size, deterministic=ctx.deterministic, ) # One None per non-tensor forward argument: softmax_scale, attn_mask_type, qkv_format, - # is_training, deterministic. Must track forward's signature exactly. + # is_training, deterministic, window_size. Must track forward's signature exactly. return ( from_frost_layout(dq, fmt), from_frost_layout(dk, fmt), @@ -2362,6 +2378,7 @@ def backward(ctx, dout): None, None, None, + None, ) @@ -2457,6 +2474,7 @@ def forward( qkv_format, self.training, self.deterministic, + window_size, ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index beb0ee20fc9..482f669a284 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1607,12 +1607,11 @@ def _frost_mask_for_window(window_size): all_gather never produces that. """ if window_size is None or tuple(window_size) == (-1, -1): - return "no_mask" - if tuple(window_size) == (-1, 0): - return "causal_bottom_right" - raise NotImplementedError( - "FROST all_gather does not support sliding window %s" % str(window_size) - ) + return "no_mask", None + # Anything with a bounded side is causal relative to the trimmed KV, and a bounded left side + # is a sliding window. Both are expressed as a band against the bottom-right diagonal, so the + # window travels with the mask type rather than needing a separate spelling per case. + return "causal_bottom_right", tuple(window_size) def cp_ag_fwd_frost_attn( @@ -1634,12 +1633,14 @@ def cp_ag_fwd_frost_attn( to_frost_layout, ) + mask_type, window = _frost_mask_for_window(window_size) out, softmax_lse = frost_attn_fwd( to_frost_layout(q_part.contiguous(), qkv_format), to_frost_layout(k_part.contiguous(), qkv_format), to_frost_layout(v_part.contiguous(), qkv_format), attn_scale=softmax_scale, - attn_mask_type=_frost_mask_for_window(window_size), + attn_mask_type=mask_type, + window_size=window, ) return from_frost_layout(out, qkv_format), softmax_lse @@ -1663,6 +1664,7 @@ def cp_ag_bwd_frost_attn( to_frost_layout, ) + mask_type, window = _frost_mask_for_window(window_size) dq, dk, dv = frost_attn_bwd( to_frost_layout(q_part.contiguous(), qkv_format), to_frost_layout(k_part.contiguous(), qkv_format), @@ -1671,7 +1673,8 @@ def cp_ag_bwd_frost_attn( softmax_lse, to_frost_layout(dout_part.contiguous(), qkv_format), attn_scale=softmax_scale, - attn_mask_type=_frost_mask_for_window(window_size), + attn_mask_type=mask_type, + window_size=window, deterministic=deterministic, ) return ( diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 11754ef7bc1..ded9bb8c8ca 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -23,9 +23,9 @@ 1. cuDNN's `use_causal_mask` is TOP-LEFT aligned and `use_causal_mask_bottom_right` is bottom-right. They coincide when SQ == SKV, so the distinction is invisible in square tests - and decisive for all_gather, which trims KV. `_MASK_MODES` lists only spellings checked - against a reference for their alignment: sdpa() ignores unknown kwargs silently, so an - unverified name would apply no mask at all and still run. + and decisive for all_gather, which trims KV. Both alignments were checked against a + reference rather than assumed, and masking is built as a diagonal band so causal, + bottom-right and sliding window come from one mechanism instead of three spellings. 2. Plan building must be cached. Building a plan is by far the most expensive cuDNN frontend call here, and dominates an execute even after cuDNN has cached the JIT and made rebuilds @@ -77,6 +77,10 @@ _SUPPORTED_ARCHS = ((10, 0), (10, 3)) _MAX_HEAD_DIM = 512 _MIN_HEAD_DIM = 257 # below this the existing cuDNN/flash backends already serve the shape +# The engine pads head_dim to a multiple of 8, so 260 is not servable even though it is in range. +# Without this it passes the gate and then fails at plan selection with a message about missing +# engines, instead of declining cleanly here. +_HEAD_DIM_MULTIPLE = 8 _cudnn = None _availability: Optional[Tuple[bool, str]] = None @@ -213,35 +217,57 @@ def _no(reason): return _availability -# cuDNN sdpa() kwargs per TE mask type. +# TE mask types this backend serves. cuDNN expresses causal, bottom-right and sliding-window +# masking as ONE mechanism -- a diagonal alignment plus a two-sided band -- rather than three +# separate flags, so that is what _mask_options builds. The legacy spellings desugar into exactly +# that: pygraph/sdpa.cpp maps 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. Building the band directly is equivalent for those two and +# additionally expresses a left bound, which is what a sliding window is. # -# These exact spellings are behaviourally verified, which matters more than it sounds: sdpa() -# takes **kwargs and SILENTLY IGNORES names it does not recognise, so a typo here would apply no -# mask at all and still build and run. Do not add an entry without checking the output against a -# reference for that alignment. -# -# Both alignments are needed. The p2p ring produces square diagonal tiles (top-left and -# bottom-right coincide there), while all_gather trims KV and relies on bottom-right alignment, -# where the two differ completely. -_MASK_MODES = { - "no_mask": {}, - "causal": {"use_causal_mask": True}, - "causal_bottom_right": {"use_causal_mask_bottom_right": True}, -} +# Both alignments are needed. The p2p ring produces square diagonal tiles, where top-left and +# bottom-right coincide, while all_gather trims KV and relies on bottom-right alignment, where +# the two differ completely. +_SUPPORTED_MASKS = ("no_mask", "causal", "causal_bottom_right") +# Sliding window as TE spells it: (left, right), -1 meaning unbounded on that side. +_NO_WINDOW = (-1, -1) -def _mask_mode(attn_mask_type: str) -> str: - """Validate a TE mask type and return its key in _MASK_MODES. - Anything not listed is rejected rather than approximated: the failure mode of guessing wrong - is silent numerical corruption, not an exception. - """ - if attn_mask_type in _MASK_MODES: - return attn_mask_type - raise NotImplementedError( - "FROST attention supports attn_mask_type in %s; got %r. Padding variants need varlen" - " support that is not implemented here." % (sorted(_MASK_MODES), attn_mask_type) - ) +def _mask_spec(attn_mask_type: str, window_size=None): + """Validate a TE mask type and window, returning the hashable spec the plan is keyed on.""" + if attn_mask_type not in _SUPPORTED_MASKS: + raise NotImplementedError( + "FROST attention supports attn_mask_type in %s; got %r" + % (str(_SUPPORTED_MASKS), attn_mask_type) + ) + window = _NO_WINDOW if window_size is None else tuple(window_size) + if len(window) != 2: + raise NotImplementedError("window_size must be a (left, right) pair; got %r" % (window,)) + if window[1] not in (-1, 0): + # A right bound past the diagonal is future context. cuDNN can express it, but no TE mask + # type asks for it, so decline rather than guess the intent. + raise NotImplementedError("FROST attention does not support a right window %r" % (window,)) + return attn_mask_type, window + + +def _mask_options(cudnn, spec): + """cuDNN sdpa kwargs for a (mask type, window) spec: a diagonal alignment plus a band.""" + attn_mask_type, window = spec + left, right = window + options = {} + if attn_mask_type in ("causal", "causal_bottom_right") or right == 0: + options["diagonal_alignment"] = ( + cudnn.diagonal_alignment.BOTTOM_RIGHT + if attn_mask_type == "causal_bottom_right" + else cudnn.diagonal_alignment.TOP_LEFT + ) + options["diagonal_band_right_bound"] = 0 + if left != -1: + # cuDNN counts the diagonal itself, TE does not, hence the +1 -- the same convention the + # C++ fused path and the Python port both use. + options["diagonal_band_left_bound"] = left + 1 + return options def is_frost_attention_supported( @@ -251,6 +277,7 @@ def is_frost_attention_supported( attn_mask_type: str, dropout: float = 0.0, attn_bias_type: str = "no_bias", + window_size: Optional[Tuple[int, int]] = None, ) -> Tuple[bool, str]: """Whether this specific attention configuration should route to FROST. @@ -268,6 +295,11 @@ def is_frost_attention_supported( ) if not _MIN_HEAD_DIM <= head_dim_qk <= _MAX_HEAD_DIM: return False, "FROST path covers head_dim in (256, 512]; got %d" % head_dim_qk + if head_dim_qk % _HEAD_DIM_MULTIPLE != 0: + return False, "FROST path needs head_dim to be a multiple of %d; got %d" % ( + _HEAD_DIM_MULTIPLE, + head_dim_qk, + ) if qkv_dtype not in (torch.bfloat16, torch.float16): return False, "FROST path supports bf16/fp16; got %s" % qkv_dtype if dropout != 0.0: @@ -275,7 +307,7 @@ def is_frost_attention_supported( if attn_bias_type != "no_bias": return False, "FROST path does not support attention bias" try: - _mask_mode(attn_mask_type) + _mask_spec(attn_mask_type, window_size) except NotImplementedError as exc: return False, str(exc) ok, reason = is_frost_attention_available() @@ -422,7 +454,7 @@ def _build_fwd(key) -> dict: v=tv, generate_stats=True, # the CP ring needs the LSE, and it is cheap attn_scale=scale, - **_MASK_MODES[mask], + **_mask_options(cudnn, mask), ) tout.set_output(True).set_dim(shq).set_stride(list(qs)) # out mirrors q tlse.set_output(True).set_dim([b, hq, sq, 1]).set_stride([hq * sq, sq, 1, 1]).set_data_type( @@ -478,7 +510,7 @@ def _build_bwd(key) -> dict: stats=handles["stats"], attn_scale=scale, use_deterministic_algorithm=deterministic, - **_MASK_MODES[mask], + **_mask_options(cudnn, mask), ) for tensor, stride in ((tdq, qs), (tdk, ks), (tdv, ks)): tensor.set_output(True).set_data_type(io_dt).set_stride(list(stride)) @@ -542,6 +574,7 @@ def frost_attn_fwd( v: torch.Tensor, attn_scale: Optional[float] = None, attn_mask_type: str = "causal", + window_size: Optional[Tuple[int, int]] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: """Forward attention via cuDNN FROST. @@ -567,7 +600,7 @@ def frost_attn_fwd( % (q.shape[1], k.shape[1]) ) - mask = _mask_mode(attn_mask_type) + mask = _mask_spec(attn_mask_type, window_size) scale = attn_scale if attn_scale is not None else q.shape[-1] ** -0.5 entry = _cached("fwd", _key(q, k, mask, scale)) tq, tk, tv, tout, tlse = entry["handles"] @@ -595,6 +628,7 @@ def frost_attn_bwd( dout: torch.Tensor, attn_scale: Optional[float] = None, attn_mask_type: str = "causal", + window_size: Optional[Tuple[int, int]] = None, deterministic: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Backward attention via cuDNN FROST. `softmax_lse` is [b, h, s] as returned by the forward.""" @@ -626,7 +660,7 @@ def frost_attn_bwd( % (tuple(softmax_lse.shape), tuple(q.shape)) ) - mask = _mask_mode(attn_mask_type) + mask = _mask_spec(attn_mask_type, window_size) scale = attn_scale if attn_scale is not None else q.shape[-1] ** -0.5 entry = _cached("bwd", _key(q, k, mask, scale, deterministic)) h = entry["handles"] diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 34774521199..5096f052600 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -1886,6 +1886,7 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt attn_mask_type=attn_mask_type, dropout=attention_dropout, attn_bias_type=core_attention_bias_type, + window_size=window_size, ) if not frost_supported: logger.debug("Disabling FrostAttention: %s", frost_reason) @@ -1902,9 +1903,6 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt if use_frost_attention and softcap is not None and softcap != 0.0: logger.debug("Disabling FrostAttention for softcap") use_frost_attention = False - if use_frost_attention and window_size not in ((-1, -1), (-1, 0)): - logger.debug("Disabling FrostAttention for sliding window %s", str(window_size)) - use_frost_attention = False if use_frost_attention and "thd" in qkv_layout: # bshd and sbhd are served directly from their own strides; thd is packed/varlen, which # needs cu_seqlens plumbing that is neither implemented nor validated here. From 441dee46321abc706574cf9536f9d8fb9f74fa50 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 09:43:02 -0700 Subject: [PATCH 20/97] fix(attention): carry the sliding window through a2a, and decline it 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/context_parallel.py | 19 ++++++++++++++++--- .../attention/dot_product_attention/utils.py | 16 ++++++++++++++++ 2 files changed, 32 insertions(+), 3 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index 482f669a284..162c56ea132 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1684,7 +1684,7 @@ def cp_ag_bwd_frost_attn( ) -def cp_a2a_fwd_frost_attn(softmax_scale, attn_mask_type, qkv_format, q, k, v): +def cp_a2a_fwd_frost_attn(softmax_scale, attn_mask_type, qkv_format, q, k, v, window_size=None): """Forward for CP a2a with the cuDNN FROST backend. The simplest of the three. After the all-to-all each rank holds the FULL sequence for a subset @@ -1703,12 +1703,23 @@ def cp_a2a_fwd_frost_attn(softmax_scale, attn_mask_type, qkv_format, q, k, v): to_frost_layout(v.contiguous(), qkv_format), attn_scale=softmax_scale, attn_mask_type=attn_mask_type, + window_size=window_size, ) return from_frost_layout(out, qkv_format), softmax_lse def cp_a2a_bwd_frost_attn( - softmax_scale, attn_mask_type, qkv_format, softmax_lse, q, k, v, out, dout, deterministic=False + softmax_scale, + attn_mask_type, + qkv_format, + softmax_lse, + q, + k, + v, + out, + dout, + deterministic=False, + window_size=None, ): """Backward for CP a2a with the cuDNN FROST backend.""" from .frost_attention import ( # pylint: disable=import-outside-toplevel @@ -1726,6 +1737,7 @@ def cp_a2a_bwd_frost_attn( to_frost_layout(dout.contiguous(), qkv_format), attn_scale=softmax_scale, attn_mask_type=attn_mask_type, + window_size=window_size, deterministic=deterministic, ) return ( @@ -5150,7 +5162,7 @@ def forward( qkv_scale_inv_format = None if use_frost_attention: out_, softmax_lse = cp_a2a_fwd_frost_attn( - softmax_scale, attn_mask_type, qkv_format, q, k, v + softmax_scale, attn_mask_type, qkv_format, q, k, v, window_size=window_size ) # Only the LSE: FROST has no dropout, so there is no RNG state to carry, and a # None in this list would have to survive the save/restore machinery. @@ -5559,6 +5571,7 @@ def backward(ctx, dout, *_args): out, dout, deterministic=ctx.deterministic, + window_size=ctx.window_size, ) elif ctx.use_fused_attention: do_format = ctx.o_format diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 5096f052600..638797d6f35 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -1927,6 +1927,22 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt # Explicit anyway: no page table reaches the backend, so a paged cache would be read raw. logger.debug("Disabling FrostAttention for KV caching") use_frost_attention = False + if ( + use_frost_attention + and context_parallel + and window_size is not None + and (window_size[0] != -1 or window_size[1] not in [-1, 0]) + and cp_comm_type in ["p2p", "a2a+p2p"] + ): + # Same rule FusedAttention carries: the p2p ring shards KV across steps, so a left bound + # measured against the full sequence does not survive the per-step tiles. all_gather and + # a2a both see a contiguous KV range and do support it. + logger.debug( + "Disabling FrostAttention as it does not support context parallelism with sliding" + " window attention and cp_comm_type = %s", + cp_comm_type, + ) + use_frost_attention = False if use_frost_attention and context_parallel: # Same two restrictions FlashAttention and FusedAttention carry above. Both are about # where the causal diagonal sits: the ring shards q and kv independently, so a mask whose From 9770ca5f02c2e1258db46aac12028395a5e620e7 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 09:49:00 -0700 Subject: [PATCH 21/97] test(attention): cover the sliding window in backward, at its boundary, 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 Signed-off-by: Nitin Vegesna --- docs/envvars.rst | 2 +- .../pytorch/attention/test_frost_attention.py | 68 +++++++++++++++++-- 2 files changed, 64 insertions(+), 6 deletions(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index 9fe25fbad85..5707111f2a5 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -191,7 +191,7 @@ longer backend-selection overview. :Type: ``int`` (0 or 1) :Default: ``1`` - :Description: Enable or disable FrostAttention backend (the cuDNN FROST CuTe-DSL SDPA kernels in cuDNN Frontend) for DotProductAttention. When set to ``0``, FrostAttention will not be used. It is the only backend serving symmetric ``head_dim`` in (256, 512] together with context parallelism; without context parallelism UnfusedDotProductAttention also covers that range, and FrostAttention is preferred over it where both are eligible. It is limited to SM100/SM103 with BF16/FP16 inputs and ``nvidia-cudnn-frontend>=1.29.0`` and ``nvidia-cutlass-dsl>=4.7.0`` installed. It supports context parallelism with ``cp_comm_type`` of ``p2p``, ``all_gather`` or ``a2a``, and declines FP8, ``thd`` layouts, dropout, attention bias, sliding window, softcap, KV caching and ``max_logit``. + :Description: Enable or disable FrostAttention backend (the cuDNN FROST CuTe-DSL SDPA kernels in cuDNN Frontend) for DotProductAttention. When set to ``0``, FrostAttention will not be used. It is the only backend serving symmetric ``head_dim`` in (256, 512] together with context parallelism; without context parallelism UnfusedDotProductAttention also covers that range, and FrostAttention is preferred over it where both are eligible. 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 supports context parallelism with ``cp_comm_type`` of ``p2p``, ``all_gather`` or ``a2a``, and sliding-window attention with ``all_gather`` or ``a2a`` (declined with ``p2p``, whose ring shards KV across steps). It declines FP8, ``thd`` layouts, dropout, attention bias, softcap, KV caching, ``max_logit``, and deterministic execution, the last because cuDNN offers no deterministic backward for these kernels. .. envvar:: NVTE_UNFUSED_ATTN diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 62ee83d21ed..933c24e09f9 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -155,7 +155,7 @@ def test_frost_forward_matches_reference(shape, mask, dtype): assert lse.dtype == torch.float32, "lse must be fp32; got %s" % lse.dtype -@pytest.mark.parametrize("window", [(256, 0), (128, 0)], ids=lambda w: "win%d" % w[0]) +@pytest.mark.parametrize("window", [(256, 0), (128, 0), (0, 0)], ids=lambda w: "win%d" % w[0]) @pytest.mark.parametrize("mask", ["causal", "causal_bottom_right"]) def test_frost_sliding_window_matches_reference(mask, window): """Sliding window against the float64 reference. @@ -196,7 +196,8 @@ def test_frost_sliding_window_matches_reference(mask, window): @pytest.mark.parametrize("shape", _SHAPES[:2], ids=lambda s: "b%d_hq%d_hkv%d_sq%d_skv%d_d%d" % s) @pytest.mark.parametrize("mask", ["no_mask", "causal"]) -def test_frost_backward_matches_reference(shape, mask): +@pytest.mark.parametrize("window", [None, (128, 0)], ids=["nowin", "win128"]) +def test_frost_backward_matches_reference(shape, mask, window): """dq/dk/dv against autograd on the same independent float64 reference.""" from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( frost_attn_bwd, @@ -211,14 +212,16 @@ def test_frost_backward_matches_reference(shape, mask): q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) scale = 1.0 / math.sqrt(d) - out, lse = frost_attn_fwd(q, k, v, attn_scale=scale, attn_mask_type=mask) + out, lse = frost_attn_fwd(q, k, v, attn_scale=scale, attn_mask_type=mask, window_size=window) dout = torch.randn_like(out) - dq, dk, dv = frost_attn_bwd(q, k, v, out, lse, dout, attn_scale=scale, attn_mask_type=mask) + dq, dk, dv = frost_attn_bwd( + q, k, v, out, lse, dout, attn_scale=scale, attn_mask_type=mask, window_size=window + ) qr = q32.detach().clone().requires_grad_(True) kr = k32.detach().clone().requires_grad_(True) vr = v32.detach().clone().requires_grad_(True) - ref_o, _ = _reference(qr, kr, vr, scale, mask) + ref_o, _ = _reference(qr, kr, vr, scale, mask, window) ref_o.backward(dout.double()) for name, got, want in (("dq", dq, qr.grad), ("dk", dk, kr.grad), ("dv", dv, vr.grad)): @@ -251,6 +254,10 @@ def test_frost_declines_unsupported_configs(): (dict(attn_bias_type="post_scale_bias"), "attention bias"), (dict(attn_mask_type="padding_causal"), "padding mask"), (dict(attn_mask_type="arbitrary"), "arbitrary mask"), + # window_size reaches _mask_spec through is_frost_attention_supported, so its validation + # is part of the selector contract rather than an internal detail. + (dict(window_size=(-1, 5)), "a right window past the diagonal"), + (dict(window_size=(128,)), "a malformed window pair"), ): cfg = dict(base) cfg.update(override) @@ -259,6 +266,57 @@ def test_frost_declines_unsupported_configs(): assert reason, "a decline must explain itself" +@pytest.mark.parametrize( + "cp_comm_type,window,expect_frost", + [ + ("all_gather", (128, 0), True), + ("a2a", (128, 0), True), + ("p2p", (128, 0), False), + ("a2a+p2p", (128, 0), False), + ("p2p", (-1, 0), True), + ("p2p", (-1, -1), True), + ], +) +def test_frost_sliding_window_selection_by_cp_comm_type(cp_comm_type, window, expect_frost): + """Which context-parallel paths may serve a sliding window. + + all_gather and a2a each see a contiguous KV range, so the window applies unchanged. The p2p + ring shards KV across steps, so a bound measured against the full sequence does not survive + the per-step tiles -- the same rule FusedAttention carries. The cases without a real window + must still select FROST, since the decline has to key on the window and not on p2p itself. + """ + from transformer_engine.pytorch.attention.dot_product_attention.utils import ( + AttentionParams, + get_attention_backend, + ) + + params = AttentionParams( + qkv_dtype=torch.bfloat16, + qkv_layout="bshd_bshd_bshd", + batch_size=2, + num_heads=8, + num_gqa_groups=4, + max_seqlen_q=4096, + max_seqlen_kv=4096, + head_dim_qk=512, + head_dim_v=512, + attn_mask_type="causal", + window_size=window, + context_parallel=True, + cp_comm_type=cp_comm_type, + is_training=True, + ) + use_frost = get_attention_backend(params)[5] + assert ( + bool(use_frost) == expect_frost + ), "cp_comm_type=%s window=%s: expected use_frost_attention=%s, got %s" % ( + cp_comm_type, + window, + expect_frost, + bool(use_frost), + ) + + def test_frost_rejects_mismatched_kv(): """k and v must agree: the graphs declare v with k's shape and stride.""" from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( From ad9dfdc633187ca6198b96f90cc0c108da928718 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 09:52:54 -0700 Subject: [PATCH 22/97] fix(attention): let the CP sliding-window asserts know FROST exists Enabling sliding window made two pre-existing assertions reachable that had never heard of this backend. Both the all_gather and a2a forwards allow a window only if FusedAttention or some FlashAttention is in play, and when FROST is selected every one of those flags is False -- so the assert survived solely on fa_utils.v2_3_plus, which reports whether flash-attn happens to be installed rather than which backend is running. Sliding window with all_gather or a2a would therefore fail or pass on an unrelated package, on exactly the Blackwell d512 box this backend exists for. Both allowlists and both messages now include FROST. Also declines a windowed non-causal mask when the q and kv lengths differ. FROST anchors the band from the mask type, so that case always lands top-left, while TE's bottom_right_diagonal defaults to True and the C++ fused path picks the alignment from it -- the two would disagree silently. Declining is better than guessing the anchor. The remaining fixes are gate hygiene. window_size now rejects a left below -1, which would otherwise build diagonal_band_left_bound=-1 and fail at plan build, and a non-iterable window declines instead of raising TypeError out of backend selection, which is not an exception the selector catches. window_size moves after deterministic in frost_attn_bwd so an existing positional caller cannot silently reinterpret one as the other. Tests cover the new gates, including the head_dim multiple-of-8 rule, which had none. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- tests/pytorch/attention/test_frost_attention.py | 5 +++++ .../dot_product_attention/context_parallel.py | 13 ++++++++----- .../dot_product_attention/frost_attention.py | 17 ++++++++++++++--- .../attention/dot_product_attention/utils.py | 16 ++++++++++++++++ 4 files changed, 43 insertions(+), 8 deletions(-) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 933c24e09f9..bf904c8a732 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -258,6 +258,11 @@ def test_frost_declines_unsupported_configs(): # is part of the selector contract rather than an internal detail. (dict(window_size=(-1, 5)), "a right window past the diagonal"), (dict(window_size=(128,)), "a malformed window pair"), + (dict(window_size=(-2, 0)), "a left window below -1"), + (dict(window_size=7), "a non-iterable window"), + # The engine pads head_dim to a multiple of 8, so an in-range but unpadded dim has to be + # declined here rather than failing later at plan selection. + (dict(head_dim_qk=260, head_dim_v=260), "head_dim not a multiple of 8"), ): cfg = dict(base) cfg.update(override) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index 162c56ea132..4a6d50cea1a 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -3661,11 +3661,12 @@ def forward( or use_fused_attention or use_flash_attn_3 or use_flash_attn_4 + or use_frost_attention or fa_utils.v2_3_plus ), ( - "cp_comm_type='all_gather' only supports SWA through FusedAttention or FlashAttention" - f" >= 2.3. Found {use_fused_attention=}, {use_flash_attn_3=}, " - f"{use_flash_attn_4=}, " + "cp_comm_type='all_gather' only supports SWA through FusedAttention, FrostAttention" + f" or FlashAttention >= 2.3. Found {use_fused_attention=}, {use_flash_attn_3=}, " + f"{use_flash_attn_4=}, {use_frost_attention=}, " f"and {fa_utils.v2_3_plus=}." ) if load_balancing_strategy is CPLoadBalancingStrategy.DUAL_CHUNK_SWAP: @@ -5004,10 +5005,12 @@ def forward( or use_fused_attention or use_flash_attn_3 or use_flash_attn_4 + or use_frost_attention or fa_utils.v2_3_plus ), ( - "cp_comm_type='a2a' only supports SWA through FusedAttention or FlashAttention >= 2.3." - f" Found {use_fused_attention=}, {use_flash_attn_3=}, {use_flash_attn_4=}, " + "cp_comm_type='a2a' only supports SWA through FusedAttention, FrostAttention or" + f" FlashAttention >= 2.3. Found {use_fused_attention=}, {use_flash_attn_3=}, " + f"{use_flash_attn_4=}, {use_frost_attention=}, " f"and {fa_utils.v2_3_plus=}." ) assert q.shape[seq_dim_qkv] % 2 == 0 and k.shape[seq_dim_qkv] % 2 == 0, ( diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index ded9bb8c8ca..6e9097e20cb 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -37,7 +37,7 @@ Numerics were validated against the criterion FlashAttention applies to itself, namely that the kernel error must stay within 2x the error bf16 inputs alone produce, across square and -rectangular, causal and non-causal shapes. +rectangular, causal and non-causal, windowed and unwindowed shapes. """ from __future__ import annotations @@ -241,9 +241,20 @@ def _mask_spec(attn_mask_type: str, window_size=None): "FROST attention supports attn_mask_type in %s; got %r" % (str(_SUPPORTED_MASKS), attn_mask_type) ) - window = _NO_WINDOW if window_size is None else tuple(window_size) + try: + window = _NO_WINDOW if window_size is None else tuple(window_size) + except TypeError: + # Raised as NotImplementedError so the selector declines instead of propagating out of + # backend selection, which is the only thing is_frost_attention_supported catches. + raise NotImplementedError( + "window_size must be a (left, right) pair; got %r" % (window_size,) + ) from None if len(window) != 2: raise NotImplementedError("window_size must be a (left, right) pair; got %r" % (window,)) + if window[0] < -1: + # cuDNN's left bound must be >= 1, so a left of -2 would build diagonal_band_left_bound=-1 + # and fail at plan build rather than declining here. + raise NotImplementedError("window_size left must be -1 or >= 0; got %r" % (window,)) if window[1] not in (-1, 0): # A right bound past the diagonal is future context. cuDNN can express it, but no TE mask # type asks for it, so decline rather than guess the intent. @@ -628,8 +639,8 @@ def frost_attn_bwd( dout: torch.Tensor, attn_scale: Optional[float] = None, attn_mask_type: str = "causal", - window_size: Optional[Tuple[int, int]] = None, deterministic: bool = False, + window_size: Optional[Tuple[int, int]] = None, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Backward attention via cuDNN FROST. `softmax_lse` is [b, h, s] as returned by the forward.""" for name, tensor in (("q", q), ("k", k), ("v", v), ("out", out), ("dout", dout)): diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 638797d6f35..c9a4962ad47 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -1927,6 +1927,22 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt # Explicit anyway: no page table reaches the backend, so a paged cache would be read raw. logger.debug("Disabling FrostAttention for KV caching") use_frost_attention = False + if ( + use_frost_attention + and window_size is not None + and window_size[0] != -1 + and "causal" not in attn_mask_type + and max_seqlen_q != max_seqlen_kv + ): + # FROST anchors the band from the mask type, so a windowed non-causal mask always lands + # top-left. TE's bottom_right_diagonal defaults to True and the C++ fused path honours it + # (fused_attn_f16_arbitrary_seqlen.cu picks the alignment from that flag), so for unequal + # q/kv lengths the two would disagree silently. Decline rather than guess the anchor. + logger.debug( + "Disabling FrostAttention for a windowed non-causal mask with max_seqlen_q != " + "max_seqlen_kv, where the diagonal anchor is ambiguous" + ) + use_frost_attention = False if ( use_frost_attention and context_parallel From 3eea2f951753b74f6c11708bb7f633c54190a3f9 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 09:54:27 -0700 Subject: [PATCH 23/97] docs(attention): correct the claimed cuDNN import-ordering hazard The docstring asserted that FROST engines register at cudnn import time, so a setdefault running after another module had already imported cudnn would be too late and leave no FROST engine. That is what the documentation implies, and it was the basis for a concern about the in-flight port of cuDNN attention to the Python API, whose shared import helper sets no such variable. Measured on B200 with cuDNN Frontend 1.29.0 and it does not hold: importing cudnn and cudnn.sdpa first with the switch unset, then setting it and building a plan, still selects a FROST engine. The switch is still set before the import, because that is what the documentation asks for and it costs nothing, but nothing depends on winning the race and the plan-name check verifies the engine either way. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../attention/dot_product_attention/frost_attention.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 6e9097e20cb..9a8d9b16b0a 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -91,9 +91,12 @@ def _import_cudnn(): """Import cuDNN Frontend with FROST engines enabled, once. - Note the ordering hazard: the engines register at import time, so if another module imported - cudnn first without the switch set, setdefault here is too late and no FROST engine exists. - _select_frost_plan catches that by checking the plan name, but only once a plan is built. + The switch is set before the import because the documentation describes the engines as + registering at import time. Measured on B200 with cuDNN Frontend 1.29.0, the ordering turns + out not to matter: importing cudnn and cudnn.sdpa first with the switch unset, then setting + it and building a plan, still selects a FROST engine. Setting it first is kept because it is + what the documentation asks for and costs nothing, but nothing here depends on winning that + race, and _select_frost_plan verifies the engine by plan name regardless. """ global _cudnn if _cudnn is None: From e062e8d4ac0c0ea43922299546dc4aecb662fcb7 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 10:16:25 -0700 Subject: [PATCH 24/97] test(attention): apply the window for every mask type in the reference The float64 reference skipped masking entirely whenever attn_mask_type was "no_mask", so a window passed alongside it was ignored. That is not TE's rule: the SWA construction in utils.py applies the window to any mask type, treats -1 as unbounded on that side, and lets a causal mask type pin the right bound to the diagonal. no_mask with (w, 0) is therefore a causal band of width w, not an unmasked attention. Caught by the B200 run, which reported dq off by 2.9 against a reference maximum of 0.83 for exactly the two no_mask windowed backward cases, while every causal windowed case passed -- the signature of a reference that is masking differently rather than a kernel that is computing wrongly. The reference now derives blocking from (left, right) plus the diagonal offset, which reproduces TE's keep rule for every mask type and window combination that carries a window. The previous form is unchanged for unwindowed causal and bottom-right, so the cases already verified on hardware keep their meaning. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_frost_attention.py | 27 ++++++++++--------- 1 file changed, 15 insertions(+), 12 deletions(-) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index bf904c8a732..cdd72551756 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -78,18 +78,21 @@ def _reference(q, k, v, scale, mask, window=None): kk = kk.repeat_interleave(rep, dim=1) vv = vv.repeat_interleave(rep, dim=1) s = (qq @ kk.transpose(-1, -2)) * scale - if mask != "no_mask": - sq, skv = qq.shape[2], kk.shape[2] - # Top-left for "causal", bottom-right for "causal_bottom_right". These coincide only when - # sq == skv, which is exactly why _SHAPES includes a rectangular case. - offset = 0 if mask == "causal" else skv - sq - blocked = torch.ones(sq, skv, device=q.device, dtype=torch.bool).triu(offset + 1) - if window is not None and window[0] != -1: - # A left window keeps only the most recent window[0] keys before the diagonal, so - # everything further back is masked as well. - blocked |= torch.ones(sq, skv, device=q.device, dtype=torch.bool).tril( - offset - window[0] - 1 - ) + sq, skv = qq.shape[2], kk.shape[2] + left, right = (-1, -1) if window is None else tuple(window) + # TE's rule, from the SWA construction in utils.py: a causal mask type pins the right bound to + # the diagonal, -1 means unbounded on that side, and a window applies to ANY mask type -- so + # no_mask with (w, 0) is a causal band of width w, not an unmasked attention. Top-left for + # "causal", bottom-right for "causal_bottom_right"; the two coincide only when sq == skv. + if mask in ("causal", "causal_bottom_right"): + right = 0 + offset = skv - sq if mask == "causal_bottom_right" else 0 + blocked = torch.zeros(sq, skv, device=q.device, dtype=torch.bool) + if right != -1: + blocked |= torch.ones(sq, skv, device=q.device, dtype=torch.bool).triu(offset + right + 1) + if left != -1: + blocked |= torch.ones(sq, skv, device=q.device, dtype=torch.bool).tril(offset - left - 1) + if bool(blocked.any()): s = s.masked_fill(blocked, float("-inf")) p = s.softmax(-1) return p @ vv, torch.logsumexp(s, dim=-1) From dd0033c6b52957b029b4972dc400f994382f17de Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 11:49:29 -0700 Subject: [PATCH 25/97] docs(attention): justify the p2p sliding-window decline from the ring, not precedent The comment said FusedAttention carries the same rule and left the reason as an assertion about per-step tiles. The actual reason is visible in the p2p path: it hardcodes the per-step window to (-1, 0) or (-1, -1) at every kernel call, so a user window is discarded there regardless of backend. all_gather by contrast computes window_size_per_step through get_kv_seq_info_after_all_gather and passes it down, and a2a sees the whole sequence after the all-to-all. Citing that is checkable; citing another backend's rule is not. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/dot_product_attention/utils.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index c9a4962ad47..f9ecc975a20 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -1950,9 +1950,11 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt and (window_size[0] != -1 or window_size[1] not in [-1, 0]) and cp_comm_type in ["p2p", "a2a+p2p"] ): - # Same rule FusedAttention carries: the p2p ring shards KV across steps, so a left bound - # measured against the full sequence does not survive the per-step tiles. all_gather and - # a2a both see a contiguous KV range and do support it. + # Same rule FusedAttention carries, and for a reason visible in the ring itself: the p2p + # path hardcodes the per-step window to (-1, 0) or (-1, -1) at every kernel call, so a + # user window is discarded there for any backend. all_gather has real machinery for this + # (window_size_per_step, from get_kv_seq_info_after_all_gather) and a2a sees the whole + # sequence, so both can serve it. logger.debug( "Disabling FrostAttention as it does not support context parallelism with sliding" " window attention and cp_comm_type = %s", From a82903b6cda70914cadad5ad5957de7b8e56b412 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 12:13:46 -0700 Subject: [PATCH 26/97] fix(attention): bind the FROST flag on the ONNX path, decline what was silently ignored DotProductAttention.forward takes a separate branch under ONNX export that skips get_attention_backend and binds the backend flags itself. Adding a seventh flag without binding it there made the availability check below read an unassigned local: UnboundLocalError on every ONNX export, on every GPU, at any head dim, and NVTE_FROST_ATTN=0 does not help because that branch never consults it. This is the one defect in the series that reaches users who will never touch head_dim 512. Five capabilities were silently ignored rather than declined, all reachable because at head_dim 512 the fused path is unavailable and FROST becomes the sole survivor of filters written to disable everything: score_mod (including the score_mod_bprop-without-score_mod case, which is meant to end in "no backend available" and was instead being rescued into a wrong answer), a quantized qkv_type carrying a nominal bf16 dtype outside an fp8 autocast, num_splits, checkpoint_core_attention, and CUDA graph capture. Each now declines, matching what the neighbouring filters do for the other backends. Lint: the new module scored 8.91 against the repo's own pylint gate, which does not disable consider-using-f-string, while the sibling flex_attention.py scores 10.00. All thirty percent-format sites are now f-strings, verified by comparing the rendered decline messages before and after. Also clears the regressions this branch introduced elsewhere -- an unused-argument pair, a used-before-assignment that three branches made unprovable, and a condition one clause over the limit. Every touched file is back to 10.00 under the pinned pylint and CI's Python. Removes a deterministic parameter accidentally added to cp_p2p_bwd_flash_attn, which never read it, and asserts in the all_gather window helper the invariant that the selector enforces in another file. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_frost_attention.py | 24 +++- .../dot_product_attention/context_parallel.py | 29 +++-- .../dot_product_attention.py | 3 + .../dot_product_attention/frost_attention.py | 114 ++++++++---------- .../attention/dot_product_attention/utils.py | 35 +++++- 5 files changed, 124 insertions(+), 81 deletions(-) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index cdd72551756..e01a3481a26 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -129,7 +129,10 @@ def test_frost_forward_matches_reference(shape, mask, dtype): torch.manual_seed(0) # Generate in fp32 so there is a true high-precision original to measure against, then cast # for the kernel. [b, h, s, d] views over bshd-contiguous memory is what the backend consumes. - mk = lambda s_, h_: torch.randn(b, s_, h_, d, device="cuda").permute(0, 2, 1, 3).contiguous() + # A bshd VIEW, which is what the backend receives: to_frost_layout permutes a bshd-contiguous + # tensor and hands the result over without a copy. Materialising with .contiguous() here would + # produce bhsd strides instead and leave the stride-keyed plan cache untested. + mk = lambda s_, h_: torch.randn(b, s_, h_, d, device="cuda").permute(0, 2, 1, 3) q32, k32, v32 = mk(sq, hq), mk(skv, hkv), mk(skv, hkv) q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) scale = 1.0 / math.sqrt(d) @@ -159,8 +162,9 @@ def test_frost_forward_matches_reference(shape, mask, dtype): @pytest.mark.parametrize("window", [(256, 0), (128, 0), (0, 0)], ids=lambda w: "win%d" % w[0]) -@pytest.mark.parametrize("mask", ["causal", "causal_bottom_right"]) -def test_frost_sliding_window_matches_reference(mask, window): +@pytest.mark.parametrize("mask", ["causal", "causal_bottom_right", "no_mask"]) +@pytest.mark.parametrize("sq,skv", [(1024, 1024), (512, 1024)], ids=["square", "rect"]) +def test_frost_sliding_window_matches_reference(mask, window, sq, skv): """Sliding window against the float64 reference. The engine advertises swa support, and cuDNN expresses a window as a left bound on the same @@ -172,10 +176,15 @@ def test_frost_sliding_window_matches_reference(mask, window): frost_attn_fwd, ) - b, hq, hkv, sq, skv, d = 2, 8, 4, 1024, 1024, 512 + # The rectangular case is the one that matters for alignment: top-left and bottom-right + # coincide when sq == skv, so a swapped alignment is invisible in square shapes. + b, hq, hkv, d = 2, 8, 4, 512 dtype = torch.bfloat16 torch.manual_seed(0) - mk = lambda s_, h_: torch.randn(b, s_, h_, d, device="cuda").permute(0, 2, 1, 3).contiguous() + # A bshd VIEW, which is what the backend receives: to_frost_layout permutes a bshd-contiguous + # tensor and hands the result over without a copy. Materialising with .contiguous() here would + # produce bhsd strides instead and leave the stride-keyed plan cache untested. + mk = lambda s_, h_: torch.randn(b, s_, h_, d, device="cuda").permute(0, 2, 1, 3) q32, k32, v32 = mk(sq, hq), mk(skv, hkv), mk(skv, hkv) q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) scale = 1.0 / math.sqrt(d) @@ -210,7 +219,10 @@ def test_frost_backward_matches_reference(shape, mask, window): b, hq, hkv, sq, skv, d = shape dtype = torch.bfloat16 torch.manual_seed(0) - mk = lambda s_, h_: torch.randn(b, s_, h_, d, device="cuda").permute(0, 2, 1, 3).contiguous() + # A bshd VIEW, which is what the backend receives: to_frost_layout permutes a bshd-contiguous + # tensor and hands the result over without a copy. Materialising with .contiguous() here would + # produce bhsd strides instead and leave the stride-keyed plan cache untested. + mk = lambda s_, h_: torch.randn(b, s_, h_, d, device="cuda").permute(0, 2, 1, 3) q32, k32, v32 = mk(sq, hq), mk(skv, hkv), mk(skv, hkv) q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) scale = 1.0 / math.sqrt(d) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index 4a6d50cea1a..e02456f20b2 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1467,7 +1467,6 @@ def cp_p2p_bwd_flash_attn( out_part, dout_part, section, - deterministic=False, ): """Per-tile backward call of CP P2P with FlashAttention backend""" if pad_between_seqs: @@ -1595,19 +1594,25 @@ def _frost_mask_for_section(attn_mask_type, section): return attn_mask_type if section in ("lower-triangle", "upper-triangle"): return "no_mask" - raise ValueError("unknown CP section %r" % section) + raise ValueError(f"unknown CP section {section!r}") def _frost_mask_for_window(window_size): """Per-step mask for the all_gather path, derived from its adjusted window. get_kv_seq_info_after_all_gather trims KV and returns a window that is BOTTOM-RIGHT aligned: - (-1, 0) means causal relative to the trimmed KV, not top-left causal. Using top-left here - would silently compute a different mask, since the two only coincide when SQ == SKV and - all_gather never produces that. + (-1, 0) means causal relative to the trimmed KV, not top-left causal. Using top-left would be + wrong wherever the two differ, which is whenever the trim leaves SKV > SQ. """ if window_size is None or tuple(window_size) == (-1, -1): return "no_mask", None + # A positive right bound is look-ahead, which none of the supported masks express. _mask_spec + # rejects it at selection time, but that is a different file, so assert the invariant here + # rather than quietly returning a causal mask that admits future keys. + assert window_size[1] in ( + -1, + 0, + ), f"all_gather produced a look-ahead window {window_size}" # Anything with a bounded side is causal relative to the trimmed KV, and a bounded left side # is a sliding window. Both are expressed as a band against the bottom-right diagonal, so the # window travels with the mask type rather than needing a separate spelling per case. @@ -1754,12 +1759,16 @@ def cp_p2p_fwd_frost_attn( q_part, k_part, v_part, - cu_seqlens_q_per_step, # noqa: ARG001 unused for bshd; matches the fused call convention - cu_seqlens_kv_per_step, # noqa: ARG001 + cu_seqlens_q_per_step, + cu_seqlens_kv_per_step, section, -): +): # pylint: disable=unused-argument """Per-tile forward call of CP P2P with the cuDNN FROST backend. + cu_seqlens_*_per_step are accepted but unused: they carry the thd offsets, and thd is + declined by the selector. They stay in the signature so the ring can call this and + cp_p2p_fwd_fused_attn with one argument list. + Returns the same 5-tuple shape as cp_p2p_fwd_fused_attn so the ring code can consume it unchanged. rng_state, attn_bias and max_logit are None: FROST supports neither dropout nor bias, and the selector declines those configurations before we get here. @@ -5562,6 +5571,10 @@ def backward(ctx, dout, *_args): fa_backward_kwargs["softcap"] = ctx.softcap dq_fp8, dk_fp8, dv_fp8 = None, None, None + # Only the fused branch below binds this, and only the fused branch reads it further + # down -- but with three branches that binding no longer dominates the read, so give it + # a definition rather than rely on the conditions staying in step. + rest = [] if ctx.use_frost_attention: dq, dk, dv = cp_a2a_bwd_frost_attn( ctx.softmax_scale, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py index b4ecff5e836..41f94352dff 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py @@ -2858,6 +2858,9 @@ def forward( use_flash_attention = False use_fused_attention = False use_unfused_attention = True + # Bound here too: the availability check below reads all four flags at this + # scope, and this branch never calls get_attention_backend. + use_frost_attention = False else: if ( _attention_backends["attention_params"] is None diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 9a8d9b16b0a..4725cc6735e 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -120,7 +120,7 @@ def _handle_for(device: torch.device): executed from different streams across ring steps. """ if device.type != "cuda": - raise ValueError("FrostAttention requires CUDA tensors; got device %s" % device) + raise ValueError(f"FrostAttention requires CUDA tensors; got device {device}") cudnn = _import_cudnn() if device.index is None: device = torch.device("cuda", torch.cuda.current_device()) @@ -185,14 +185,12 @@ def _no(reason): # and raising from _select_frost_plan once a plan is built. return _no("CUDNN_FRONTEND_ENABLE_FROST_ENGINES=0 disables the FROST engines") if torch.cuda.get_device_capability() not in _SUPPORTED_ARCHS: - return _no( - "cuDNN FROST head_dim>256 kernels are SM100/SM103 only; found sm%d%d" - % torch.cuda.get_device_capability() - ) + major, minor = torch.cuda.get_device_capability() + return _no(f"cuDNN FROST head_dim>256 kernels are SM100/SM103 only; found sm{major}{minor}") try: _import_cudnn() except ImportError as exc: - return _no("nvidia-cudnn-frontend not importable: %s" % exc) + return _no(f"nvidia-cudnn-frontend not importable: {exc}") # Decline on positive evidence that FROST cannot work: a version below a floor, or a package # that is absent outright. A version that is present but unparseable is NOT evidence, so it @@ -200,20 +198,19 @@ def _no(reason): frontend, frontend_raw = _pkg_version("nvidia-cudnn-frontend", _cudnn) if frontend is not None and frontend < _MIN_CUDNN_FRONTEND: return _no( - "nvidia-cudnn-frontend %s registers no sm100 backward engine; >= %s is required" - " (1.28.0 ships the d512 forward only, so this would otherwise raise on the first" - " backward rather than here)" % (frontend_raw, _MIN_CUDNN_FRONTEND) + f"nvidia-cudnn-frontend {frontend_raw} registers no sm100 backward engine; >=" + f" {_MIN_CUDNN_FRONTEND} is required (1.28.0 ships the d512 forward only, so this would" + " otherwise raise on the first backward rather than here)" ) cutlass, cutlass_raw = _pkg_version("nvidia-cutlass-dsl") if cutlass_raw is None: - return _no("nvidia-cutlass-dsl not installed (FROST requires >= %s)" % _MIN_CUTLASS_DSL) + return _no(f"nvidia-cutlass-dsl not installed (FROST requires >= {_MIN_CUTLASS_DSL})") if cutlass is not None and cutlass < _MIN_CUTLASS_DSL: # Worth being loud: this combination fails by silently declining, not by raising. return _no( - "nvidia-cutlass-dsl %s is below the FROST floor %s; FROST engines would be" - " silently skipped in favour of ordinary cuDNN backend plans" - % (cutlass_raw, _MIN_CUTLASS_DSL) + f"nvidia-cutlass-dsl {cutlass_raw} is below the FROST floor {_MIN_CUTLASS_DSL}; FROST" + " engines would be silently skipped in favour of ordinary cuDNN backend plans" ) _availability = (True, "") @@ -241,8 +238,8 @@ def _mask_spec(attn_mask_type: str, window_size=None): """Validate a TE mask type and window, returning the hashable spec the plan is keyed on.""" if attn_mask_type not in _SUPPORTED_MASKS: raise NotImplementedError( - "FROST attention supports attn_mask_type in %s; got %r" - % (str(_SUPPORTED_MASKS), attn_mask_type) + f"FROST attention supports attn_mask_type in {str(_SUPPORTED_MASKS)}; got" + f" {attn_mask_type!r}" ) try: window = _NO_WINDOW if window_size is None else tuple(window_size) @@ -250,18 +247,18 @@ def _mask_spec(attn_mask_type: str, window_size=None): # Raised as NotImplementedError so the selector declines instead of propagating out of # backend selection, which is the only thing is_frost_attention_supported catches. raise NotImplementedError( - "window_size must be a (left, right) pair; got %r" % (window_size,) + f"window_size must be a (left, right) pair; got {window_size!r}" ) from None if len(window) != 2: - raise NotImplementedError("window_size must be a (left, right) pair; got %r" % (window,)) + raise NotImplementedError(f"window_size must be a (left, right) pair; got {window!r}") if window[0] < -1: # cuDNN's left bound must be >= 1, so a left of -2 would build diagonal_band_left_bound=-1 # and fail at plan build rather than declining here. - raise NotImplementedError("window_size left must be -1 or >= 0; got %r" % (window,)) + raise NotImplementedError(f"window_size left must be -1 or >= 0; got {window!r}") if window[1] not in (-1, 0): # A right bound past the diagonal is future context. cuDNN can express it, but no TE mask # type asks for it, so decline rather than guess the intent. - raise NotImplementedError("FROST attention does not support a right window %r" % (window,)) + raise NotImplementedError(f"FROST attention does not support a right window {window!r}") return attn_mask_type, window @@ -303,19 +300,19 @@ def is_frost_attention_supported( them should pay that cost or have their engine pool changed underneath them. """ if head_dim_qk != head_dim_v: - return False, "FROST path requires symmetric head_dim; got %d/%d" % ( - head_dim_qk, - head_dim_v, - ) + return False, f"FROST path requires symmetric head_dim; got {head_dim_qk}/{head_dim_v}" if not _MIN_HEAD_DIM <= head_dim_qk <= _MAX_HEAD_DIM: - return False, "FROST path covers head_dim in (256, 512]; got %d" % head_dim_qk + return False, f"FROST path covers head_dim in (256, 512]; got {head_dim_qk}" if head_dim_qk % _HEAD_DIM_MULTIPLE != 0: - return False, "FROST path needs head_dim to be a multiple of %d; got %d" % ( - _HEAD_DIM_MULTIPLE, - head_dim_qk, + return ( + False, + ( + f"FROST path needs head_dim to be a multiple of {_HEAD_DIM_MULTIPLE}; got" + f" {head_dim_qk}" + ), ) if qkv_dtype not in (torch.bfloat16, torch.float16): - return False, "FROST path supports bf16/fp16; got %s" % qkv_dtype + return False, f"FROST path supports bf16/fp16; got {qkv_dtype}" if dropout != 0.0: return False, "FROST path does not support dropout" if attn_bias_type != "no_bias": @@ -342,8 +339,8 @@ def to_frost_layout(t: torch.Tensor, qkv_format: str) -> torch.Tensor: if qkv_format == "sbhd": # [s, b, h, d] -> [b, h, s, d] return t.permute(1, 2, 0, 3) raise NotImplementedError( - "FROST attention supports qkv_format 'bshd' and 'sbhd'; got %r." - " thd needs varlen support that is not implemented here." % qkv_format + f"FROST attention supports qkv_format 'bshd' and 'sbhd'; got {qkv_format!r}. thd needs" + " varlen support that is not implemented here." ) @@ -354,7 +351,7 @@ def from_frost_layout(t: torch.Tensor, qkv_format: str) -> torch.Tensor: if qkv_format == "sbhd": # [b, h, s, d] -> [s, b, h, d] return t.permute(2, 0, 1, 3) raise NotImplementedError( - "FROST attention supports qkv_format 'bshd' and 'sbhd'; got %r." % qkv_format + f"FROST attention supports qkv_format 'bshd' and 'sbhd'; got {qkv_format!r}." ) @@ -374,11 +371,11 @@ def _check_layout(name: str, t: torch.Tensor) -> None: dimension is contiguous, which the kernels assume. """ if t.dim() != 4: - raise ValueError("%s must be 4D [b, h, s, d]; got %s" % (name, tuple(t.shape))) + raise ValueError(f"{name} must be 4D [b, h, s, d]; got {tuple(t.shape)}") if t.stride(3) != 1: raise ValueError( - "%s must have a contiguous head dimension; got shape %s stride %s" - % (name, tuple(t.shape), tuple(t.stride())) + f"{name} must have a contiguous head dimension; got shape {tuple(t.shape)} stride" + f" {tuple(t.stride())}" ) @@ -390,7 +387,7 @@ def _check_dtype(name: str, t: torch.Tensor, expected: torch.dtype) -> None: matters most: it arrives from autograd and is not this module's to control. """ if t.dtype != expected: - raise ValueError("%s must be %s to match q; got %s" % (name, expected, t.dtype)) + raise ValueError(f"{name} must be {expected} to match q; got {t.dtype}") def _check_kv_match(k: torch.Tensor, v: torch.Tensor) -> None: @@ -402,11 +399,11 @@ def _check_kv_match(k: torch.Tensor, v: torch.Tensor) -> None: and is purely a guard against a silent wrong answer. """ if k.shape != v.shape: - raise ValueError("k and v must have the same shape; got %s and %s" % (k.shape, v.shape)) + raise ValueError(f"k and v must have the same shape; got {k.shape} and {v.shape}") if k.stride() != v.stride(): raise ValueError( - "k and v must have the same layout; got strides %s and %s" - % (tuple(k.stride()), tuple(v.stride())) + f"k and v must have the same layout; got strides {tuple(k.stride())} and" + f" {tuple(v.stride())}" ) @@ -425,17 +422,12 @@ def _select_frost_plan(graph, token: str, what: str): # Both versions, because either floor can cause this and blaming one misdirects. Looked # up defensively: this is the message explaining a failure, so it must not raise itself. raise RuntimeError( - "no cuDNN FROST %s engine was offered (looked for %r). Candidate plans: %s." - " nvidia-cudnn-frontend=%s (floor %s), nvidia-cutlass-dsl=%s (floor %s)." - % ( - what, - token, - names[:6], - _pkg_version("nvidia-cudnn-frontend", _cudnn)[1] or "unknown", - _MIN_CUDNN_FRONTEND, - _pkg_version("nvidia-cutlass-dsl")[1] or "unknown", - _MIN_CUTLASS_DSL, - ) + f"no cuDNN FROST {what} engine was offered (looked for {token!r}). Candidate plans:" + f" {names[:6]}." + f" nvidia-cudnn-frontend={_pkg_version('nvidia-cudnn-frontend', _cudnn)[1] or 'unknown'} (floor" + f" {_MIN_CUDNN_FRONTEND})," + f" nvidia-cutlass-dsl={_pkg_version('nvidia-cutlass-dsl')[1] or 'unknown'} (floor" + f" {_MIN_CUTLASS_DSL})." ) graph.select_plan(hits[0]) graph.check_support() @@ -605,13 +597,10 @@ def frost_attn_fwd( if k.shape[0] != q.shape[0] or k.shape[3] != q.shape[3]: # The graph declares k and v with q's batch and head_dim, so a mismatch would bind a # differently shaped buffer to that node and read the wrong elements silently. - raise ValueError( - "k must match q in batch and head_dim; got q %s and k %s" % (q.shape, k.shape) - ) + raise ValueError(f"k must match q in batch and head_dim; got q {q.shape} and k {k.shape}") if q.shape[1] % k.shape[1] != 0: raise ValueError( - "num_heads must be divisible by num_gqa_groups; got %d and %d" - % (q.shape[1], k.shape[1]) + f"num_heads must be divisible by num_gqa_groups; got {q.shape[1]} and {k.shape[1]}" ) mask = _mask_spec(attn_mask_type, window_size) @@ -653,25 +642,20 @@ def frost_attn_bwd( # The same shape assumptions the forward makes, plus o/dO, which the graph declares with q's # shape. The forward runs first in autograd, but the CP ring calls this directly. if k.shape[0] != q.shape[0] or k.shape[3] != q.shape[3]: - raise ValueError( - "k must match q in batch and head_dim; got q %s and k %s" % (q.shape, k.shape) - ) + raise ValueError(f"k must match q in batch and head_dim; got q {q.shape} and k {k.shape}") if q.shape[1] % k.shape[1] != 0: raise ValueError( - "num_heads must be divisible by num_gqa_groups; got %d and %d" - % (q.shape[1], k.shape[1]) + f"num_heads must be divisible by num_gqa_groups; got {q.shape[1]} and {k.shape[1]}" ) for name, tensor in (("out", out), ("dout", dout)): if tensor.shape != q.shape: - raise ValueError( - "%s must have q's shape; got %s and %s" % (name, tensor.shape, q.shape) - ) + raise ValueError(f"{name} must have q's shape; got {tensor.shape} and {q.shape}") if softmax_lse.dtype != torch.float32: - raise ValueError("softmax_lse must be fp32; got %s" % softmax_lse.dtype) + raise ValueError(f"softmax_lse must be fp32; got {softmax_lse.dtype}") if tuple(softmax_lse.shape[:3]) != tuple(q.shape[:3]): raise ValueError( - "softmax_lse must be [b, h, s] matching q; got %s and %s" - % (tuple(softmax_lse.shape), tuple(q.shape)) + f"softmax_lse must be [b, h, s] matching q; got {tuple(softmax_lse.shape)} and" + f" {tuple(q.shape)}" ) mask = _mask_spec(attn_mask_type, window_size) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index f9ecc975a20..50031022ba6 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -1917,6 +1917,35 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt # cuDNN ships a deterministic d512 backward. logger.debug("Disabling FrostAttention as its backward has no deterministic cuDNN plan") use_frost_attention = False + if use_frost_attention and (has_score_mod or has_score_mod_bprop): + # The score_mod filter above disables flash, fused and unfused, and at head_dim 512 the + # fused path is unavailable anyway -- so without this FROST would be the sole survivor + # and would compute plain attention with the callback silently dropped. That includes the + # score_mod_bprop-without-score_mod case, which is meant to end in "no backend available". + logger.debug("Disabling FrostAttention for score_mod") + use_frost_attention = False + if use_frost_attention and qkv_type is not torch.Tensor: + # Every other backend filters on the tensor class, not just the dtype: a quantized tensor + # can carry a nominal bf16 dtype outside an fp8 autocast, and the fp8 guard below keys on + # the autocast flag rather than the type. + logger.debug("Disabling FrostAttention for qkv_type = %s", qkv_type) + use_frost_attention = False + if use_frost_attention and num_splits != 1: + # Declined for the same reason the fused and unfused paths are: silently ignoring it + # would change the computation the caller asked for. + logger.debug("Disabling FrostAttention for num_splits = %s", num_splits) + use_frost_attention = False + if use_frost_attention and checkpoint_core_attention: + # The backend FROST displaces at this head dim is unfused, which does honour activation + # recompute. Selecting FROST would silently remove it, which is a memory regression + # rather than a wrong answer, but not one the caller asked for. + logger.debug("Disabling FrostAttention for checkpoint_core_attention") + use_frost_attention = False + if use_frost_attention and cuda_graph: + # Plan lookup and lazy handle creation are host-side work on the first call, which is + # hazardous inside a capture. Not validated under capture, so decline rather than guess. + logger.debug("Disabling FrostAttention for CUDA graph capture") + use_frost_attention = False if use_frost_attention and return_max_logit: # FrostAttention returns the context layer alone, where UnfusedDotProductAttention returns # (context, max_logit). Selecting it here would break the caller's unpack. @@ -1943,11 +1972,13 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt "max_seqlen_kv, where the diagonal anchor is ambiguous" ) use_frost_attention = False + has_sliding_window = window_size is not None and ( + window_size[0] != -1 or window_size[1] not in [-1, 0] + ) if ( use_frost_attention and context_parallel - and window_size is not None - and (window_size[0] != -1 or window_size[1] not in [-1, 0]) + and has_sliding_window and cp_comm_type in ["p2p", "a2a+p2p"] ): # Same rule FusedAttention carries, and for a reason visible in the ring itself: the p2p From 591955de632f87589d233abdc96b64be4e2a83a5 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 13:05:17 -0700 Subject: [PATCH 27/97] fix(attention): read qkv_type from attention_params, not the rebound local The new qkv_type decline read a local that get_attention_backend rebinds far earlier: the fused-attention dtype spec assigns qkv_type, o_type, do_type and dqkv_type from spec, so by the time the FROST guards run the name holds an NVTE dtype enum rather than the tensor class. Comparing that against torch.Tensor is unequal for every input, so FROST was declined unconditionally -- the selector reported "Disabling FrostAttention for qkv_type = 6" and no backend at all for head_dim 512. Reading attention_params.qkv_type is unambiguous and cannot be shadowed. Audited the other names these guards read for the same hazard: only window_size is also rebound, at the check_set_window_size normalisation, which is the canonical value every neighbouring filter uses and is the right one to read. Caught by test_frost_sliding_window_selection_by_cp_comm_type, which exists because a reviewer pointed out the selector rules had no coverage at all. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/dot_product_attention/utils.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 50031022ba6..ea3b267c646 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -1924,11 +1924,15 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt # score_mod_bprop-without-score_mod case, which is meant to end in "no backend available". logger.debug("Disabling FrostAttention for score_mod") use_frost_attention = False - if use_frost_attention and qkv_type is not torch.Tensor: + if use_frost_attention and attention_params.qkv_type is not torch.Tensor: # Every other backend filters on the tensor class, not just the dtype: a quantized tensor # can carry a nominal bf16 dtype outside an fp8 autocast, and the fp8 guard below keys on # the autocast flag rather than the type. - logger.debug("Disabling FrostAttention for qkv_type = %s", qkv_type) + # + # Read from attention_params, not the local: the fused-attention dtype spec rebinds + # qkv_type to an NVTE dtype enum well before this point, so the local compares unequal to + # torch.Tensor for every input and would decline FROST unconditionally. + logger.debug("Disabling FrostAttention for qkv_type = %s", attention_params.qkv_type) use_frost_attention = False if use_frost_attention and num_splits != 1: # Declined for the same reason the fused and unfused paths are: silently ignoring it From 6832a9b90d8384fc24c5ecb17c7e2c6b8b44961c Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 16 Sep 2026 13:13:25 -0700 Subject: [PATCH 28/97] test(attention): cover the ONNX-export branch on hardware that can run it The ONNX fix has never executed. The section meant to verify it ran test_onnx_export.py, which imports onnxruntime -- absent from the container -- so it failed at collection and proved nothing. The bug was an UnboundLocalError, not anything about ONNX serialization: the export branch skips get_attention_backend and binds the backend flags by hand, and the availability check below it reads all of them. Entering export mode and running one ordinary head_dim-64 attention reproduces it without onnxruntime. That test must not be Blackwell-gated -- the bug hit every user on every GPU -- so the module-level pytestmark becomes a named decorator applied to the six tests that genuinely need FROST, leaving the new one to run wherever there is a CUDA device. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_frost_attention.py | 37 ++++++++++++++++++- 1 file changed, 36 insertions(+), 1 deletion(-) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index e01a3481a26..ddd46bd5277 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -51,7 +51,10 @@ def _frost_availability(): # that is supposed to cover FROST turns a silent skip into a loud failure. if os.getenv("NVTE_FROST_TEST_REQUIRED", "0") == "1" and _SKIP is not None: raise RuntimeError("NVTE_FROST_TEST_REQUIRED=1, but FrostAttention is unavailable: %s" % _SKIP) -pytestmark = pytest.mark.skipif(_SKIP is not None, reason=str(_SKIP)) +# Applied per test rather than as a module-level pytestmark: the ONNX-export regression +# below guards a code path that runs on every GPU, so gating it on Blackwell would skip it +# exactly where the bug it covers can still occur. +requires_frost = pytest.mark.skipif(_SKIP is not None, reason=str(_SKIP)) # head_dim 512 is the whole point of the backend; 320 checks the interior of the (256, 512] range # rather than only its endpoint. @@ -116,6 +119,7 @@ def _floor(q32, k32, v32, scale, mask, dtype, window=None): ) +@requires_frost @pytest.mark.parametrize("shape", _SHAPES, ids=lambda s: "b%d_hq%d_hkv%d_sq%d_skv%d_d%d" % s) @pytest.mark.parametrize("mask", ["no_mask", "causal", "causal_bottom_right"]) @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) @@ -161,6 +165,7 @@ def test_frost_forward_matches_reference(shape, mask, dtype): assert lse.dtype == torch.float32, "lse must be fp32; got %s" % lse.dtype +@requires_frost @pytest.mark.parametrize("window", [(256, 0), (128, 0), (0, 0)], ids=lambda w: "win%d" % w[0]) @pytest.mark.parametrize("mask", ["causal", "causal_bottom_right", "no_mask"]) @pytest.mark.parametrize("sq,skv", [(1024, 1024), (512, 1024)], ids=["square", "rect"]) @@ -206,6 +211,7 @@ def test_frost_sliding_window_matches_reference(mask, window, sq, skv): assert not torch.equal(out, full), "window %s produced the same output as no window" % (window,) +@requires_frost @pytest.mark.parametrize("shape", _SHAPES[:2], ids=lambda s: "b%d_hq%d_hkv%d_sq%d_skv%d_d%d" % s) @pytest.mark.parametrize("mask", ["no_mask", "causal"]) @pytest.mark.parametrize("window", [None, (128, 0)], ids=["nowin", "win128"]) @@ -252,6 +258,7 @@ def test_frost_backward_matches_reference(shape, mask, window): ) +@requires_frost def test_frost_declines_unsupported_configs(): """The selector must decline what the kernels do not serve, rather than computing wrongly.""" from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( @@ -286,6 +293,7 @@ def test_frost_declines_unsupported_configs(): assert reason, "a decline must explain itself" +@requires_frost @pytest.mark.parametrize( "cp_comm_type,window,expect_frost", [ @@ -337,6 +345,7 @@ def test_frost_sliding_window_selection_by_cp_comm_type(cp_comm_type, window, ex ) +@requires_frost def test_frost_rejects_mismatched_kv(): """k and v must agree: the graphs declare v with k's shape and stride.""" from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( @@ -360,3 +369,29 @@ def test_frost_rejects_mismatched_kv(): frost_attn_fwd(q, k, v_odd) with pytest.raises(ValueError, match="match q"): frost_attn_fwd(q, k, k.to(torch.float32)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a CUDA device") +def test_dot_product_attention_runs_in_onnx_export_mode(): + """The ONNX-export branch must bind every backend flag the availability check reads. + + Deliberately not gated on FROST: that branch skips get_attention_backend entirely and sets the + flags by hand, so leaving use_frost_attention unbound there raised UnboundLocalError for every + user on every GPU, whether or not FROST could run. A plain head_dim-64 config reproduces it -- + the failure is in the selector bookkeeping, not in any kernel. + """ + from transformer_engine.pytorch import DotProductAttention + from transformer_engine.pytorch.export import onnx_export + + b, h, s, d = 2, 4, 128, 64 + dtype = torch.bfloat16 + qkv = [torch.randn(s, b, h, d, device="cuda", dtype=dtype) for _ in range(3)] + block = DotProductAttention( + h, d, qkv_format="sbhd", attn_mask_type="causal", attention_dropout=0.0 + ).to(dtype=dtype, device="cuda") + + with onnx_export(enabled=True): + out = block(*qkv) + + assert out.numel() == s * b * h * d + assert torch.isfinite(out).all() From ca6c95ad496b68268e281cbc366784f5e5319282 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 17 Sep 2026 01:47:03 -0700 Subject: [PATCH 29/97] feat(attention): allow FrostAttention with cp_comm_type=a2a+p2p The decline said a2a+p2p "is not wired up". It is. context_parallel.py dispatches `cp_comm_type in ["p2p", "a2a+p2p"]` to the same AttnFuncWithCPAndKVP2P and passes use_frost_attention into it, where FrostAttention is called at all four forward section sites and all four backward sites. The a2a stage is flash_attn_a2a_communicate: a redistribution between sequence- and head-sharding that invokes no attention kernel. Under a2a+p2p the per-step calls are therefore the ordinary p2p section calls with fewer heads per rank. What was actually true is that it was untested. a2a+p2p needs four ranks, an a2a subgroup crossed with a p2p subgroup, and every FrostAttention CP arm ran on a pool of two. Declining an untested path is defensible; describing it as unwired was not, and it would have misled anyone deciding whether to enable it. The sliding-window decline for a2a+p2p stays and is unrelated: it rings across sub-groups, so a window measured against the full sequence still does not survive the per-step tiles, exactly as with plain p2p. Test coverage extends to four ranks for this case only, and asserts the a2a divisibility requirement rather than relying on the current configs happening to satisfy it. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- docs/envvars.rst | 2 +- .../attention/test_attention_with_cp.py | 21 +++++++++++++++---- .../attention/dot_product_attention/utils.py | 6 +++++- 3 files changed, 23 insertions(+), 6 deletions(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index 5707111f2a5..e51c8568285 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -191,7 +191,7 @@ longer backend-selection overview. :Type: ``int`` (0 or 1) :Default: ``1`` - :Description: Enable or disable FrostAttention backend (the cuDNN FROST CuTe-DSL SDPA kernels in cuDNN Frontend) for DotProductAttention. When set to ``0``, FrostAttention will not be used. It is the only backend serving symmetric ``head_dim`` in (256, 512] together with context parallelism; without context parallelism UnfusedDotProductAttention also covers that range, and FrostAttention is preferred over it where both are eligible. 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 supports context parallelism with ``cp_comm_type`` of ``p2p``, ``all_gather`` or ``a2a``, and sliding-window attention with ``all_gather`` or ``a2a`` (declined with ``p2p``, whose ring shards KV across steps). It declines FP8, ``thd`` layouts, dropout, attention bias, softcap, KV caching, ``max_logit``, and deterministic execution, the last because cuDNN offers no deterministic backward for these kernels. + :Description: Enable or disable FrostAttention backend (the cuDNN FROST CuTe-DSL SDPA kernels in cuDNN Frontend) for DotProductAttention. When set to ``0``, FrostAttention will not be used. It is the only backend serving symmetric ``head_dim`` in (256, 512] together with context parallelism; without context parallelism UnfusedDotProductAttention also covers that range, and FrostAttention is preferred over it where both are eligible. 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 supports context parallelism with ``cp_comm_type`` of ``p2p``, ``all_gather``, ``a2a`` or ``a2a+p2p``, and sliding-window attention with ``all_gather`` or ``a2a`` (declined with ``p2p`` and ``a2a+p2p``, whose ring shards KV across steps). It declines FP8, ``thd`` layouts, dropout, attention bias, softcap, KV caching, ``max_logit``, and deterministic execution, the last because cuDNN offers no deterministic backward for these kernels. .. envvar:: NVTE_UNFUSED_ATTN diff --git a/tests/pytorch/attention/test_attention_with_cp.py b/tests/pytorch/attention/test_attention_with_cp.py index d1dd818721f..08acfeccbaf 100644 --- a/tests/pytorch/attention/test_attention_with_cp.py +++ b/tests/pytorch/attention/test_attention_with_cp.py @@ -778,12 +778,17 @@ def _frost_availability(): @pytest.mark.parametrize("model", model_configs_frost_attn.keys()) @pytest.mark.parametrize("qkv_format", ["bshd", "sbhd"]) -@pytest.mark.parametrize("cp_comm_type", ["p2p", "all_gather", "a2a"]) +@pytest.mark.parametrize("cp_comm_type", ["p2p", "all_gather", "a2a", "a2a+p2p"]) def test_cp_with_frost_attention(cp_pool, model, qkv_format, cp_comm_type): """Context parallelism at head_dim 512, which no other backend serves. - thd and a2a+p2p are excluded because the backend declines them: thd needs varlen support that - is not implemented, and a2a+p2p is not wired up. + thd is excluded because the backend declines it: it needs varlen support that is not + implemented. + + a2a+p2p needs four ranks rather than two -- an a2a subgroup crossed with a p2p subgroup -- and + exercises no new attention code: it dispatches to the same AttnFuncWithCPAndKVP2P as plain p2p, + with an a2a communication stage on either side of the ring. It is covered here so that claim is + measured rather than assumed. """ reason = _frost_availability() if reason is not None: @@ -793,7 +798,15 @@ def test_cp_with_frost_attention(cp_pool, model, qkv_format, cp_comm_type): config.context_parallel = True config.cp_comm_type = cp_comm_type - pool = cp_pool(2) + # a2a requires num_heads and num_gqa_groups divisible by the a2a subgroup size; every config + # here satisfies that, but assert rather than rely on it staying true. + if cp_comm_type == "a2a+p2p": + assert config.num_heads % 2 == 0 and config.num_gqa_groups % 2 == 0, ( + f"cp_comm_type=a2a+p2p needs num_heads ({config.num_heads}) and num_gqa_groups" + f" ({config.num_gqa_groups}) divisible by the a2a subgroup size" + ) + + pool = cp_pool(4 if cp_comm_type == "a2a+p2p" else 2) _submit( pool, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index ea3b267c646..999150741a7 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -2021,9 +2021,13 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt "p2p", "all_gather", "a2a", + "a2a+p2p", ) ): - # p2p (ring), all_gather and a2a are wired up in context_parallel.py; a2a+p2p is not. + # a2a+p2p needs no separate wiring: it dispatches to AttnFuncWithCPAndKVP2P, the same class + # as plain p2p, and its a2a stage is flash_attn_a2a_communicate -- a redistribution between + # sequence- and head-sharding that calls no attention kernel. The per-step calls are the + # ordinary p2p section calls with fewer heads per rank. # Non-p2p types matter for Gemma-4: TE refuses sliding-window attention with p2p, and the # model has sliding layers, so those layers need all_gather or a2a. logger.debug( From b4cdcb4c7ff2c836de2dc62c75c072201d246f8e Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 17 Sep 2026 02:17:08 -0700 Subject: [PATCH 30/97] fix(attention): handle a list-valued cp_group in FrostAttention.forward cp_comm_type="a2a+p2p" passes cp_group as [a2a_group, p2p_group]. FrostAttention computed context_parallel with a one-liner that assumed a single group, so get_distributed_world_size received a list and raised TypeError: unhashable type: 'list' at backends.py in FrostAttention.forward, before any attention ran. The signature already declared Optional[Union[dist_group_type, List[dist_group_type]]]; the body did not honour it. Now the same form FlashAttention and FusedAttention use a few hundred lines above and below: multiply the sub-group sizes when a list arrives. Found by enabling a2a+p2p and running it, after the previous commit claimed on code-reading grounds that the path was already complete. It was reachable, but it crashed on the first line of the forward. All six a2a+p2p arms failed deterministically while the eighteen existing arms passed. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../attention/dot_product_attention/backends.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index e67d012b2ba..c15e3ea04a1 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -2431,7 +2431,16 @@ def forward( """Forward pass. Routes through the CP ring when a cp_group is present.""" assert self.attention_dropout == 0.0, "FrostAttention does not support dropout" - context_parallel = cp_group is not None and get_distributed_world_size(cp_group) != 1 + # Same form as FlashAttention and FusedAttention above. cp_group is a list of two groups + # for cp_comm_type="a2a+p2p", and passing that list to get_distributed_world_size raises + # TypeError: unhashable type: 'list'. + cp_size = 1 + if isinstance(cp_group, dist_group_type): + cp_size = get_distributed_world_size(cp_group) + elif isinstance(cp_group, list): + for group in cp_group: + cp_size *= get_distributed_world_size(group) + context_parallel = cp_size > 1 if context_parallel: output = attn_forward_func_with_cp( self.training, From 9a8b47454409df58be166a0704d30d7fed67403f Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 17 Sep 2026 09:19:56 -0700 Subject: [PATCH 31/97] test(attention): cover fp16 in the backward and under context parallelism The backend serves BF16 and FP16, but coverage was uneven: only the forward numerics ran both dtypes. The backward and all 24 context-parallel configurations were bf16 only. That is the wrong way round. fp16 has a far narrower exponent range than bf16, and the two places it would show first are exactly the two that were untested: the gradient of a softmax subtracts similarly sized terms, and the ring correction exponentiates a difference of log-sum-exp values across steps. The backward test now runs both dtypes. Context parallelism gains one fp16 arm per comm type rather than a doubled matrix -- one model, one layout, three cases. a2a+p2p is omitted from the fp16 arm because it would need a second four-rank pool for a dtype that exercises no additional code path. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../attention/test_attention_with_cp.py | 31 +++++++++++++++++++ .../pytorch/attention/test_frost_attention.py | 11 +++++-- 2 files changed, 39 insertions(+), 3 deletions(-) diff --git a/tests/pytorch/attention/test_attention_with_cp.py b/tests/pytorch/attention/test_attention_with_cp.py index 08acfeccbaf..ce0b79d51e1 100644 --- a/tests/pytorch/attention/test_attention_with_cp.py +++ b/tests/pytorch/attention/test_attention_with_cp.py @@ -820,6 +820,37 @@ def test_cp_with_frost_attention(cp_pool, model, qkv_format, cp_comm_type): ) +@pytest.mark.parametrize("cp_comm_type", ["p2p", "all_gather", "a2a"]) +def test_cp_with_frost_attention_fp16(cp_pool, cp_comm_type): + """One fp16 arm per comm type, since the matrix above is bf16 throughout. + + The backend serves BF16 and FP16, but every context-parallel configuration was covered in bf16 + only. fp16 has a far narrower exponent range, and the ring correction exponentiates a difference + of log-sum-exp values across steps, so a range problem would surface here rather than in the + non-CP numerics. One model and one layout keeps the cost to three cases rather than doubling + the matrix; a2a+p2p is omitted because it would need a second four-rank pool for a dtype that + exercises no additional code path. + """ + reason = _frost_availability() + if reason is not None: + pytest.skip(reason) + + config = model_configs_frost_attn["cp_hd512_0"] + config.context_parallel = True + config.cp_comm_type = cp_comm_type + + _submit( + cp_pool(2), + dtype="fp16", + model="cp_hd512_0", + qkv_format="bshd", + kernel_backend="FrostAttention", + cp_comm_type=cp_comm_type, + is_training=True, + log_level=pytest_logging_level, + ) + + @pytest.mark.skipif(get_cudnn_version() < (8, 9, 7), reason="cuDNN 8.9.7+ is required.") @pytest.mark.skipif( get_device_compute_capability() < (9, 0), reason="FusedAttention THD requires sm90+." diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index ddd46bd5277..0ecee3b0f63 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -215,15 +215,20 @@ def test_frost_sliding_window_matches_reference(mask, window, sq, skv): @pytest.mark.parametrize("shape", _SHAPES[:2], ids=lambda s: "b%d_hq%d_hkv%d_sq%d_skv%d_d%d" % s) @pytest.mark.parametrize("mask", ["no_mask", "causal"]) @pytest.mark.parametrize("window", [None, (128, 0)], ids=["nowin", "win128"]) -def test_frost_backward_matches_reference(shape, mask, window): - """dq/dk/dv against autograd on the same independent float64 reference.""" +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_frost_backward_matches_reference(shape, mask, window, dtype): + """dq/dk/dv against autograd on the same independent float64 reference. + + Both dtypes, not just bf16: fp16 has a much narrower exponent range, and the backward is where + that would show first -- the gradient of a softmax involves a subtraction of similarly sized + terms, so a range problem surfaces there before it surfaces in the forward. + """ from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( frost_attn_bwd, frost_attn_fwd, ) b, hq, hkv, sq, skv, d = shape - dtype = torch.bfloat16 torch.manual_seed(0) # A bshd VIEW, which is what the backend receives: to_frost_layout permutes a bshd-contiguous # tensor and hands the result over without a copy. Materialising with .contiguous() here would From 97905757edcfdb39fa507e8c2f9fb5d1df0784cd Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 17 Sep 2026 13:20:00 -0700 Subject: [PATCH 32/97] docs(attention): narrow the FrostAttention availability claim It read as the only backend serving symmetric head_dim in (256, 512] with context parallelism. That was true when written and is now imprecise: Dao-AILab/flash-attention#2877 adds symmetric D512 kernels to FA4, and with the window-sentinel fix in #3532 that path works too. Qualified to released components, which is the claim that actually holds -- #2877 is unmerged and unreviewed. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- docs/envvars.rst | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index e51c8568285..e3bc619210b 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -191,7 +191,7 @@ longer backend-selection overview. :Type: ``int`` (0 or 1) :Default: ``1`` - :Description: Enable or disable FrostAttention backend (the cuDNN FROST CuTe-DSL SDPA kernels in cuDNN Frontend) for DotProductAttention. When set to ``0``, FrostAttention will not be used. It is the only backend serving symmetric ``head_dim`` in (256, 512] together with context parallelism; without context parallelism UnfusedDotProductAttention also covers that range, and FrostAttention is preferred over it where both are eligible. 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 supports context parallelism with ``cp_comm_type`` of ``p2p``, ``all_gather``, ``a2a`` or ``a2a+p2p``, and sliding-window attention with ``all_gather`` or ``a2a`` (declined with ``p2p`` and ``a2a+p2p``, whose ring shards KV across steps). It declines FP8, ``thd`` layouts, dropout, attention bias, softcap, KV caching, ``max_logit``, and deterministic execution, the last because cuDNN offers no deterministic backward for these kernels. + :Description: Enable or disable FrostAttention backend (the cuDNN FROST CuTe-DSL SDPA kernels in cuDNN Frontend) for DotProductAttention. When set to ``0``, FrostAttention will not be used. From released components it is the only backend serving symmetric ``head_dim`` in (256, 512] together with context parallelism; without context parallelism UnfusedDotProductAttention also covers that range, and FrostAttention is preferred over it where both are eligible. 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 supports context parallelism with ``cp_comm_type`` of ``p2p``, ``all_gather``, ``a2a`` or ``a2a+p2p``, and sliding-window attention with ``all_gather`` or ``a2a`` (declined with ``p2p`` and ``a2a+p2p``, whose ring shards KV across steps). It declines FP8, ``thd`` layouts, dropout, attention bias, softcap, KV caching, ``max_logit``, and deterministic execution, the last because cuDNN offers no deterministic backward for these kernels. .. envvar:: NVTE_UNFUSED_ATTN From 9a548fb8add19991325e413c548013d67ae276df Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Sun, 20 Sep 2026 22:33:56 -0700 Subject: [PATCH 33/97] docs(attention): mark FrostAttention experimental and trim review comments Addresses review feedback on the module docstring and on comments that read as PR-specific once merged. - FrostAttention and frost_attention are marked experimental and subject to change, including possible consolidation into FusedAttention: the underlying cuDNN FROST engines are themselves experimental. - The module docstring drops the backend-by-backend motivation, which duplicates the PR description and dates quickly, and keeps the three kernel properties that constrain the code. - The selector's FROST rationale block in utils.py is reduced to two lines. - Records why select_plan precedes check_support: check_support is scoped to the selected plan, so calling it first would answer for whichever plan the heuristic ranked at index 0, and pinning is what makes build_plans strict rather than letting it walk on to a non-FROST plan. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/backends.py | 11 ++-- .../dot_product_attention/frost_attention.py | 59 ++++++++----------- .../attention/dot_product_attention/utils.py | 11 +--- 3 files changed, 35 insertions(+), 46 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index c15e3ea04a1..ed45c9aa75c 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -2385,10 +2385,13 @@ def backward(ctx, dout): class FrostAttention(torch.nn.Module): """cuDNN FROST attention for symmetric head_dim in (256, 512] on SM100/SM103. - This is the only backend that serves that head-dim range together with context parallelism, - which is what Gemma-4 global layers need. Deliberately narrow: no FP8, no bias, no dropout, - no softmax offset, no paging. get_attention_backend declines all of those before selecting - this backend, so anything reaching here should already be supported. + **Experimental and subject to change**, including the possibility of being folded into + FusedAttention: the underlying cuDNN FROST engines are themselves experimental. + + This is the only backend that serves that head-dim range together with context parallelism. + Deliberately narrow: no FP8, no bias, no dropout, no softmax offset, no paging. + get_attention_backend declines all of those before selecting this backend, so anything + reaching here should already be supported. """ def __init__( diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 4725cc6735e..51b1436ca57 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -4,40 +4,27 @@ """cuDNN FROST attention backend for head_dim in (256, 512] on SM100/SM103. -Why this exists. Gemma-4 global layers use symmetric head_dim=512, and no backend TE can select -today serves both that head dim and context parallelism: FlashAttention 2/3 cap at 256, FA4 is -gated off at symmetric 512, the C++ cuDNN fused path is refused a graph by cuDNN above 256, and -the unfused path supports 512 but cannot do CP. cuDNN Frontend 1.29.0 ships CuTe-DSL ("FROST") -SDPA kernels that do serve symmetric 512 forward and backward on Blackwell. - -Why a separate Python backend rather than teaching the existing C++ fused path. The 256 ceiling -there is not a TE check -- the f16 dispatch applies no head-dim test and simply asks cuDNN to -build a graph -- so the natural question is why the new engines cannot just be picked up. They -cannot: FROST engines are registered at Python import time behind -CUDNN_FRONTEND_ENABLE_FROST_ENGINES and require the nvidia-cutlass-dsl Python package, while -TE's C++ builds against cuDNN Frontend headers only. Reaching them therefore requires a Python -graph, which is what this module is. - -Three properties of these kernels were verified on Blackwell before this was written, and each -one constrains the code: - -1. cuDNN's `use_causal_mask` is TOP-LEFT aligned and `use_causal_mask_bottom_right` is - bottom-right. They coincide when SQ == SKV, so the distinction is invisible in square tests - and decisive for all_gather, which trims KV. Both alignments were checked against a - reference rather than assumed, and masking is built as a diagonal band so causal, - bottom-right and sliding window come from one mechanism instead of three spellings. - -2. Plan building must be cached. Building a plan is by far the most expensive cuDNN frontend - call here, and dominates an execute even after cuDNN has cached the JIT and made rebuilds - cheap, so a per-call build would leave training build-bound. Hence `_PLAN_CACHE`. - -3. The forward LSE is natural-log logsumexp in fp32, shaped [b, h, s, 1]. Squeezed to [b, h, s] - it is exactly what the CP ring correction in context_parallel.py consumes, which is what - makes ring attention over these kernels valid at all. - -Numerics were validated against the criterion FlashAttention applies to itself, namely that the -kernel error must stay within 2x the error bf16 inputs alone produce, across square and -rectangular, causal and non-causal, windowed and unwindowed shapes. +**Experimental and subject to change.** The engines this wraps are themselves experimental in +cuDNN Frontend, and if the fused path gains these shapes this backend may be folded into it. + +Why a separate Python backend rather than teaching the existing C++ fused path: FROST engines are +registered at Python import time behind CUDNN_FRONTEND_ENABLE_FROST_ENGINES and require the +nvidia-cutlass-dsl Python package, while TE's C++ builds against cuDNN Frontend headers only. +Reaching them requires a Python graph, which is what this module is. + +Three properties of these kernels were verified on Blackwell, and each constrains the code: + +1. cuDNN's causal masking is TOP_LEFT aligned unless bottom-right is requested. The two coincide + when SQ == SKV, so the distinction is invisible in square tests and decisive for all_gather, + which trims KV. Masking is built as a diagonal band so causal, bottom-right and sliding window + come from one mechanism. + +2. Plan building must be cached. It dominates an execute even after cuDNN has cached the JIT, so a + per-call build would leave training build-bound. Hence `_PLAN_CACHE`. + +3. The forward LSE is natural-log logsumexp in fp32, shaped [b, h, s, 1]. Squeezed to [b, h, s] it + is what the CP ring correction in context_parallel.py consumes, which is what makes ring + attention over these kernels valid at all. """ from __future__ import annotations @@ -429,6 +416,10 @@ def _select_frost_plan(graph, token: str, what: str): f" nvidia-cutlass-dsl={_pkg_version('nvidia-cutlass-dsl')[1] or 'unknown'} (floor" f" {_MIN_CUTLASS_DSL})." ) + # select_plan before check_support, not after: check_support is scoped to the *selected* + # plan, so calling it first would answer for whichever plan the heuristic ranked at index 0. + # Pinning also makes build_plans strict -- a decline raises instead of walking on to a + # non-FROST plan, which is the fallback this selection exists to prevent. graph.select_plan(hits[0]) graph.check_support() graph.build_plans() diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 999150741a7..6edf91b01fd 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -1866,15 +1866,10 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt ), ) FlashAttentionUtils.warning_printed = True - # cuDNN FROST (CuTe-DSL SDPA in cuDNN Frontend >= 1.29.0) is the only backend that serves - # symmetric head_dim in (256, 512] with context parallelism on SM100/SM103. Every other option - # stops short: FA2/FA3 cap at 256, FA4 is disabled at symmetric 512 above, the C++ cuDNN fused - # path is refused a graph by cuDNN above 256, and UnfusedDotProductAttention supports 512 but - # not context parallelism. Without this, Gemma-4 global layers with CP > 1 select no backend at - # all. + # FROST serves symmetric head_dim in (256, 512]; it is the only backend that also does + # context parallelism there. Experimental, and declined per-shape below. if use_frost_attention: - # Local import: frost_attention pulls in cudnn lazily, so this stays cheap and keeps - # TE importable on systems without cudnn-frontend installed. + # Local import: keeps TE importable without cudnn-frontend installed. from .frost_attention import ( # pylint: disable=import-outside-toplevel is_frost_attention_supported, ) From fe34dccd19df10cd73134d7db2a25f9bbace1b7d Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Sun, 20 Sep 2026 22:53:41 -0700 Subject: [PATCH 34/97] docs(attention): mark the FrostAttention backend experimental in envvars The docstrings say so; the place users actually read about the backend did not. Follows the wording used for the module and class, including that it may be folded into FusedAttention. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- docs/envvars.rst | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index e3bc619210b..d6723c0c219 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -191,7 +191,7 @@ longer backend-selection overview. :Type: ``int`` (0 or 1) :Default: ``1`` - :Description: Enable or disable FrostAttention backend (the cuDNN FROST CuTe-DSL SDPA kernels in cuDNN Frontend) for DotProductAttention. When set to ``0``, FrostAttention will not be used. From released components it is the only backend serving symmetric ``head_dim`` in (256, 512] together with context parallelism; without context parallelism UnfusedDotProductAttention also covers that range, and FrostAttention is preferred over it where both are eligible. 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 supports context parallelism with ``cp_comm_type`` of ``p2p``, ``all_gather``, ``a2a`` or ``a2a+p2p``, and sliding-window attention with ``all_gather`` or ``a2a`` (declined with ``p2p`` and ``a2a+p2p``, whose ring shards KV across steps). It declines FP8, ``thd`` layouts, dropout, attention bias, softcap, KV caching, ``max_logit``, and deterministic execution, the last because cuDNN offers no deterministic backward for these kernels. + :Description: Enable or disable FrostAttention backend (the cuDNN FROST CuTe-DSL SDPA kernels in cuDNN Frontend) for DotProductAttention. **This backend is experimental and subject to change**, including the possibility of being folded into FusedAttention; the underlying cuDNN FROST engines are themselves experimental. When set to ``0``, FrostAttention will not be used. From released components it is the only backend serving symmetric ``head_dim`` in (256, 512] together with context parallelism; without context parallelism UnfusedDotProductAttention also covers that range, and FrostAttention is preferred over it where both are eligible. 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 supports context parallelism with ``cp_comm_type`` of ``p2p``, ``all_gather``, ``a2a`` or ``a2a+p2p``, and sliding-window attention with ``all_gather`` or ``a2a`` (declined with ``p2p`` and ``a2a+p2p``, whose ring shards KV across steps). It declines FP8, ``thd`` layouts, dropout, attention bias, softcap, KV caching, ``max_logit``, and deterministic execution, the last because cuDNN offers no deterministic backward for these kernels. .. envvar:: NVTE_UNFUSED_ATTN From 7bbccff92ca33516fa007bebe0248cbd0aa587ca Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 30 Sep 2026 19:08:46 -0700 Subject: [PATCH 35/97] refactor(attention): extract the shared cuDNN pygraph plumbing flex_attention.py and frost_attention.py both build cuDNN graphs from Python and had grown near-identical copies of everything around the SDPA node. Review on #3527 asked for this specifically. New cudnn_pygraph.py holds the common part, with no attention semantics in it: the frontend import, one handle per device rebound to PyTorch's current stream on every call, the BHSD dim/stride description of an SBHD/BSHD tensor, plan finalization, and execution. Both backends now delegate, keeping their existing private function names so nothing referencing them has to change. Two details the shared code has to preserve rather than unify: - CUDNN_FRONTEND_ENABLE_FROST_ENGINES is not additive. It also ranks FROST ahead of the backend engines everywhere, so it is opt-in per caller and flex must not set it. - Plan choice differs on purpose. flex takes heur_mode A plus FALLBACK with HEURISTICS_CHOICE; FROST pins a plan by name and raises if no FROST engine is offered, because without the pin build_plans walks on to a fallback, which at these head dims is the wrong kernel rather than a slower one. finalize_plans serves both through require_plan_token. The failure hint stays lazy: resolving package versions is only worth doing when explaining a failure, not on every plan build. Net 149 lines of duplication removed from the two backends for a 204 line shared module. No behaviour change intended; both paths still need a GPU run to confirm. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/cudnn_pygraph.py | 208 ++++++++++++++++++ .../dot_product_attention/flex_attention.py | 95 ++------ .../dot_product_attention/frost_attention.py | 105 +++------ 3 files changed, 259 insertions(+), 149 deletions(-) create mode 100644 transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py new file mode 100644 index 00000000000..08e29f10389 --- /dev/null +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -0,0 +1,208 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Shared cuDNN Frontend Python-graph plumbing. + +Two attention backends build cuDNN graphs from Python: flex_attention.py, for score_mod, and +frost_attention.py, for the CuTe-DSL SDPA kernels at head_dim in (256, 512]. They differ in the +SDPA node they build, which cannot be shared because cuDNN treats a score_mod and a diagonal band +as mutually exclusive, but everything around that node is the same work: importing the frontend, +holding one handle per device on PyTorch's current stream, describing an SBHD/BSHD tensor in the +BHSD form cuDNN wants, finalizing plans, and executing. + +This module is that common part. It contains no attention semantics. +""" + +from typing import Any, Dict, Optional, Sequence, Tuple + +import os + +import torch + + +_cudnn = None +_handles: Dict[torch.device, Any] = {} + + +def import_cudnn_frontend(enable_frost_engines: bool = False): + """Import cuDNN Frontend once, optionally with the FROST engines registered. + + ``enable_frost_engines`` is not merely additive: the switch also ranks FROST ahead of the + backend engines everywhere, so a caller that does not want FROST must not ask for it. + + The switch is set before the import because the engines are documented as registering at + import time. Measured on B200 with cuDNN Frontend 1.29.0 the ordering turns out not to + matter, but setting it first is what the documentation asks for and costs nothing. Callers + that require a FROST engine should verify by plan name rather than rely on the switch, which + is what ``finalize_plans(require_plan_token=...)`` does. + """ + global _cudnn # pylint: disable=global-statement + if _cudnn is None: + if enable_frost_engines: + os.environ.setdefault("CUDNN_FRONTEND_ENABLE_FROST_ENGINES", "1") + try: + import cudnn # pylint: disable=import-outside-toplevel + + if enable_frost_engines: + # pylint: disable=import-outside-toplevel,unused-import + import cudnn.sdpa # noqa: F401 + except ImportError as exc: + raise ImportError( + "cuDNN frontend Python package not found. " + "Install it with: pip install nvidia-cudnn-frontend" + ) from exc + + _cudnn = cudnn + return _cudnn + + +def handle_for(device: torch.device, *, backend_name: str = "cuDNN attention"): + """A cuDNN handle for ``device``, rebound to PyTorch's current stream on every call. + + Without the rebinding, cuDNN runs on its handle's own stream while the tensors and workspace + are allocated on PyTorch's current stream, and nothing orders the two. That is not + hypothetical: the p2p context-parallel ring issues attention inside + ``with torch.cuda.stream(cp_stream)``, so on alternating ring steps the kernel and its buffers + would otherwise be on different streams. The same cached plan is executed from different + streams across steps, so this has to happen per call rather than once per handle. + """ + if device.type != "cuda": + raise ValueError(f"{backend_name} requires CUDA tensors; got device {device}") + cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() + if device.index is None: + device = torch.device("cuda", torch.cuda.current_device()) + with torch.cuda.device(device): + handle = _handles.get(device) + if handle is None: + handle = cudnn.create_handle() + _handles[device] = handle + cudnn.set_stream(handle=handle, stream=torch.cuda.current_stream(device).cuda_stream) + return handle + + +def io_data_type(cudnn, dtype: torch.dtype, *, backend_name: str = "cuDNN attention"): + """Map a torch dtype to the cuDNN frontend enum, for the dtypes these backends accept.""" + if dtype == torch.float16: + return cudnn.data_type.HALF + if dtype == torch.bfloat16: + return cudnn.data_type.BFLOAT16 + raise ValueError(f"{backend_name} only supports FP16/BF16 tensors, got {dtype}") + + +def build_pygraph(dtype: torch.dtype, device: torch.device, *, + backend_name: str = "cuDNN attention"): + """A cuDNN frontend graph for F16/BF16 SDPA, bound to this device's stream-current handle.""" + cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() + return cudnn.pygraph( + io_data_type=io_data_type(cudnn, dtype, backend_name=backend_name), + intermediate_data_type=cudnn.data_type.FLOAT, + compute_data_type=cudnn.data_type.FLOAT, + handle=handle_for(device, backend_name=backend_name), + ) + + +def bhsd_dim_stride( + tensor: torch.Tensor, tensor_format: str +) -> Tuple[Tuple[int, ...], Tuple[int, ...]]: + """Describe an SBHD/BSHD tensor as cuDNN frontend's logical BHSD form. + + No copy and no permute: the strides are handed to cuDNN as they are, which is what lets both + layouts be served directly. sbhd matters because that is what Megatron uses internally. + """ + if tensor_format == "sbhd": + return ( + (tensor.shape[1], tensor.shape[2], tensor.shape[0], tensor.shape[3]), + (tensor.stride(1), tensor.stride(2), tensor.stride(0), tensor.stride(3)), + ) + if tensor_format == "bshd": + return ( + (tensor.shape[0], tensor.shape[2], tensor.shape[1], tensor.shape[3]), + (tensor.stride(0), tensor.stride(2), tensor.stride(1), tensor.stride(3)), + ) + raise ValueError(f"Only SBHD/BSHD tensor formats are supported, got {tensor_format}.") + + +def bhsd_graph_tensor(graph, tensor: torch.Tensor, tensor_format: str): + """Create a cuDNN graph tensor with BHSD dims and the tensor's own strides.""" + dim, stride = bhsd_dim_stride(tensor, tensor_format) + return graph.tensor(dim=dim, stride=stride, data_type=tensor.dtype) + + +def finalize_plans( + graph, + *, + heuristics: Optional[Sequence[Any]] = None, + build_policy: Any = None, + require_plan_token: Optional[str] = None, + not_found_hint: Any = "", +) -> Tuple[int, Optional[str]]: + """Create plans, optionally pin one by name, build, and return (workspace size, plan name). + + ``require_plan_token`` makes the choice strict: only a plan whose name contains the token is + acceptable, and anything else raises. That is not a stylistic preference. Without a pin, + ``build_plans`` walks the ranked list and finalizes the first plan that builds, so a graph + that a specialised engine declines would quietly run on a fallback instead, which for the + FROST head-dim range is the wrong kernel rather than a slower one. + + The pin must precede ``check_support``: that call is scoped to the *selected* plan, so running + it first would answer for whichever plan the heuristic happened to rank at index 0. + """ + cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() + + graph.validate() + graph.build_operation_graph() + + if heuristics is None: + heuristics = [cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK] + + if require_plan_token is None: + try: + graph.create_execution_plans(list(heuristics)) + graph.check_support() + except cudnn.cudnnGraphNotSupportedError as exc: + raise RuntimeError(f"cuDNN SDPA graph is not supported: {exc}") from exc + if build_policy is None: + build_policy = cudnn.build_plan_policy.HEURISTICS_CHOICE + graph.build_plans(build_policy) + return max(graph.get_workspace_size(), 1), None + + graph.create_execution_plans(list(heuristics)) + names = [graph.get_plan_name_at_index(i) for i in range(graph.get_execution_plan_count())] + hits = [i for i, n in enumerate(names) if require_plan_token in n] + if not hits: + # Callable hints are resolved only here: a caller may want to look up package versions to + # explain the failure, and that work should not happen on the success path. + hint = not_found_hint() if callable(not_found_hint) else not_found_hint + raise RuntimeError( + f"no cuDNN engine matching {require_plan_token!r} was offered." + f" Candidate plans: {names[:6]}.{(' ' + hint) if hint else ''}" + ) + graph.select_plan(hits[0]) + graph.check_support() + graph.build_plans() + return max(graph.get_workspace_size(), 1), names[hits[0]] + + +def selected_plan_name(graph, index: int = 0) -> str: + """Name of the plan at ``index``, for logging and for asserting which engine answered.""" + return graph.get_plan_name_at_index(index) + + +def execute_graph( + graph, + variant_pack: Dict[Any, torch.Tensor], + workspace_size: int, + device: torch.device, + *, + backend_name: str = "cuDNN attention", +): + """Execute a built graph on this device's stream-current handle.""" + if device.type == "cuda" and device.index is None: + device = torch.device("cuda", torch.cuda.current_device()) + workspace = torch.empty(workspace_size, device=device, dtype=torch.uint8) + graph.execute( + variant_pack, + workspace, + handle=handle_for(device, backend_name=backend_name), + ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index b9593b42d9b..e5decacd84f 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -5,52 +5,40 @@ """cuDNN-backed Flex Attention helpers.""" from dataclasses import dataclass -import importlib import inspect from typing import Any, Callable, Dict, Optional, Tuple import torch -_cudnn_score_mod_handles: Dict[torch.device, Any] = {} +from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph + +# The handle cache lives in cudnn_pygraph now; the alias keeps the old name working. +_cudnn_score_mod_handles = cudnn_pygraph._handles # pylint: disable=protected-access _cudnn_score_mod_graph_cache: Dict[Tuple[Any, ...], Any] = {} _SCORE_MOD_UNCACHEABLE = object() +_BACKEND = "Flex Attention" + def _import_cudnn_frontend(): """Import the cuDNN frontend Python package.""" - try: - return importlib.import_module("cudnn") - except ImportError as exc: - raise ImportError( - "cuDNN frontend Python package not found. " - "Install it with: pip install nvidia-cudnn-frontend" - ) from exc + # Without the FROST engines: enabling them also ranks them ahead of the backend engines + # everywhere, which would change which plan this path runs. + return cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=False) def _bhsd_dim_stride( tensor: torch.Tensor, tensor_format: str ) -> Tuple[Tuple[int, ...], Tuple[int, ...]]: """Describe an SBHD/BSHD tensor as cuDNN frontend's logical BHSD format.""" - if tensor_format == "sbhd": - return ( - (tensor.shape[1], tensor.shape[2], tensor.shape[0], tensor.shape[3]), - (tensor.stride(1), tensor.stride(2), tensor.stride(0), tensor.stride(3)), - ) - if tensor_format == "bshd": - return ( - (tensor.shape[0], tensor.shape[2], tensor.shape[1], tensor.shape[3]), - (tensor.stride(0), tensor.stride(2), tensor.stride(1), tensor.stride(3)), - ) - raise ValueError(f"Flex Attention only supports SBHD/BSHD tensor formats, got {tensor_format}.") + return cudnn_pygraph.bhsd_dim_stride(tensor, tensor_format) def _bhsd_graph_tensor(graph, tensor: torch.Tensor, tensor_format: str): """Create a cuDNN graph tensor with BHSD dims and TE-layout strides.""" - dim, stride = _bhsd_dim_stride(tensor, tensor_format) - return graph.tensor(dim=dim, stride=stride, data_type=tensor.dtype) + return cudnn_pygraph.bhsd_graph_tensor(graph, tensor, tensor_format) -# score_mod graph cache helpers. def _freeze_score_mod_cache_key(value: Any) -> Any: """Convert a user-provided score_mod graph key into a hashable structure.""" if isinstance(value, torch.Tensor): @@ -194,40 +182,13 @@ def _wrapped_score_mod(sdpa_graph, score_tensor): def _get_cudnn_current_stream_handle(cudnn, device: torch.device): """Return a cuDNN handle for device, bound to PyTorch's current stream.""" - if device.type != "cuda": - raise ValueError(f"Flex Attention only supports CUDA tensors, got device {device}.") - if device.index is None: - device = torch.device("cuda", torch.cuda.current_device()) - - handle = _cudnn_score_mod_handles.get(device) - with torch.cuda.device(device): - if handle is None: - handle = cudnn.create_handle() - _cudnn_score_mod_handles[device] = handle - - stream = torch.cuda.current_stream(device).cuda_stream - cudnn.set_stream(handle=handle, stream=stream) - return handle + del cudnn # the shared helper resolves the module itself + return cudnn_pygraph.handle_for(device, backend_name=_BACKEND) def _build_cudnn_pygraph(dtype: torch.dtype, device: torch.device): """Create a cuDNN frontend Python graph for F16/BF16 SDPA.""" - cudnn = _import_cudnn_frontend() - - if dtype == torch.float16: - io_data_type = cudnn.data_type.HALF - elif dtype == torch.bfloat16: - io_data_type = cudnn.data_type.BFLOAT16 - else: - raise ValueError(f"Flex Attention only supports FP16/BF16 tensors, got {dtype}.") - - graph = cudnn.pygraph( - io_data_type=io_data_type, - intermediate_data_type=cudnn.data_type.FLOAT, - compute_data_type=cudnn.data_type.FLOAT, - handle=_get_cudnn_current_stream_handle(cudnn, device), - ) - return graph + return cudnn_pygraph.build_pygraph(dtype, device, backend_name=_BACKEND) @dataclass @@ -265,17 +226,8 @@ class _CudnnScoreModBwdGraphEntry: def _finalize_cudnn_graph(graph) -> int: """Build a cuDNN frontend Python graph and return its workspace size.""" - cudnn = _import_cudnn_frontend() - - graph.validate() - graph.build_operation_graph() - try: - graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) - graph.check_support() - except cudnn.cudnnGraphNotSupportedError as exc: - raise RuntimeError(f"cuDNN Flex Attention SDPA graph is not supported: {exc}") from exc - graph.build_plans(cudnn.build_plan_policy.HEURISTICS_CHOICE) - return max(graph.get_workspace_size(), 1) + workspace_size, _ = cudnn_pygraph.finalize_plans(graph) + return workspace_size def _execute_cudnn_graph( @@ -285,19 +237,8 @@ def _execute_cudnn_graph( device: torch.device, ): """Execute a built cuDNN frontend Python graph.""" - cudnn = _import_cudnn_frontend() - - if device.type == "cuda" and device.index is None: - device = torch.device("cuda", torch.cuda.current_device()) - workspace = torch.empty( - workspace_size, - device=device, - dtype=torch.uint8, - ) - graph.execute( - variant_pack, - workspace, - handle=_get_cudnn_current_stream_handle(cudnn, device), + cudnn_pygraph.execute_graph( + graph, variant_pack, workspace_size, device, backend_name=_BACKEND ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 51b1436ca57..fc6e4327420 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -37,6 +37,8 @@ import torch from packaging.version import InvalidVersion, Version as PkgVersion +from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph + __all__ = [ "is_frost_attention_available", "is_frost_attention_supported", @@ -72,52 +74,22 @@ _cudnn = None _availability: Optional[Tuple[bool, str]] = None _PLAN_CACHE: dict = {} -_HANDLES: dict = {} +_HANDLES = cudnn_pygraph._handles # pylint: disable=protected-access def _import_cudnn(): - """Import cuDNN Frontend with FROST engines enabled, once. - - The switch is set before the import because the documentation describes the engines as - registering at import time. Measured on B200 with cuDNN Frontend 1.29.0, the ordering turns - out not to matter: importing cudnn and cudnn.sdpa first with the switch unset, then setting - it and building a plan, still selects a FROST engine. Setting it first is kept because it is - what the documentation asks for and costs nothing, but nothing here depends on winning that - race, and _select_frost_plan verifies the engine by plan name regardless. - """ - global _cudnn - if _cudnn is None: - # Must be set before the import: the engines are registered at import time. - os.environ.setdefault("CUDNN_FRONTEND_ENABLE_FROST_ENGINES", "1") - import cudnn # pylint: disable=import-outside-toplevel - import cudnn.sdpa # noqa: F401 pylint: disable=import-outside-toplevel,unused-import + """Import cuDNN Frontend with the FROST engines registered. - _cudnn = cudnn - return _cudnn + The switch has to be set before the import because the engines register at import time, and it + also ranks FROST ahead of the backend engines, so only this backend asks for it. + _select_frost_plan verifies the engine by plan name regardless, rather than trusting the flag. + """ + return cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=True) def _handle_for(device: torch.device): - """A cuDNN handle for `device`, bound to PyTorch's current stream on it. - - Without this, cuDNN runs on its default handle's stream while the tensors and workspace are - allocated on PyTorch's current stream, and nothing orders the two. That is not hypothetical - here: the p2p CP ring issues attention inside `with torch.cuda.stream(cp_stream)`, so on - alternating ring steps the kernel and its buffers would be on different streams. Re-binding - on every call is what flex_attention.py does, and is required because the same cached plan is - executed from different streams across ring steps. - """ - if device.type != "cuda": - raise ValueError(f"FrostAttention requires CUDA tensors; got device {device}") - cudnn = _import_cudnn() - if device.index is None: - device = torch.device("cuda", torch.cuda.current_device()) - with torch.cuda.device(device): - handle = _HANDLES.get(device) - if handle is None: - handle = cudnn.create_handle() - _HANDLES[device] = handle - cudnn.set_stream(handle=handle, stream=torch.cuda.current_stream(device).cuda_stream) - return handle + """A cuDNN handle for `device`, bound to PyTorch's current stream on every call.""" + return cudnn_pygraph.handle_for(device, backend_name="FrostAttention") def _device_from_key(device_key) -> torch.device: @@ -401,29 +373,25 @@ def _select_frost_plan(graph, token: str, what: str): dims the non-FROST plans do not exist, so an unnoticed fallback would either fail obscurely or quietly serve a different shape. """ - cudnn = _import_cudnn() - graph.create_execution_plans([cudnn.heur_mode.A]) - names = [graph.get_plan_name_at_index(i) for i in range(graph.get_execution_plan_count())] - hits = [i for i, n in enumerate(names) if token in n] - if not hits: - # Both versions, because either floor can cause this and blaming one misdirects. Looked - # up defensively: this is the message explaining a failure, so it must not raise itself. - raise RuntimeError( - f"no cuDNN FROST {what} engine was offered (looked for {token!r}). Candidate plans:" - f" {names[:6]}." - f" nvidia-cudnn-frontend={_pkg_version('nvidia-cudnn-frontend', _cudnn)[1] or 'unknown'} (floor" - f" {_MIN_CUDNN_FRONTEND})," - f" nvidia-cutlass-dsl={_pkg_version('nvidia-cutlass-dsl')[1] or 'unknown'} (floor" - f" {_MIN_CUTLASS_DSL})." + # Both versions, because either floor can cause this and blaming one misdirects. Looked up + # defensively: this explains a failure, so it must not raise itself. + def hint(): + return ( + f"nvidia-cudnn-frontend=" + f"{_pkg_version('nvidia-cudnn-frontend', _cudnn)[1] or 'unknown'}" + f" (floor {_MIN_CUDNN_FRONTEND})," + f" nvidia-cutlass-dsl={_pkg_version('nvidia-cutlass-dsl')[1] or 'unknown'}" + f" (floor {_MIN_CUTLASS_DSL})." ) - # select_plan before check_support, not after: check_support is scoped to the *selected* - # plan, so calling it first would answer for whichever plan the heuristic ranked at index 0. - # Pinning also makes build_plans strict -- a decline raises instead of walking on to a - # non-FROST plan, which is the fallback this selection exists to prevent. - graph.select_plan(hits[0]) - graph.check_support() - graph.build_plans() - return names[hits[0]] + + cudnn = _import_cudnn() + _, name = cudnn_pygraph.finalize_plans( + graph, + heuristics=[cudnn.heur_mode.A], + require_plan_token=token, + not_found_hint=hint, + ) + return name def _build_fwd(key) -> dict: @@ -432,14 +400,10 @@ def _build_fwd(key) -> dict: # deterministic is unused here: it selects a backward algorithm. Callers pass False for the # forward so the two never split the forward cache. *_device, b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks, _deterministic = key - io_dt = _cudnn_dtype(dtype) shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] - graph = cudnn.pygraph( - io_data_type=io_dt, - intermediate_data_type=cudnn.data_type.FLOAT, - compute_data_type=cudnn.data_type.FLOAT, - handle=_handle_for(_device_from_key(_device)), + graph = cudnn_pygraph.build_pygraph( + dtype, _device_from_key(_device), backend_name="FrostAttention" ) tq = graph.tensor(name="q", dim=shq, stride=list(qs)) tk = graph.tensor(name="k", dim=shkv, stride=list(ks)) @@ -475,11 +439,8 @@ def _build_bwd(key) -> dict: io_dt = _cudnn_dtype(dtype) shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] - graph = cudnn.pygraph( - io_data_type=io_dt, - intermediate_data_type=cudnn.data_type.FLOAT, - compute_data_type=cudnn.data_type.FLOAT, - handle=_handle_for(_device_from_key(_device)), + graph = cudnn_pygraph.build_pygraph( + dtype, _device_from_key(_device), backend_name="FrostAttention" ) handles = {} # o and dO share q's layout; k, v and their grads share k's. From 42fa333e522fbddac0adeafcd1525c06be10aea8 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 30 Sep 2026 19:11:48 -0700 Subject: [PATCH 36/97] feat(attention): let the flex cuDNN graphs carry a diagonal-band mask Groundwork for serving head_dim 512 through flex_attention.py's graph code, per the review discussion on #3527. Measured on B200: with a FROST plan pinned, that path already runs d512 correctly for no_mask, and with a band injected it is correct across causal, bottom-right and sliding window, square and rectangular, forward and backward, against a float64 reference. The builders, their getters and both cache keys now take an optional (attn_mask_type, window) pair. It defaults to None, which keeps the existing score_mod behaviour exactly: the node is still built with use_causal_mask=False and no band. Three things worth calling out: - The mask is part of the cache key. flex's key had no mask field, so once a band exists, two graphs differing only in mask type would collide and the second would silently reuse the first. That is a wrong answer, not a cache miss. - A band and a score_mod cannot share a graph; cuDNN rejects the pair outright. _mask_or_score_mod_kwargs refuses it here with a clearer message than the frontend's. - The band translation moved to cudnn_pygraph, including the off-by-one: cuDNN's left bound counts the diagonal and TE's window_size does not, so a window of w is a left bound of w + 1. Getting that wrong drops one token of context per layer and no shape-level test would see it. The builders and their cache keys are splatted from one positional tuple, so a static check that the getter, builder and key signatures still line up is part of the verification rather than something to eyeball. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/cudnn_pygraph.py | 27 +++++++++- .../dot_product_attention/flex_attention.py | 51 ++++++++++++++++--- .../dot_product_attention/frost_attention.py | 15 +----- 3 files changed, 72 insertions(+), 21 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index 08e29f10389..ae64de77ddb 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -11,7 +11,8 @@ holding one handle per device on PyTorch's current stream, describing an SBHD/BSHD tensor in the BHSD form cuDNN wants, finalizing plans, and executing. -This module is that common part. It contains no attention semantics. +This module is that common part: the plumbing, plus the one piece of shared attention +vocabulary, translating a TE mask type and window into cuDNN's diagonal band. """ from typing import Any, Dict, Optional, Sequence, Tuple @@ -129,6 +130,30 @@ def bhsd_graph_tensor(graph, tensor: torch.Tensor, tensor_format: str): return graph.tensor(dim=dim, stride=stride, data_type=tensor.dtype) +def diagonal_band_kwargs(cudnn, attn_mask_type: str, window: Tuple[int, int]) -> Dict[str, Any]: + """cuDNN sdpa kwargs for a TE (mask type, window): a diagonal alignment plus a band. + + Note the off-by-one. cuDNN's left bound counts the diagonal itself and TE's window_size does + not, so a window of w becomes a left bound of w + 1. Passing it through unconverted silently + drops one token of context per layer, which no shape-level test would catch. + + These kwargs are mutually exclusive with score_mod: cuDNN rejects a graph carrying both with + "Attention score mod enabled and hence other subgraphs are disabled". + """ + left, right = window + opts: Dict[str, Any] = {} + if attn_mask_type in ("causal", "causal_bottom_right") or right == 0: + opts["diagonal_alignment"] = ( + cudnn.diagonal_alignment.BOTTOM_RIGHT + if attn_mask_type == "causal_bottom_right" + else cudnn.diagonal_alignment.TOP_LEFT + ) + opts["diagonal_band_right_bound"] = 0 + if left != -1: + opts["diagonal_band_left_bound"] = left + 1 + return opts + + def finalize_plans( graph, *, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index e5decacd84f..f441853788f 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -161,6 +161,26 @@ def _score_mod_bhsd_tensor_metadata(tensor: torch.Tensor, tensor_format: str) -> return (dim, stride, tensor.dtype, _score_mod_device_key(tensor.device)) +def _mask_or_score_mod_kwargs( + mask_spec: Optional[Tuple[str, Tuple[int, int]]], wrapped_score_mod +) -> Dict[str, Any]: + """SDPA kwargs for exactly one of a diagonal band or a score_mod. + + cuDNN rejects a graph carrying both ("Attention score mod enabled and hence other subgraphs + are disabled"), so this refuses the combination here with a clearer message than the frontend + gives, rather than building a graph that cannot be served. + """ + if mask_spec is None: + return {"use_causal_mask": False, "score_mod": wrapped_score_mod} + if wrapped_score_mod is not None: + raise ValueError( + "a diagonal-band mask and a score_mod cannot be combined in one cuDNN SDPA graph; " + f"got mask_spec={mask_spec!r} alongside a score_mod" + ) + cudnn = _import_cudnn_frontend() + return cudnn_pygraph.diagonal_band_kwargs(cudnn, mask_spec[0], mask_spec[1]) + + def _make_cudnn_graph_tensor_dict(graph, tensors: Optional[Dict[str, torch.Tensor]]): """Create cuDNN graph tensors matching runtime tensors.""" if tensors is None: @@ -254,6 +274,7 @@ def _cudnn_score_mod_fwd_cache_key( score_mod_tensors: Optional[Dict[str, torch.Tensor]], output_layer: torch.Tensor, stats: Optional[torch.Tensor], + mask_spec: Optional[Tuple[str, Tuple[int, int]]] = None, ) -> Optional[Tuple[Any, ...]]: """Pre-build cache key for score_mod fprop execution plans. @@ -276,6 +297,9 @@ def _cudnn_score_mod_fwd_cache_key( _score_mod_bhsd_tensor_metadata(output_layer, q_format), _score_mod_tensor_metadata(stats) if stats is not None else None, _score_mod_tensor_dict_metadata(score_mod_tensors), + # The mask belongs in the key. Without it two graphs differing only in mask type collide + # and the second silently reuses the first, which is a wrong answer rather than a miss. + mask_spec, ) @@ -294,6 +318,7 @@ def _cudnn_score_mod_bwd_cache_key( score_mod_tensors: Optional[Dict[str, torch.Tensor]], score_mod_bprop_tensors: Optional[Dict[str, torch.Tensor]], deterministic: bool, + mask_spec: Optional[Tuple[str, Tuple[int, int]]] = None, ) -> Optional[Tuple[Any, ...]]: """Pre-build cache key for score_mod bprop execution plans.""" score_mod_key = _score_mod_callback_cache_key(score_mod) @@ -316,6 +341,7 @@ def _cudnn_score_mod_bwd_cache_key( _score_mod_tensor_metadata(stats), _score_mod_tensor_dict_metadata(score_mod_tensors), _score_mod_tensor_dict_metadata(score_mod_bprop_tensors), + mask_spec, ) @@ -331,8 +357,15 @@ def _build_cudnn_score_mod_fwd_graph( score_mod_tensors: Optional[Dict[str, torch.Tensor]], output_layer: torch.Tensor, stats: Optional[torch.Tensor], + mask_spec: Optional[Tuple[str, Tuple[int, int]]] = None, ) -> _CudnnScoreModFwdGraphEntry: - """Build a cached cuDNN frontend graph for score_mod fprop.""" + """Build a cached cuDNN frontend graph for score_mod fprop. + + ``mask_spec`` is an optional (attn_mask_type, window) pair. When given, the SDPA node carries + cuDNN's diagonal band instead of the unmasked default, which is how a backend without a + score_mod expresses causal, bottom-right and sliding-window attention. The two are mutually + exclusive: cuDNN rejects a graph carrying both. + """ cudnn = _import_cudnn_frontend() graph = _build_cudnn_pygraph(query_layer.dtype, query_layer.device) @@ -344,6 +377,7 @@ def _build_cudnn_score_mod_fwd_graph( wrapped_score_mod = _wrap_score_mod(score_mod, score_mod_graph_tensors) output_dim, output_stride = _bhsd_dim_stride(output_layer, q_format) + sdpa_kwargs = _mask_or_score_mod_kwargs(mask_spec, wrapped_score_mod) output, stats_tensor = graph.sdpa( name="te_score_mod_sdpa", q=q, @@ -351,8 +385,7 @@ def _build_cudnn_score_mod_fwd_graph( v=v, generate_stats=is_training, attn_scale=attn_scale, - use_causal_mask=False, - score_mod=wrapped_score_mod, + **sdpa_kwargs, ) output.set_output(True).set_dim(output_dim).set_stride(output_stride) @@ -389,6 +422,7 @@ def _get_cudnn_score_mod_fwd_graph( score_mod_tensors: Optional[Dict[str, torch.Tensor]], output_layer: torch.Tensor, stats: Optional[torch.Tensor], + mask_spec: Optional[Tuple[str, Tuple[int, int]]] = None, ) -> _CudnnScoreModFwdGraphEntry: """Return a cached cuDNN frontend graph for score_mod fprop.""" build_args = ( @@ -403,6 +437,7 @@ def _get_cudnn_score_mod_fwd_graph( score_mod_tensors, output_layer, stats, + mask_spec, ) key = _cudnn_score_mod_fwd_cache_key(*build_args) if key is None: @@ -429,8 +464,11 @@ def _build_cudnn_score_mod_bwd_graph( score_mod_tensors: Optional[Dict[str, torch.Tensor]], score_mod_bprop_tensors: Optional[Dict[str, torch.Tensor]], deterministic: bool, + mask_spec: Optional[Tuple[str, Tuple[int, int]]] = None, ) -> _CudnnScoreModBwdGraphEntry: - """Build a cached cuDNN frontend graph for score_mod bprop.""" + """Build a cached cuDNN frontend graph for score_mod bprop. See the fprop builder for + ``mask_spec``; the backward must carry the same mask as the forward or the gradients are + computed against a different attention.""" graph = _build_cudnn_pygraph(query_layer.dtype, query_layer.device) q = _bhsd_graph_tensor(graph, query_layer, q_format) k = _bhsd_graph_tensor(graph, key_layer, kv_format) @@ -463,8 +501,7 @@ def _build_cudnn_score_mod_bwd_graph( dO=d_output, stats=stats_tensor, attn_scale=attn_scale, - use_causal_mask=False, - score_mod=wrapped_score_mod, + **_mask_or_score_mod_kwargs(mask_spec, wrapped_score_mod), score_mod_bprop=wrapped_score_mod_bprop, use_deterministic_algorithm=deterministic, ) @@ -505,6 +542,7 @@ def _get_cudnn_score_mod_bwd_graph( score_mod_tensors: Optional[Dict[str, torch.Tensor]], score_mod_bprop_tensors: Optional[Dict[str, torch.Tensor]], deterministic: bool, + mask_spec: Optional[Tuple[str, Tuple[int, int]]] = None, ) -> _CudnnScoreModBwdGraphEntry: """Return a cached cuDNN frontend graph for score_mod bprop.""" build_args = ( @@ -522,6 +560,7 @@ def _get_cudnn_score_mod_bwd_graph( score_mod_tensors, score_mod_bprop_tensors, deterministic, + mask_spec, ) key = _cudnn_score_mod_bwd_cache_key(*build_args) if key is None: diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index fc6e4327420..10f07a88cd0 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -224,20 +224,7 @@ def _mask_spec(attn_mask_type: str, window_size=None): def _mask_options(cudnn, spec): """cuDNN sdpa kwargs for a (mask type, window) spec: a diagonal alignment plus a band.""" attn_mask_type, window = spec - left, right = window - options = {} - if attn_mask_type in ("causal", "causal_bottom_right") or right == 0: - options["diagonal_alignment"] = ( - cudnn.diagonal_alignment.BOTTOM_RIGHT - if attn_mask_type == "causal_bottom_right" - else cudnn.diagonal_alignment.TOP_LEFT - ) - options["diagonal_band_right_bound"] = 0 - if left != -1: - # cuDNN counts the diagonal itself, TE does not, hence the +1 -- the same convention the - # C++ fused path and the Python port both use. - options["diagonal_band_left_bound"] = left + 1 - return options + return cudnn_pygraph.diagonal_band_kwargs(cudnn, attn_mask_type, window) def is_frost_attention_supported( From 084f5282efa948252466a7a33b9c8c87c4bd3467 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 30 Sep 2026 20:11:25 -0700 Subject: [PATCH 37/97] fix(attention): stop preparing the FROST graphs twice The extraction moved validate() and build_operation_graph() into cudnn_pygraph.finalize_plans, taking them from flex's _finalize_cudnn_graph where they lived. FROST's builders called both themselves, because its old _select_frost_plan did not, so after the refactor each FROST graph was validated and lowered twice. Found by counting what the two builders would still share if they were merged, not by a test: the duplicate pair showed up as lines present in one builder and absent from the other for no reason. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../attention/dot_product_attention/frost_attention.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 10f07a88cd0..7874de3aa48 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -408,8 +408,6 @@ def _build_fwd(key) -> dict: tlse.set_output(True).set_dim([b, hq, sq, 1]).set_stride([hq * sq, sq, 1, 1]).set_data_type( cudnn.data_type.FLOAT ) - graph.validate() - graph.build_operation_graph() plan = _select_frost_plan(graph, _FROST_FWD_PLAN_TOKEN, "forward") return { "graph": graph, @@ -459,8 +457,6 @@ def _build_bwd(key) -> dict: ) for tensor, stride in ((tdq, qs), (tdk, ks), (tdv, ks)): tensor.set_output(True).set_data_type(io_dt).set_stride(list(stride)) - graph.validate() - graph.build_operation_graph() plan = _select_frost_plan(graph, _FROST_BWD_PLAN_TOKEN, "backward") handles["dq"], handles["dk"], handles["dv"] = tdq, tdk, tdv return { From 1f43e88d67d22cf951901cc78f714cffa6149d54 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 30 Sep 2026 20:43:11 -0700 Subject: [PATCH 38/97] fix(attention): keep the flex builder call shape, and frost's cudnn handle Two regressions from the extraction, both caught by running the suites. flex: widening build_args to carry mask_spec changed the arity of every builder call, which broke the cache tests that substitute their own builder ("fake_build() takes 11 positional arguments but 12 were given"). mask_spec is now passed only when it is set, so the call shape is byte-identical for every existing caller and the mocks keep working. frost: _import_cudnn stopped assigning the module's _cudnn global once it delegated, leaving it permanently None. _pkg_version falls back to that module's __version__ when distribution metadata is unavailable, which is how a source or vendored install avoids being misreported as absent, so the binding is restored. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/flex_attention.py | 18 ++++++++++-------- .../dot_product_attention/frost_attention.py | 6 +++++- 2 files changed, 15 insertions(+), 9 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index f441853788f..8b3bcca5e7b 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -437,14 +437,16 @@ def _get_cudnn_score_mod_fwd_graph( score_mod_tensors, output_layer, stats, - mask_spec, ) - key = _cudnn_score_mod_fwd_cache_key(*build_args) + # Only when set: an unconditional extra argument would change the call shape for every + # existing caller, including the tests that substitute their own builder. + extra = {} if mask_spec is None else {"mask_spec": mask_spec} + key = _cudnn_score_mod_fwd_cache_key(*build_args, **extra) if key is None: - return _build_cudnn_score_mod_fwd_graph(*build_args) + return _build_cudnn_score_mod_fwd_graph(*build_args, **extra) entry = _cudnn_score_mod_graph_cache.get(key) if entry is None: - entry = _build_cudnn_score_mod_fwd_graph(*build_args) + entry = _build_cudnn_score_mod_fwd_graph(*build_args, **extra) _cudnn_score_mod_graph_cache[key] = entry return entry @@ -560,14 +562,14 @@ def _get_cudnn_score_mod_bwd_graph( score_mod_tensors, score_mod_bprop_tensors, deterministic, - mask_spec, ) - key = _cudnn_score_mod_bwd_cache_key(*build_args) + extra = {} if mask_spec is None else {"mask_spec": mask_spec} + key = _cudnn_score_mod_bwd_cache_key(*build_args, **extra) if key is None: - return _build_cudnn_score_mod_bwd_graph(*build_args) + return _build_cudnn_score_mod_bwd_graph(*build_args, **extra) entry = _cudnn_score_mod_graph_cache.get(key) if entry is None: - entry = _build_cudnn_score_mod_bwd_graph(*build_args) + entry = _build_cudnn_score_mod_bwd_graph(*build_args, **extra) _cudnn_score_mod_graph_cache[key] = entry return entry diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 7874de3aa48..57426a08e3f 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -84,7 +84,11 @@ def _import_cudnn(): also ranks FROST ahead of the backend engines, so only this backend asks for it. _select_frost_plan verifies the engine by plan name regardless, rather than trusting the flag. """ - return cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=True) + global _cudnn # pylint: disable=global-statement + # Kept bound: _pkg_version falls back to the module's __version__ when distribution metadata + # is unavailable, which is how a source or vendored install avoids being misreported. + _cudnn = cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=True) + return _cudnn def _handle_for(device: torch.device): From 469bad9475426001954d21e5b2295b8d9b210211 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 1 Oct 2026 01:19:49 -0700 Subject: [PATCH 39/97] fix(attention): enable the FROST engines whichever backend imports cuDNN first Sharing one cuDNN import between flex and frost made the FROST switch first-caller-wins: the enabling sat inside the "already imported?" memo, so a process that ran a score_mod layer first left frost with a cuDNN offering it no engine. That surfaced as "no cuDNN engine matching 'sdpa_fwd_prefill_sm100' was offered" on the first head_dim 512 forward, pointing at package versions that were fine. Each backend had its own import before the extraction, so this was new. Enabling late is sound: cuDNN Frontend 1.29.0 reads the switch per graph, inside engines/manifest.py offered_ids(), not at import time. Also corrects three comments that asserted things the code does not do: - the engines do not register at import time, the switch is read at planning time - cuDNN rejects score_mod plus a diagonal band in its backward node only; the forward composes both silently, so refusing the pair is our choice and the reason is forward/backward symmetry - flex cannot opt out of the FROST ranking by not asking for it, since the switch is process-wide Adds the signature-alignment test an earlier commit message claimed but never committed, and a test for the ordering above. Both are CPU-only. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_flex_attention.py | 29 ++++++++ .../pytorch/attention/test_frost_attention.py | 64 +++++++++++++++++ .../dot_product_attention/cudnn_pygraph.py | 69 +++++++++++++------ .../dot_product_attention/flex_attention.py | 6 +- .../dot_product_attention/frost_attention.py | 11 +-- 5 files changed, 152 insertions(+), 27 deletions(-) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index beed4069917..98233b1d215 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -705,3 +705,32 @@ def test_dot_product_attention_score_mod(dtype, qkv_format, score_mod_case, scal torch.testing.assert_close(q.grad, q_ref.grad, **tols) torch.testing.assert_close(k.grad, k_ref.grad, **tols) torch.testing.assert_close(v.grad, v_ref.grad, **tols) + + +@pytest.mark.parametrize("direction", ["fwd", "bwd"]) +def test_score_mod_graph_signatures_stay_aligned(direction): + """The cache key, the builder and the getter are splatted from one positional tuple. + + `_get_cudnn_score_mod_*_graph` passes the same `build_args` tuple to the cache key and to the + builder, so the three parameter lists have to stay in the same order. Nothing enforced that, + and the failure is quiet in the worst direction: a parameter inserted in one signature and not + another shifts the rest by one, and a shifted *cache key* is not a crash, it is two different + configurations sharing a cached graph. + + No GPU: this reads signatures only. + """ + import inspect + + names = [ + getattr(flex_attention, "_cudnn_score_mod_%s_cache_key" % direction), + getattr(flex_attention, "_build_cudnn_score_mod_%s_graph" % direction), + getattr(flex_attention, "_get_cudnn_score_mod_%s_graph" % direction), + ] + signatures = [list(inspect.signature(fn).parameters) for fn in names] + reference = signatures[0] + for fn, params in zip(names[1:], signatures[1:]): + assert params == reference, ( + "%s takes %s but _cudnn_score_mod_%s_cache_key takes %s; these are splatted from one" + " positional tuple and must stay in the same order" + % (fn.__name__, params, direction, reference) + ) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 0ecee3b0f63..7a3f6ca0bcb 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -400,3 +400,67 @@ def test_dot_product_attention_runs_in_onnx_export_mode(): assert out.numel() == s * b * h * d assert torch.isfinite(out).all() + + +def test_frost_engines_are_enabled_even_if_flex_imported_cudnn_first(): + """Enabling the FROST engines must not depend on which backend touched cuDNN first. + + flex_attention and frost_attention share one cuDNN import in cudnn_pygraph. flex asks for the + import without the FROST engines and frost asks with them, so if the enabling sat inside the + "already imported?" memo, a process that ran a score_mod layer first would leave FROST with a + cuDNN that offers it no engine. That surfaces far from its cause, as "no cuDNN engine matching + 'sdpa_fwd_prefill_sm100' was offered" on the first head_dim 512 forward, with a hint pointing + at package versions that are in fact fine. + + No GPU and no real cuDNN: a stub stands in for the package, because what is under test is the + order-dependence of our own wrapper. It also has to run in-process with the globals reset, + since the real order is decided once per process and pytest gives us no second one. + """ + import sys + import types + + from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph + + env = "CUDNN_FRONTEND_ENABLE_FROST_ENGINES" + saved = ( + cudnn_pygraph._cudnn, + cudnn_pygraph._frost_engines_enabled, + os.environ.get(env), + sys.modules.get("cudnn"), + sys.modules.get("cudnn.sdpa"), + ) + try: + stub = types.ModuleType("cudnn") + stub.sdpa = types.ModuleType("cudnn.sdpa") + sys.modules["cudnn"] = stub + sys.modules["cudnn.sdpa"] = stub.sdpa + cudnn_pygraph._cudnn = None + cudnn_pygraph._frost_engines_enabled = False + os.environ.pop(env, None) + + # flex first, which must not enable anything. + cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=False) + assert env not in os.environ, "the non-FROST caller must not set the switch" + assert not cudnn_pygraph.frost_engines_enabled() + + # frost second, on an already-imported cuDNN. This is the case that used to be skipped. + cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=True) + assert os.environ.get(env) == "1", "FROST was requested after the import and not enabled" + assert cudnn_pygraph.frost_engines_enabled() + finally: + ( + cudnn_pygraph._cudnn, + cudnn_pygraph._frost_engines_enabled, + prior_env, + prior_cudnn, + prior_sdpa, + ) = saved + if prior_env is None: + os.environ.pop(env, None) + else: + os.environ[env] = prior_env + for name, module in (("cudnn", prior_cudnn), ("cudnn.sdpa", prior_sdpa)): + if module is None: + sys.modules.pop(name, None) + else: + sys.modules[name] = module diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index ae64de77ddb..7ad0a32233e 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -23,31 +23,33 @@ _cudnn = None +_frost_engines_enabled = False _handles: Dict[torch.device, Any] = {} def import_cudnn_frontend(enable_frost_engines: bool = False): - """Import cuDNN Frontend once, optionally with the FROST engines registered. + """Import cuDNN Frontend, enabling the FROST engines if this caller needs them. ``enable_frost_engines`` is not merely additive: the switch also ranks FROST ahead of the backend engines everywhere, so a caller that does not want FROST must not ask for it. - The switch is set before the import because the engines are documented as registering at - import time. Measured on B200 with cuDNN Frontend 1.29.0 the ordering turns out not to - matter, but setting it first is what the documentation asks for and costs nothing. Callers - that require a FROST engine should verify by plan name rather than rely on the switch, which - is what ``finalize_plans(require_plan_token=...)`` does. + The enabling is deliberately outside the import memo. Both backends call this, and whichever + one reaches it first would otherwise decide for the process: with the flag inside the memo, a + flex call would cache the module with FROST off and every later FROST call would get a cuDNN + that offers no FROST engine, which surfaces much later as "no cuDNN engine matching ... was + offered". Enabling late is sound because the switch is read per graph rather than at import: + in cuDNN Frontend 1.29.0 ``engines/manifest.py`` consults the environment inside + ``offered_ids()``, reached from ``engines_for(graph)`` on every ``create_execution_plans``. + + Note the switch is process-wide and never unset, so enabling it for FROST also reorders the + candidates a concurrent score_mod graph sees. Callers that require a particular engine should + verify by plan name rather than rely on the switch, which is what + ``finalize_plans(require_plan_token=...)`` does. """ - global _cudnn # pylint: disable=global-statement + global _cudnn, _frost_engines_enabled # pylint: disable=global-statement if _cudnn is None: - if enable_frost_engines: - os.environ.setdefault("CUDNN_FRONTEND_ENABLE_FROST_ENGINES", "1") try: import cudnn # pylint: disable=import-outside-toplevel - - if enable_frost_engines: - # pylint: disable=import-outside-toplevel,unused-import - import cudnn.sdpa # noqa: F401 except ImportError as exc: raise ImportError( "cuDNN frontend Python package not found. " @@ -55,9 +57,22 @@ def import_cudnn_frontend(enable_frost_engines: bool = False): ) from exc _cudnn = cudnn + + if enable_frost_engines and not _frost_engines_enabled: + os.environ.setdefault("CUDNN_FRONTEND_ENABLE_FROST_ENGINES", "1") + # pylint: disable=import-outside-toplevel,unused-import + import cudnn.sdpa # noqa: F401 + + _frost_engines_enabled = True + return _cudnn +def frost_engines_enabled() -> bool: + """Whether this process has enabled the FROST engines through ``import_cudnn_frontend``.""" + return _frost_engines_enabled + + def handle_for(device: torch.device, *, backend_name: str = "cuDNN attention"): """A cuDNN handle for ``device``, rebound to PyTorch's current stream on every call. @@ -137,8 +152,11 @@ def diagonal_band_kwargs(cudnn, attn_mask_type: str, window: Tuple[int, int]) -> not, so a window of w becomes a left bound of w + 1. Passing it through unconverted silently drops one token of context per layer, which no shape-level test would catch. - These kwargs are mutually exclusive with score_mod: cuDNN rejects a graph carrying both with - "Attention score mod enabled and hence other subgraphs are disabled". + These kwargs are mutually exclusive with score_mod. cuDNN enforces that in the backward node + only ("Attention score mod enabled and hence other subgraphs are disabled"); its forward node + composes the two without complaint. Callers must still refuse the pair on both sides, because + forward and backward have to carry the same mask or the gradients belong to a different + attention than the output does. """ left, right = window opts: Dict[str, Any] = {} @@ -166,12 +184,21 @@ def finalize_plans( ``require_plan_token`` makes the choice strict: only a plan whose name contains the token is acceptable, and anything else raises. That is not a stylistic preference. Without a pin, - ``build_plans`` walks the ranked list and finalizes the first plan that builds, so a graph - that a specialised engine declines would quietly run on a fallback instead, which for the - FROST head-dim range is the wrong kernel rather than a slower one. - - The pin must precede ``check_support``: that call is scoped to the *selected* plan, so running - it first would answer for whichever plan the heuristic happened to rank at index 0. + ``build_plans`` walks the ranked list from index 0 and finalizes the first plan that builds, + logging each decline at INFO, so a graph that the intended engine declines runs on whatever + cuDNN ranked next with nothing in the return value to say so. At head_dim 512 that matters in + the forward, where an ordinary engine may well build and compute a different function from the + FROST kernel. The backward is self-limiting, since no non-FROST d512 backward exists, so an + unpinned backward would fail loudly on its own. + + The token is matched as a substring rather than by equality on purpose: cuDNN has already + collapsed per-head-dim engine names (``..._d512`` and friends) into a single row once, and the + substring test survived that. + + Pinning also changes what ``check_support`` means. Selecting a plan sets cuDNN's internal + ``_plan_pinned``, and only then is a decline fatal; unpinned, cuDNN records the decline and + keeps walking. So the pin has to come first both because the check is scoped to the selected + plan and because it is what makes the check binding at all. """ cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index 8b3bcca5e7b..b500031039f 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -22,8 +22,10 @@ def _import_cudnn_frontend(): """Import the cuDNN frontend Python package.""" - # Without the FROST engines: enabling them also ranks them ahead of the backend engines - # everywhere, which would change which plan this path runs. + # This path does not ask for the FROST engines, but asking is all it controls: the switch is + # process-wide, so a FrostAttention call elsewhere in the process, or a user setting + # CUDNN_FRONTEND_ENABLE_FROST_ENGINES themselves, still ranks FROST ahead of the backend + # engines for the graphs built here. return cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=False) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 57426a08e3f..1fef234841b 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -360,15 +360,18 @@ def _check_kv_match(k: torch.Tensor, v: torch.Tensor) -> None: def _select_frost_plan(graph, token: str, what: str): """Select a plan whose name proves a FROST engine was chosen. - Falling back to whatever plan happens to be first would defeat the purpose: at these head - dims the non-FROST plans do not exist, so an unnoticed fallback would either fail obscurely - or quietly serve a different shape. + Falling back to whatever plan happens to be first would defeat the purpose. A too-old + nvidia-cutlass-dsl makes the FROST engines decline silently, and in the forward an ordinary + engine may then build and compute something else; the pin turns that into a named error at + the first forward rather than a wrong number or a backward that fails later for no visible + reason. """ # Both versions, because either floor can cause this and blaming one misdirects. Looked up # defensively: this explains a failure, so it must not raise itself. def hint(): return ( - f"nvidia-cudnn-frontend=" + f"Wanted the FROST {what} engine." + f" nvidia-cudnn-frontend=" f"{_pkg_version('nvidia-cudnn-frontend', _cudnn)[1] or 'unknown'}" f" (floor {_MIN_CUDNN_FRONTEND})," f" nvidia-cutlass-dsl={_pkg_version('nvidia-cutlass-dsl')[1] or 'unknown'}" From 626bde2e17e86f5351249c5a737054ca2941b253 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 1 Oct 2026 01:28:12 -0700 Subject: [PATCH 40/97] test(attention): cover the flex mask_spec path on CPU mask_spec was threaded through both builders, both cache keys and an exclusivity check with no test and no caller, so the only thing exercising it was the FROST suite, which skips on every machine without a Blackwell GPU. These run anywhere: the band translation for causal, bottom-right and sliding window including the off-by-one cuDNN needs, the refusal when a mask and a score_mod arrive together, and that the no-mask path is unchanged. Also corrects the exclusivity docstring: cuDNN refuses the pair in its backward node, not generally. Refusing it on both sides is our choice, because the backward must carry the same mask as the forward. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_flex_attention.py | 57 +++++++++++++++++++ .../dot_product_attention/flex_attention.py | 7 ++- 2 files changed, 61 insertions(+), 3 deletions(-) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index 98233b1d215..bcde093e50b 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -734,3 +734,60 @@ def test_score_mod_graph_signatures_stay_aligned(direction): " positional tuple and must stay in the same order" % (fn.__name__, params, direction, reference) ) + + +@pytest.mark.parametrize( + "mask_spec,expected", + [ + # Causal: top-left aligned, right bound pinned to the diagonal, no left bound. + (("causal", (-1, 0)), {"diagonal_alignment": "TOP_LEFT", "diagonal_band_right_bound": 0}), + # Bottom-right causal, which is what KV trimming produces whenever SKV > SQ. + ( + ("causal_bottom_right", (-1, 0)), + {"diagonal_alignment": "BOTTOM_RIGHT", "diagonal_band_right_bound": 0}, + ), + # Sliding window. cuDNN's left bound counts the diagonal itself and TE's window_size does + # not, so 511 must arrive as 512. Getting this wrong drops one token of context per layer + # and no shape-level test would notice. + ( + ("causal", (511, 0)), + { + "diagonal_alignment": "TOP_LEFT", + "diagonal_band_right_bound": 0, + "diagonal_band_left_bound": 512, + }, + ), + # No mask at all: no alignment, no bounds. + (("no_mask", (-1, -1)), {}), + ], +) +def test_mask_spec_translates_to_a_diagonal_band(mask_spec, expected): + """A mask_spec must become the cuDNN band kwargs, and never a score_mod. + + No GPU: this builds no graph, it checks the kwargs the graph would be given. + """ + cudnn = flex_attention._import_cudnn_frontend() + got = flex_attention._mask_or_score_mod_kwargs(mask_spec, None) + + assert "score_mod" not in got and "use_causal_mask" not in got + for key, want in expected.items(): + if key == "diagonal_alignment": + assert got[key] == getattr(cudnn.diagonal_alignment, want) + else: + assert got[key] == want + assert set(got) == set(expected) + + +def test_mask_spec_and_score_mod_cannot_be_combined(): + """cuDNN's backward refuses the pair, so flex must refuse it before building the graph.""" + with pytest.raises(ValueError, match="cannot be combined"): + flex_attention._mask_or_score_mod_kwargs(("causal", (-1, 0)), lambda *a, **k: None) + + +def test_no_mask_spec_still_takes_the_score_mod_path(): + """The default path must be byte-identical to what it was before mask_spec existed.""" + sentinel = object() + assert flex_attention._mask_or_score_mod_kwargs(None, sentinel) == { + "use_causal_mask": False, + "score_mod": sentinel, + } diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index b500031039f..881f6d22598 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -168,9 +168,10 @@ def _mask_or_score_mod_kwargs( ) -> Dict[str, Any]: """SDPA kwargs for exactly one of a diagonal band or a score_mod. - cuDNN rejects a graph carrying both ("Attention score mod enabled and hence other subgraphs - are disabled"), so this refuses the combination here with a clearer message than the frontend - gives, rather than building a graph that cannot be served. + cuDNN rejects the pair in its backward node ("Attention score mod enabled and hence other + subgraphs are disabled") while its forward node composes both silently. Refusing it here on + both sides is deliberate: the backward has to carry the same mask as the forward, or the + gradients belong to a different attention than the output does. """ if mask_spec is None: return {"use_causal_mask": False, "score_mod": wrapped_score_mod} From 40b8a2aa89aad1cdb6141d08742b9fdc7fd189a3 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 1 Oct 2026 01:36:06 -0700 Subject: [PATCH 41/97] fix(attention): say why a pinned cuDNN engine declined the graph The strict path has two failures that read very differently. The engine not being offered is already explained. The engine being offered and then refusing was escaping bare, so the message said neither which engine judged the graph unservable nor what constraint it missed, although cuDNN puts its reason in the exception. Now framed with the plan name, cuDNN's reason and the same version hint the other failure gets. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_frost_attention.py | 53 +++++++++++++++++++ .../dot_product_attention/cudnn_pygraph.py | 15 +++++- 2 files changed, 66 insertions(+), 2 deletions(-) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 7a3f6ca0bcb..305c78f2739 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -464,3 +464,56 @@ def test_frost_engines_are_enabled_even_if_flex_imported_cudnn_first(): sys.modules.pop(name, None) else: sys.modules[name] = module + + +def test_pinned_plan_decline_reports_the_engine_reason(): + """A pinned engine that refuses the graph must say why, not raise bare. + + The name lookup failing and the engine declining after selection are the two ways the strict + path fails, and they read very differently: the second means the engine was there and judged + this graph unservable, so cuDNN's own reason is the only thing identifying which constraint + was missed. Without this the exception escaped with neither the reason framed nor the version + hint attached. + + No GPU: a stub graph stands in, raising the real cuDNN exception type. + """ + from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph + + cudnn = cudnn_pygraph.import_cudnn_frontend() + + class _DeclinedGraph: + """Offers the wanted plan, then refuses it at check_support.""" + + def validate(self): + pass + + def build_operation_graph(self): + pass + + def create_execution_plans(self, _heuristics): + pass + + def get_execution_plan_count(self): + return 1 + + def get_plan_name_at_index(self, _i): + return "sdpa_fwd_prefill_sm100" + + def select_plan(self, _i): + pass + + def check_support(self): + raise cudnn.cudnnGraphNotSupportedError("head_dim 512 needs SM100; this is SM90") + + with pytest.raises(RuntimeError) as excinfo: + cudnn_pygraph.finalize_plans( + _DeclinedGraph(), + heuristics=[cudnn.heur_mode.A], + require_plan_token="sdpa_fwd_prefill_sm100", + not_found_hint="nvidia-cutlass-dsl=4.8.0.", + ) + + message = str(excinfo.value) + assert "sdpa_fwd_prefill_sm100" in message, "the message must name the engine that declined" + assert "needs SM100" in message, "cuDNN's own reason must survive" + assert "nvidia-cutlass-dsl" in message, "the version hint must be attached here too" diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index 7ad0a32233e..a2e20d01539 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -231,8 +231,19 @@ def finalize_plans( f" Candidate plans: {names[:6]}.{(' ' + hint) if hint else ''}" ) graph.select_plan(hits[0]) - graph.check_support() - graph.build_plans() + # The engine is pinned, so a decline here is the engine's own verdict on this graph and cuDNN + # puts its reason in the exception. Surface that rather than letting it escape bare: a plan + # that was offered and then refused is the harder failure to read, and the reason is the only + # thing that says which constraint was missed. + try: + graph.check_support() + graph.build_plans() + except cudnn.cudnnGraphNotSupportedError as exc: + hint = not_found_hint() if callable(not_found_hint) else not_found_hint + raise RuntimeError( + f"cuDNN engine {names[hits[0]]!r} was offered but declined this graph:" + f" {exc}{(' ' + hint) if hint else ''}" + ) from exc return max(graph.get_workspace_size(), 1), names[hits[0]] From 812d4753f2dd5a2a50bf7ba36a0b621127933848 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 1 Oct 2026 08:54:33 +0000 Subject: [PATCH 42/97] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- .../pytorch/attention/dot_product_attention/cudnn_pygraph.py | 5 +++-- .../attention/dot_product_attention/flex_attention.py | 4 +--- .../attention/dot_product_attention/frost_attention.py | 3 ++- 3 files changed, 6 insertions(+), 6 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index a2e20d01539..a10ae5f508a 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -106,8 +106,9 @@ def io_data_type(cudnn, dtype: torch.dtype, *, backend_name: str = "cuDNN attent raise ValueError(f"{backend_name} only supports FP16/BF16 tensors, got {dtype}") -def build_pygraph(dtype: torch.dtype, device: torch.device, *, - backend_name: str = "cuDNN attention"): +def build_pygraph( + dtype: torch.dtype, device: torch.device, *, backend_name: str = "cuDNN attention" +): """A cuDNN frontend graph for F16/BF16 SDPA, bound to this device's stream-current handle.""" cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() return cudnn.pygraph( diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index 881f6d22598..5a15c234a5c 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -260,9 +260,7 @@ def _execute_cudnn_graph( device: torch.device, ): """Execute a built cuDNN frontend Python graph.""" - cudnn_pygraph.execute_graph( - graph, variant_pack, workspace_size, device, backend_name=_BACKEND - ) + cudnn_pygraph.execute_graph(graph, variant_pack, workspace_size, device, backend_name=_BACKEND) def _cudnn_score_mod_fwd_cache_key( diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 1fef234841b..d8c235c200a 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -366,12 +366,13 @@ def _select_frost_plan(graph, token: str, what: str): the first forward rather than a wrong number or a backward that fails later for no visible reason. """ + # Both versions, because either floor can cause this and blaming one misdirects. Looked up # defensively: this explains a failure, so it must not raise itself. def hint(): return ( f"Wanted the FROST {what} engine." - f" nvidia-cudnn-frontend=" + " nvidia-cudnn-frontend=" f"{_pkg_version('nvidia-cudnn-frontend', _cudnn)[1] or 'unknown'}" f" (floor {_MIN_CUDNN_FRONTEND})," f" nvidia-cutlass-dsl={_pkg_version('nvidia-cutlass-dsl')[1] or 'unknown'}" From c981a0dab44f4e0da52c9a7bd4a5b44d15c2cdf9 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 1 Oct 2026 02:04:08 -0700 Subject: [PATCH 43/97] test(attention): skip the new cuDNN-frontend tests when the package is absent The frontend is an optional dependency and L0 now runs both modules, so a test that imports it unguarded fails the job on a machine where FROST is simply unavailable. Matches the guard the existing score_mod test already uses. Two of the six new tests could reach the import; the rest either return before it, use only inspect, or stub the module into sys.modules first. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- tests/pytorch/attention/test_flex_attention.py | 8 ++++++-- tests/pytorch/attention/test_frost_attention.py | 5 ++++- 2 files changed, 10 insertions(+), 3 deletions(-) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index bcde093e50b..ba94e4db0b8 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -764,9 +764,13 @@ def test_score_mod_graph_signatures_stay_aligned(direction): def test_mask_spec_translates_to_a_diagonal_band(mask_spec, expected): """A mask_spec must become the cuDNN band kwargs, and never a score_mod. - No GPU: this builds no graph, it checks the kwargs the graph would be given. + No GPU: this builds no graph, it checks the kwargs the graph would be given. The frontend is + an optional dependency, so skip rather than fail where it is absent. """ - cudnn = flex_attention._import_cudnn_frontend() + try: + cudnn = flex_attention._import_cudnn_frontend() + except ImportError: + pytest.skip("cuDNN frontend Python package is required for the diagonal-band kwargs.") got = flex_attention._mask_or_score_mod_kwargs(mask_spec, None) assert "score_mod" not in got and "use_causal_mask" not in got diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 305c78f2739..9492267e78a 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -479,7 +479,10 @@ def test_pinned_plan_decline_reports_the_engine_reason(): """ from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph - cudnn = cudnn_pygraph.import_cudnn_frontend() + try: + cudnn = cudnn_pygraph.import_cudnn_frontend() + except ImportError: + pytest.skip("cuDNN frontend Python package is required for the decline-reason path.") class _DeclinedGraph: """Offers the wanted plan, then refuses it at check_support.""" From 73a205bc2f67ec2567ef4ea997e263baed234f9d Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 1 Oct 2026 02:45:06 -0700 Subject: [PATCH 44/97] fix(attention): stop flex graphs running on a FROST engine flex declined to ask for the FROST engines but could not avoid them: the switch that offers them is process-wide, so once anything enables it they are ranked ahead of the backend engines for flex's graphs too. They accept a score_mod graph, pass check_support, build, and then compute without the callback. Measured on B200 with cuDNN Frontend 1.29.0, score_mod bf16, switch on. A FROST plan ranks at index 0 and an unpinned build selects it at every head dim, and the result tracks a float64 reference computed WITHOUT the bias: head_dim unpinned with the engines barred 64 dropped correct 128 dropped correct 256 dropped correct 512 dropped correct So flex bars them explicitly now, through a new exclude_plan_tokens on the shared finalizer. It is inert where those engines are not on offer, which is every process that has not enabled them. Also stops the FROST availability probe from enabling them. It only reads a version off the module, and enabling there reordered plan selection for the whole process even when the checks that follow went on to decline FROST. Upstream cause is nvbug 6856051: the FROST engines do have a score_mod capability gate, but it reads a dict key the sdpa() path never writes, so it never fires. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_flex_attention.py | 90 +++++++++++++++++++ .../dot_product_attention/cudnn_pygraph.py | 12 +++ .../dot_product_attention/flex_attention.py | 13 ++- .../dot_product_attention/frost_attention.py | 18 ++-- 4 files changed, 125 insertions(+), 8 deletions(-) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index ba94e4db0b8..52114e0a0e7 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -795,3 +795,93 @@ def test_no_mask_spec_still_takes_the_score_mod_path(): "use_causal_mask": False, "score_mod": sentinel, } + + +def test_flex_bars_the_frost_engines(): + """flex must tell cuDNN not to use a FROST engine, not merely decline to ask for them. + + The switch that offers those engines is process-wide, so a FrostAttention call elsewhere in the + process, or a user setting CUDNN_FRONTEND_ENABLE_FROST_ENGINES, puts them ahead of the backend + engines for these graphs too. They accept a score_mod graph, pass check_support, build, and + then compute without the callback. + + No GPU: this checks the instruction is passed, not what cuDNN does with it. + """ + from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph + + seen = {} + + def fake_finalize(graph, **kwargs): + seen.update(kwargs) + return 4096, None + + original = cudnn_pygraph.finalize_plans + cudnn_pygraph.finalize_plans = fake_finalize + try: + assert flex_attention._finalize_cudnn_graph(object()) == 4096 + finally: + cudnn_pygraph.finalize_plans = original + + excluded = seen.get("exclude_plan_tokens") + assert excluded, "flex did not ask cuDNN to exclude any engine" + assert "sdpa_fwd_prefill_sm100" in excluded and "sdpa_bwd_sm100" in excluded, excluded + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required.") +def test_frost_switch_does_not_change_what_flex_computes(): + """Enabling the FROST engines must not change flex's output. + + This is the property the silent drop violated: with the engines on, an unpinned build selected + a FROST plan at every head dim measured on B200, and that plan returns plain attention with the + score_mod discarded. Comparing flex against itself across the switch needs no reference and no + knowledge of which plan ran; if the two differ, a different kernel answered. + """ + try: + flex_attention._import_cudnn_frontend() + except ImportError: + pytest.skip("cuDNN frontend Python package is required for score_mod attention.") + + env = "CUDNN_FRONTEND_ENABLE_FROST_ENGINES" + saved = os.environ.get(env) + torch.manual_seed(0) + b, h, s, d = 2, 4, 512, 64 + dtype = torch.bfloat16 if is_bf16_available() else torch.float16 + q, k, v = (torch.randn(b, s, h, d, device="cuda", dtype=dtype) for _ in range(3)) + + def bias_score_mod(score_mod_graph, score_tensor, _tensors): + """score += (row - col). Self-contained, and large enough that dropping it is obvious.""" + cudnn = flex_attention._import_cudnn_frontend() + row = score_mod_graph.gen_index(input=score_tensor, axis=2) + row.set_data_type(cudnn.data_type.INT32) + col = score_mod_graph.gen_index(input=score_tensor, axis=3) + col.set_data_type(cudnn.data_type.INT32) + bias = score_mod_graph.sub(a=row, b=col, compute_data_type=cudnn.data_type.FLOAT) + bias.set_data_type(cudnn.data_type.FLOAT) + return score_mod_graph.add( + a=score_tensor, b=bias, compute_data_type=cudnn.data_type.FLOAT + ) + + def run(): + flex_attention._cudnn_score_mod_graph_cache.clear() + return flex_attention.FusedAttentionWithScoreModFunc.apply( + False, q, k, v, "bshd", "bshd", d**-0.5, bias_score_mod, None, None, None, False + ) + + try: + os.environ.pop(env, None) + without = run() + os.environ[env] = "1" + with_engines = run() + finally: + flex_attention._cudnn_score_mod_graph_cache.clear() + if saved is None: + os.environ.pop(env, None) + else: + os.environ[env] = saved + + torch.testing.assert_close( + with_engines, + without, + msg=lambda m: "flex computed something different with the FROST engines enabled, which" + " means a FROST plan answered and dropped the score_mod:\n" + m, + ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index a10ae5f508a..7a07a08484c 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -180,6 +180,7 @@ def finalize_plans( build_policy: Any = None, require_plan_token: Optional[str] = None, not_found_hint: Any = "", + exclude_plan_tokens: Optional[Sequence[str]] = None, ) -> Tuple[int, Optional[str]]: """Create plans, optionally pin one by name, build, and return (workspace size, plan name). @@ -196,6 +197,12 @@ def finalize_plans( collapsed per-head-dim engine names (``..._d512`` and friends) into a single row once, and the substring test survived that. + ``exclude_plan_tokens`` is the opposite instruction, for a caller that must NOT run on a + particular engine. It is needed because the FROST engine switch is process-wide: a caller that + declines to ask for those engines still gets them ranked first once anything else in the + process has enabled them. Measured on B200 at head_dim 64, 128, 256 and 512, a FROST plan + ranks at index 0 for a score_mod graph and an unpinned build selects it every time. + Pinning also changes what ``check_support`` means. Selecting a plan sets cuDNN's internal ``_plan_pinned``, and only then is a decline fatal; unpinned, cuDNN records the decline and keeps walking. So the pin has to come first both because the check is scoped to the selected @@ -212,6 +219,11 @@ def finalize_plans( if require_plan_token is None: try: graph.create_execution_plans(list(heuristics)) + if exclude_plan_tokens: + # Bar the named engines before the walk, so build_plans falls through to the first + # entry that is both unbarred and buildable. Inert when those engines are not on + # offer, which is every process that has not enabled them. + graph.deselect_engines(list(exclude_plan_tokens)) graph.check_support() except cudnn.cudnnGraphNotSupportedError as exc: raise RuntimeError(f"cuDNN SDPA graph is not supported: {exc}") from exc diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index 5a15c234a5c..82305550795 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -247,9 +247,20 @@ class _CudnnScoreModBwdGraphEntry: workspace_size: int +# cuDNN FROST SDPA engine names. These are barred here, not merely left unasked for: the switch +# that offers them is process-wide, so any FrostAttention call elsewhere in the process, or a user +# setting CUDNN_FRONTEND_ENABLE_FROST_ENGINES, puts them ahead of the backend engines for these +# graphs too. They accept a score_mod graph, pass check_support, build, and then compute without +# the callback. Measured on B200 with cuDNN Frontend 1.29.0: a FROST plan ranks at index 0 at +# head_dim 64, 128, 256 and 512, and an unpinned build selects it and returns plain attention. +_FROST_PLAN_TOKENS = ("sdpa_fwd_prefill_sm100", "sdpa_bwd_sm100") + + def _finalize_cudnn_graph(graph) -> int: """Build a cuDNN frontend Python graph and return its workspace size.""" - workspace_size, _ = cudnn_pygraph.finalize_plans(graph) + workspace_size, _ = cudnn_pygraph.finalize_plans( + graph, exclude_plan_tokens=_FROST_PLAN_TOKENS + ) return workspace_size diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index d8c235c200a..b300be40a78 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -77,17 +77,18 @@ _HANDLES = cudnn_pygraph._handles # pylint: disable=protected-access -def _import_cudnn(): - """Import cuDNN Frontend with the FROST engines registered. +def _import_cudnn(enable_frost_engines: bool = True): + """Import cuDNN Frontend, registering the FROST engines unless told not to. - The switch has to be set before the import because the engines register at import time, and it - also ranks FROST ahead of the backend engines, so only this backend asks for it. - _select_frost_plan verifies the engine by plan name regardless, rather than trusting the flag. + The switch is process-wide and ranks FROST ahead of the backend engines for every cuDNN Python + graph afterwards, including other backends’ graphs, so it is set only where FROST is actually + used. _select_frost_plan verifies the engine by plan name regardless, rather than trusting the + flag. """ global _cudnn # pylint: disable=global-statement # Kept bound: _pkg_version falls back to the module's __version__ when distribution metadata # is unavailable, which is how a source or vendored install avoids being misreported. - _cudnn = cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=True) + _cudnn = cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=enable_frost_engines) return _cudnn @@ -151,7 +152,10 @@ def _no(reason): major, minor = torch.cuda.get_device_capability() return _no(f"cuDNN FROST head_dim>256 kernels are SM100/SM103 only; found sm{major}{minor}") try: - _import_cudnn() + # Without the engines: this only needs the module to read a version off it, and enabling + # here would reorder plan selection for the whole process even when the checks below go on + # to decline FROST, which is all cost and no benefit. The use sites enable it. + _import_cudnn(enable_frost_engines=False) except ImportError as exc: return _no(f"nvidia-cudnn-frontend not importable: {exc}") From c46dffc77baf3f49f9e6fee57b396f62241c7693 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 1 Oct 2026 10:02:23 +0000 Subject: [PATCH 45/97] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/pytorch/attention/test_flex_attention.py | 8 +++----- .../attention/dot_product_attention/flex_attention.py | 4 +--- 2 files changed, 4 insertions(+), 8 deletions(-) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index 52114e0a0e7..466d8d902c6 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -857,9 +857,7 @@ def bias_score_mod(score_mod_graph, score_tensor, _tensors): col.set_data_type(cudnn.data_type.INT32) bias = score_mod_graph.sub(a=row, b=col, compute_data_type=cudnn.data_type.FLOAT) bias.set_data_type(cudnn.data_type.FLOAT) - return score_mod_graph.add( - a=score_tensor, b=bias, compute_data_type=cudnn.data_type.FLOAT - ) + return score_mod_graph.add(a=score_tensor, b=bias, compute_data_type=cudnn.data_type.FLOAT) def run(): flex_attention._cudnn_score_mod_graph_cache.clear() @@ -882,6 +880,6 @@ def run(): torch.testing.assert_close( with_engines, without, - msg=lambda m: "flex computed something different with the FROST engines enabled, which" - " means a FROST plan answered and dropped the score_mod:\n" + m, + msg=lambda m: "flex computed something different with the FROST engines enabled, which means a FROST plan answered and dropped the score_mod:\n" + + m, ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index 82305550795..6df8655f0a1 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -258,9 +258,7 @@ class _CudnnScoreModBwdGraphEntry: def _finalize_cudnn_graph(graph) -> int: """Build a cuDNN frontend Python graph and return its workspace size.""" - workspace_size, _ = cudnn_pygraph.finalize_plans( - graph, exclude_plan_tokens=_FROST_PLAN_TOKENS - ) + workspace_size, _ = cudnn_pygraph.finalize_plans(graph, exclude_plan_tokens=_FROST_PLAN_TOKENS) return workspace_size From 2cfd6ff45f8629490604b773ab5c6d37c6b8117f Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 1 Oct 2026 03:10:43 -0700 Subject: [PATCH 46/97] test(attention): skip the FROST switch test where it cannot detect anything Guarded only on CUDA, the test passed on every machine where the FROST engines are absent or decline on arch: both runs get a backend plan and agree regardless of what flex does. It now skips unless the engines are actually reachable, so a pass means something. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- tests/pytorch/attention/test_flex_attention.py | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index 466d8d902c6..61fbab80d47 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -841,6 +841,19 @@ def test_frost_switch_does_not_change_what_flex_computes(): except ImportError: pytest.skip("cuDNN frontend Python package is required for score_mod attention.") + # Without this the test is vacuous nearly everywhere: where the FROST engines are absent or + # decline on arch, both runs get a backend plan and agree no matter what flex does. The + # engines themselves are found lazily at planning time, so the switch works whenever it is + # set, but they still have to exist. + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + is_frost_attention_available, + ) + + frost_ok, frost_reason = is_frost_attention_available() + if not frost_ok: + pytest.skip("the FROST engines must be reachable for this to test anything: %s" + % frost_reason) + env = "CUDNN_FRONTEND_ENABLE_FROST_ENGINES" saved = os.environ.get(env) torch.manual_seed(0) From 636d7a651d2adef57349ece488e9a2e8973a0422 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 1 Oct 2026 10:13:20 +0000 Subject: [PATCH 47/97] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/pytorch/attention/test_flex_attention.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index 61fbab80d47..42236812b28 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -851,8 +851,9 @@ def test_frost_switch_does_not_change_what_flex_computes(): frost_ok, frost_reason = is_frost_attention_available() if not frost_ok: - pytest.skip("the FROST engines must be reachable for this to test anything: %s" - % frost_reason) + pytest.skip( + "the FROST engines must be reachable for this to test anything: %s" % frost_reason + ) env = "CUDNN_FRONTEND_ENABLE_FROST_ENGINES" saved = os.environ.get(env) From 69a253b530e592ac0554d932adcdccb2267b3861 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Sun, 4 Oct 2026 03:05:21 -0700 Subject: [PATCH 48/97] fix(attention): index the FROST p2p results by the alternating slot Merging main brought #2916, which shrank the P2P forward's per-step result lists from cp_size to two alternating slots and converted the fused and flash branches to [i % 2]. The FROST branch was added separately and kept [i], so from the third ring step it indexed past the end: p2p and a2a+p2p would raise IndexError at CP=4. out_per_step, softmax_lse_per_step and max_logit_per_step are the two-slot lists; rng_states and attn_biases are still cp_size long and keep [i]. The all_gather path is unchanged, where the fused branch also uses [i]. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/context_parallel.py | 24 +++++++++---------- 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index b2a2b56bf6f..a547532d5e7 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -2291,11 +2291,11 @@ def forward( q_inputs[i % 2] = q_part if use_frost_attention: ( - out_per_step[i], - softmax_lse_per_step[i], + out_per_step[i % 2], + softmax_lse_per_step[i % 2], rng_states[i], attn_biases[i], - max_logit_per_step[i], + max_logit_per_step[i % 2], ) = cp_p2p_fwd_frost_attn( *frost_attn_inputs, *prepare_outputs, section ) @@ -2330,11 +2330,11 @@ def forward( q_inputs[i % 2] = q_part if use_frost_attention: ( - out_per_step[i], - softmax_lse_per_step[i], + out_per_step[i % 2], + softmax_lse_per_step[i % 2], rng_states[i], attn_biases[i], - max_logit_per_step[i], + max_logit_per_step[i % 2], ) = cp_p2p_fwd_frost_attn( *frost_attn_inputs, *prepare_outputs, section ) @@ -2369,11 +2369,11 @@ def forward( q_inputs[i % 2] = q_part if use_frost_attention: ( - out_per_step[i], - softmax_lse_per_step[i], + out_per_step[i % 2], + softmax_lse_per_step[i % 2], rng_states[i], attn_biases[i], - max_logit_per_step[i], + max_logit_per_step[i % 2], ) = cp_p2p_fwd_frost_attn( *frost_attn_inputs, *prepare_outputs, section ) @@ -2409,11 +2409,11 @@ def forward( q_inputs[i % 2] = q_part if use_frost_attention: ( - out_per_step[i], - softmax_lse_per_step[i], + out_per_step[i % 2], + softmax_lse_per_step[i % 2], rng_states[i], attn_biases[i], - max_logit_per_step[i], + max_logit_per_step[i % 2], ) = cp_p2p_fwd_frost_attn(*frost_attn_inputs, *prepare_outputs, section) elif use_fused_attention: ( From 7a558193c7bbdb785dd758066df72a7960484150 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Mon, 5 Oct 2026 17:47:18 -0700 Subject: [PATCH 49/97] refactor(attention): make FROST a FusedAttention sub-backend FROST was a fourth top-level backend, which meant plumbing a use_frost_attention flag through dot_product_attention.py, backends.py and context_parallel.py and duplicating the per-ring-step mask and layout handling the fused path already has. It is now FusedAttnBackend["FROST"], selected inside _get_fused_attn_backend when the C++ sub-backends decline, and dispatched inside cpp_extensions.fused_attn.fused_attn_fwd/bwd. frost_attention.py gained fused_attn_fwd/bwd behind those signatures, so FusedAttnFunc and the context-parallel ring reach the kernels without knowing which sub-backend they got. dot_product_attention.py needs no change at all. backends.py and context_parallel.py each keep the selected sub-backend instead of re-deriving F16_arbitrary_seqlen, which is the only reason they change: nine sites discarded it. Several FROST declines now come from the existing fused filters rather than their own copies -- the context-parallel mask and window restrictions, the score_mod filter, and fp8 -- and the all_gather path's bottom-right mask rewrite applies to FROST for free. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- docs/envvars.rst | 9 +- .../attention/run_attention_with_cp.py | 17 +- .../pytorch/attention/test_frost_attention.py | 127 +++-- .../attention/test_mixed_thd_attention.py | 2 +- tests/pytorch/test_torch_compile.py | 1 - tests/pytorch/utils.py | 2 - .../dot_product_attention/backends.py | 216 +------- .../dot_product_attention/context_parallel.py | 481 ++---------------- .../dot_product_attention.py | 55 +- .../dot_product_attention/frost_attention.py | 323 ++++++++++-- .../attention/dot_product_attention/utils.py | 210 +------- .../pytorch/cpp_extensions/fused_attn.py | 93 +++- 12 files changed, 560 insertions(+), 976 deletions(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index 9a7933f6d1e..1df2ed5bcca 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -178,10 +178,9 @@ Then it applies a performance-based preference order among the remaining eligibl In PyTorch, the broad preference order is ``FlashAttention > FusedAttention > UnfusedDotProductAttention`` on supported pre-Hopper GPUs such as Ampere/Ada, and ``FusedAttention > FlashAttention > UnfusedDotProductAttention`` on Hopper and newer GPUs, -including Blackwell. On Blackwell SM100/SM103 the order is ``FusedAttention > FlashAttention > -FrostAttention > UnfusedDotProductAttention``; FrostAttention only becomes eligible for -symmetric ``head_dim`` in (256, 512], which flash and fused attention do not serve, so the -backend it can displace is UnfusedDotProductAttention. In JAX, Transformer Engine uses cuDNN +including Blackwell. On Blackwell SM100/SM103, FusedAttention has an extra sub-backend, FROST, +which is selected only for symmetric ``head_dim`` in (256, 512] and only when the cuDNN +sub-backends decline; it does not change the order above. In JAX, Transformer Engine uses cuDNN fused attention when ``NVTE_FUSED_ATTN=1`` and an eligible cuDNN kernel is available; otherwise it falls back to the JAX-native implementation. See :doc:`examples/attention/attention` for a longer backend-selection overview. @@ -220,7 +219,7 @@ longer backend-selection overview. :Type: ``int`` (0 or 1) :Default: ``1`` - :Description: Enable or disable FrostAttention backend (the cuDNN FROST CuTe-DSL SDPA kernels in cuDNN Frontend) for DotProductAttention. **This backend is experimental and subject to change**, including the possibility of being folded into FusedAttention; the underlying cuDNN FROST engines are themselves experimental. When set to ``0``, FrostAttention will not be used. From released components it is the only backend serving symmetric ``head_dim`` in (256, 512] together with context parallelism; without context parallelism UnfusedDotProductAttention also covers that range, and FrostAttention is preferred over it where both are eligible. 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 supports context parallelism with ``cp_comm_type`` of ``p2p``, ``all_gather``, ``a2a`` or ``a2a+p2p``, and sliding-window attention with ``all_gather`` or ``a2a`` (declined with ``p2p`` and ``a2a+p2p``, whose ring shards KV across steps). It declines FP8, ``thd`` layouts, dropout, attention bias, softcap, KV caching, ``max_logit``, and deterministic execution, the last because cuDNN offers no deterministic backward for these kernels. + :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. .. envvar:: NVTE_UNFUSED_ATTN diff --git a/tests/pytorch/attention/run_attention_with_cp.py b/tests/pytorch/attention/run_attention_with_cp.py index 176e813ea53..342084474c3 100644 --- a/tests/pytorch/attention/run_attention_with_cp.py +++ b/tests/pytorch/attention/run_attention_with_cp.py @@ -278,8 +278,9 @@ def run_dpa_with_cp( else: assert False, f"{model=} is not a known FusedAttention CP config!" if kernel_backend == "FrostAttention": - # Leave NVTE_FLASH_ATTN and NVTE_FUSED_ATTN at 0: FROST is the only backend that serves - # head_dim > 256, so get_attention_backend selects it on its own. + # FROST is a sub-backend of FusedAttention, so NVTE_FUSED_ATTN has to stay on. Flash is + # left off; nothing else serves head_dim > 256, so the selector reaches FROST on its own. + os.environ["NVTE_FUSED_ATTN"] = "1" os.environ["NVTE_FROST_ATTN"] = "1" if model in model_configs_frost_attn: config = copy.deepcopy(model_configs_frost_attn[model]) @@ -606,18 +607,20 @@ def run_dpa_with_cp( fp8_output=fp8_mha, ) if kernel_backend == "FrostAttention": - # Assert the backend actually used, not just the one requested. FROST is currently - # the only selectable backend for these configs -- flash and fused are env-gated off + # Assert the sub-backend actually used, not just the one requested. FROST is + # currently the only selectable backend for these configs -- flash is env-gated off # and CP disables unfused -- so a silent substitution is impossible today and this # would pass by construction. It is here so it stops passing if that stops being # true, rather than quietly testing some other kernel. from transformer_engine.pytorch.attention.dot_product_attention.dot_product_attention import ( # pylint: disable=import-outside-toplevel _attention_backends, ) + # pylint: disable-next=import-outside-toplevel + from transformer_engine.pytorch.cpp_extensions.fused_attn import FusedAttnBackend - assert _attention_backends[ - "use_frost_attention" - ], "expected FrostAttention to be selected, got %s" % (_attention_backends,) + assert ( + _attention_backends["fused_attention_backend"] == FusedAttnBackend.FROST + ), "expected the FROST sub-backend to be selected, got %s" % (_attention_backends,) if config.return_max_logit: out_, max_logit_ = out_ if is_training: diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 9492267e78a..45c301f703d 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -263,41 +263,99 @@ def test_frost_backward_matches_reference(shape, mask, window, dtype): ) +def _frost_params(**overrides): + """A FusedAttentionParams for a config FROST serves, with fields overridable by name.""" + from transformer_engine.pytorch.attention.dot_product_attention.utils import ( + FusedAttentionParams, + ) + from transformer_engine.pytorch.cpp_extensions.fused_attn import ( + AttnBiasType, + AttnMaskType, + QKVFormat, + QKVLayout, + SoftmaxType, + ) + from transformer_engine.pytorch.constants import TE_DType + + fields = dict( + head_dim_qk=512, + head_dim_v=512, + qkv_dtype=TE_DType[torch.bfloat16], + attn_mask_type=AttnMaskType["causal"], + bias_type=AttnBiasType["no_bias"], + softmax_type=SoftmaxType["vanilla"], + qkv_layout=QKVLayout["bshd_bshd_bshd"], + o_format=QKVFormat["bshd"], + window_size_left=-1, + window_size_right=-1, + bottom_right_diagonal=False, + ) + fields.update(overrides) + return FusedAttentionParams(**fields) + + @requires_frost def test_frost_declines_unsupported_configs(): """The selector must decline what the kernels do not serve, rather than computing wrongly.""" from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( is_frost_attention_supported, ) + from transformer_engine.pytorch.cpp_extensions.fused_attn import ( + AttnBiasType, + AttnMaskType, + FusedAttnBackend, + QKVFormat, + QKVLayout, + ) + from transformer_engine.pytorch.constants import TE_DType - base = dict(head_dim_qk=512, head_dim_v=512, qkv_dtype=torch.bfloat16, attn_mask_type="causal") - assert is_frost_attention_supported(**base)[0], "the supported case must be accepted" + assert ( + is_frost_attention_supported(_frost_params())[0] == FusedAttnBackend.FROST + ), "the supported case must be accepted" for override, why in ( (dict(head_dim_qk=256, head_dim_v=256), "head_dim at the exclusive lower bound"), (dict(head_dim_v=256), "asymmetric head_dim"), - (dict(qkv_dtype=torch.float32), "fp32"), + (dict(qkv_dtype=TE_DType[torch.float32]), "fp32"), (dict(dropout=0.1), "dropout"), - (dict(attn_bias_type="post_scale_bias"), "attention bias"), - (dict(attn_mask_type="padding_causal"), "padding mask"), - (dict(attn_mask_type="arbitrary"), "arbitrary mask"), - # window_size reaches _mask_spec through is_frost_attention_supported, so its validation - # is part of the selector contract rather than an internal detail. - (dict(window_size=(-1, 5)), "a right window past the diagonal"), - (dict(window_size=(128,)), "a malformed window pair"), - (dict(window_size=(-2, 0)), "a left window below -1"), - (dict(window_size=7), "a non-iterable window"), + (dict(bias_type=AttnBiasType["post_scale_bias"]), "attention bias"), + (dict(attn_mask_type=AttnMaskType["padding_causal"]), "padding mask"), + (dict(qkv_layout=QKVLayout["thd_thd_thd"], o_format=QKVFormat["thd"]), "thd layout"), + (dict(o_format=QKVFormat["sbhd"]), "an output format that differs from the input"), + (dict(num_pages_k=4, num_pages_v=4), "paged KV"), + (dict(return_max_logit=True), "max_logit"), + (dict(cuda_graph=True), "CUDA graph capture"), + (dict(deterministic=True, is_training=True), "a deterministic backward"), + # window_size reaches _mask_spec through the selector, so its validation is part of the + # selector contract rather than an internal detail. + (dict(window_size_right=5), "a right window past the diagonal"), + (dict(window_size_left=-2, window_size_right=0), "a left window below -1"), # The engine pads head_dim to a multiple of 8, so an in-range but unpadded dim has to be # declined here rather than failing later at plan selection. (dict(head_dim_qk=260, head_dim_v=260), "head_dim not a multiple of 8"), ): - cfg = dict(base) - cfg.update(override) - ok, reason = is_frost_attention_supported(**cfg) - assert not ok, "%s must be declined" % why + backend, reason = is_frost_attention_supported(_frost_params(**override)) + assert backend == FusedAttnBackend.No_Backend, "%s must be declined" % why assert reason, "a decline must explain itself" +@requires_frost +def test_frost_mask_spec_rejects_malformed_windows(): + """_mask_spec is the only validation between a caller-supplied window and a built band.""" + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + _mask_spec, + ) + + for window, why in ( + ((128,), "a malformed window pair"), + (7, "a non-iterable window"), + ((-1, 5), "a right window past the diagonal"), + ((-2, 0), "a left window below -1"), + ): + with pytest.raises(NotImplementedError): + _mask_spec("causal", window), why + + @requires_frost @pytest.mark.parametrize( "cp_comm_type,window,expect_frost", @@ -339,14 +397,17 @@ def test_frost_sliding_window_selection_by_cp_comm_type(cp_comm_type, window, ex cp_comm_type=cp_comm_type, is_training=True, ) - use_frost = get_attention_backend(params)[5] + from transformer_engine.pytorch.cpp_extensions.fused_attn import FusedAttnBackend + + use_fused, fused_backend = get_attention_backend(params)[2:4] + use_frost = bool(use_fused) and fused_backend == FusedAttnBackend.FROST assert ( - bool(use_frost) == expect_frost - ), "cp_comm_type=%s window=%s: expected use_frost_attention=%s, got %s" % ( + use_frost == expect_frost + ), "cp_comm_type=%s window=%s: expected the FROST sub-backend=%s, got %s" % ( cp_comm_type, window, expect_frost, - bool(use_frost), + use_frost, ) @@ -376,32 +437,6 @@ def test_frost_rejects_mismatched_kv(): frost_attn_fwd(q, k, k.to(torch.float32)) -@pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a CUDA device") -def test_dot_product_attention_runs_in_onnx_export_mode(): - """The ONNX-export branch must bind every backend flag the availability check reads. - - Deliberately not gated on FROST: that branch skips get_attention_backend entirely and sets the - flags by hand, so leaving use_frost_attention unbound there raised UnboundLocalError for every - user on every GPU, whether or not FROST could run. A plain head_dim-64 config reproduces it -- - the failure is in the selector bookkeeping, not in any kernel. - """ - from transformer_engine.pytorch import DotProductAttention - from transformer_engine.pytorch.export import onnx_export - - b, h, s, d = 2, 4, 128, 64 - dtype = torch.bfloat16 - qkv = [torch.randn(s, b, h, d, device="cuda", dtype=dtype) for _ in range(3)] - block = DotProductAttention( - h, d, qkv_format="sbhd", attn_mask_type="causal", attention_dropout=0.0 - ).to(dtype=dtype, device="cuda") - - with onnx_export(enabled=True): - out = block(*qkv) - - assert out.numel() == s * b * h * d - assert torch.isfinite(out).all() - - def test_frost_engines_are_enabled_even_if_flex_imported_cudnn_first(): """Enabling the FROST engines must not depend on which backend touched cuDNN first. diff --git a/tests/pytorch/attention/test_mixed_thd_attention.py b/tests/pytorch/attention/test_mixed_thd_attention.py index d4618db126d..d665df4ceff 100644 --- a/tests/pytorch/attention/test_mixed_thd_attention.py +++ b/tests/pytorch/attention/test_mixed_thd_attention.py @@ -453,7 +453,7 @@ def test_thd_mask_type_runtime_dispatch_uses_backend_selection(monkeypatch): def fake_get_attention_backend(attention_params): observed_params.append(attention_params) available_backends = [False, attention_params.attn_mask_type == "padding", False] - return False, None, available_backends[1], None, False, False, available_backends + return False, None, available_backends[1], None, False, available_backends monkeypatch.setattr(dpa_module.dpa_utils, "get_attention_backend", fake_get_attention_backend) padded_policies, grouped_policies = DotProductAttention._partition_thd_mask_policies( diff --git a/tests/pytorch/test_torch_compile.py b/tests/pytorch/test_torch_compile.py index fd4412e26d8..eae6f0a8a24 100644 --- a/tests/pytorch/test_torch_compile.py +++ b/tests/pytorch/test_torch_compile.py @@ -1294,7 +1294,6 @@ def fn(x, params): fused_attention_backend, use_unfused_attention, _, - _, ) = dpa_utils.get_attention_backend(params) # Encode the full selection (enabled backends + fused sub-backend) in # the tensor value: without a tensor op dynamo skips the frame entirely diff --git a/tests/pytorch/utils.py b/tests/pytorch/utils.py index 6a8d50bf19a..62917a5c8d0 100644 --- a/tests/pytorch/utils.py +++ b/tests/pytorch/utils.py @@ -452,7 +452,6 @@ def test(): use_fused_attention, fused_attention_backend, use_unfused_attention, - _use_frost_attention, available_backends, ) = get_attention_backend(attention_params) # Check if FA3 is an available backend when num_splits != 1 @@ -466,7 +465,6 @@ def test(): _attention_backends["flash_attention_backend"] = flash_attention_backend _attention_backends["fused_attention_backend"] = fused_attention_backend _attention_backends["use_unfused_attention"] = use_unfused_attention - _attention_backends["use_frost_attention"] = _use_frost_attention _attention_backends["backend_selection_requires_update"] = False return available_backends, flash_attention_backend, fused_attention_backend diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index 49845211ec7..3a81b5142bc 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -2036,8 +2036,13 @@ def _fused_attn_setup_ctx( bwd_args.softmax_type = fwd_args.softmax_type bwd_args.window_size = fwd_args.window_size bwd_args.bottom_right_diagonal = fwd_args.bottom_right_diagonal + # FROST has to survive this: it is a python sub-backend, so re-deriving F16_arbitrary_seqlen + # here would send its backward to the C++ path, which does not serve these head dims. + saved_fused_attention_backend = ctx_attrs["fused_attention_backend"] bwd_args.fused_attention_backend = ( - ctx_attrs["fused_attention_backend"] if fp8 else FusedAttnBackend["F16_arbitrary_seqlen"] + saved_fused_attention_backend + if fp8 or saved_fused_attention_backend == FusedAttnBackend["FROST"] + else FusedAttnBackend["F16_arbitrary_seqlen"] ) bwd_args.use_FAv2_bwd = fwd_args.use_FAv2_bwd bwd_args.deterministic = fwd_args.deterministic @@ -2369,210 +2374,6 @@ def backward(ctx, d_out, *_args): return (*_fused_attn_backward_impl(bwd_args), None) -class FrostAttnFunc(torch.autograd.Function): - """Autograd wrapper around the cuDNN FROST kernels, for the non-context-parallel path. - - The CP path does not go through here: context_parallel.py calls frost_attn_fwd/bwd per ring - step itself, because the ring has to interleave those calls with KV exchange and LSE - correction rather than treating attention as one opaque autograd node. - """ - - @staticmethod - def forward( - ctx, - q, - k, - v, - softmax_scale, - attn_mask_type, - qkv_format, - is_training, - deterministic, - window_size, - ): - # pylint: disable=missing-function-docstring - from .frost_attention import ( # pylint: disable=import-outside-toplevel - frost_attn_fwd, - from_frost_layout, - to_frost_layout, - ) - - # .contiguous() first: the graphs are built from each tensor's actual strides, so an - # arbitrary incoming layout would key a separate plan per layout and require k and v to - # agree. Normalising here keeps one plan per shape. - q_f = to_frost_layout(q.contiguous(), qkv_format) - k_f = to_frost_layout(k.contiguous(), qkv_format) - v_f = to_frost_layout(v.contiguous(), qkv_format) - out_f, softmax_lse = frost_attn_fwd( - q_f, - k_f, - v_f, - attn_scale=softmax_scale, - attn_mask_type=attn_mask_type, - window_size=window_size, - ) - out = from_frost_layout(out_f, qkv_format) - if is_training: - ctx.save_for_backward(q_f, k_f, v_f, out_f, softmax_lse) - ctx.softmax_scale = softmax_scale - ctx.attn_mask_type = attn_mask_type - ctx.window_size = window_size - ctx.qkv_format = qkv_format - ctx.unflattened_shape = out.shape - ctx.deterministic = deterministic - # TE attention modules return the heads flattened into the last dimension - # ([b, s, h*d] for bshd), matching FlashAttention and FusedAttention. Returning the - # unflattened [b, s, h, d] makes autograd reject the incoming grad on shape mismatch. - return out.reshape(out.shape[0], out.shape[1], -1) - - @staticmethod - def backward(ctx, dout): - # pylint: disable=missing-function-docstring - from .frost_attention import ( # pylint: disable=import-outside-toplevel - frost_attn_bwd, - from_frost_layout, - to_frost_layout, - ) - - q_f, k_f, v_f, out_f, softmax_lse = ctx.saved_tensors - fmt = ctx.qkv_format - # dout arrives flattened, matching what forward returned; restore [b, s, h, d]. - dout = dout.reshape(ctx.unflattened_shape) - dq, dk, dv = frost_attn_bwd( - q_f, - k_f, - v_f, - out_f, - softmax_lse, - to_frost_layout(dout.contiguous(), fmt), - attn_scale=ctx.softmax_scale, - attn_mask_type=ctx.attn_mask_type, - window_size=ctx.window_size, - deterministic=ctx.deterministic, - ) - # One None per non-tensor forward argument: softmax_scale, attn_mask_type, qkv_format, - # is_training, deterministic, window_size. Must track forward's signature exactly. - return ( - from_frost_layout(dq, fmt), - from_frost_layout(dk, fmt), - from_frost_layout(dv, fmt), - None, - None, - None, - None, - None, - None, - ) - - -class FrostAttention(torch.nn.Module): - """cuDNN FROST attention for symmetric head_dim in (256, 512] on SM100/SM103. - - **Experimental and subject to change**, including the possibility of being folded into - FusedAttention: the underlying cuDNN FROST engines are themselves experimental. - - This is the only backend that serves that head-dim range together with context parallelism. - Deliberately narrow: no FP8, no bias, no dropout, no softmax offset, no paging. - get_attention_backend declines all of those before selecting this backend, so anything - reaching here should already be supported. - """ - - def __init__( - self, - softmax_scale: float, - attention_type: str = "self", - layer_number: Optional[int] = None, - deterministic: bool = False, - **kwargs, # attention_dropout / attention_dropout_ctx: accepted, must be unused - ) -> None: - super().__init__() - self.softmax_scale = softmax_scale - self.attention_type = attention_type - self.layer_number = 1 if layer_number is None else layer_number - self.deterministic = deterministic - self.attention_dropout = kwargs.get("attention_dropout", 0.0) - - def forward( - self, - query_layer: torch.Tensor, - key_layer: torch.Tensor, - value_layer: torch.Tensor, - qkv_format: str = "bshd", - cu_seqlens_q: Optional[torch.Tensor] = None, - cu_seqlens_kv: Optional[torch.Tensor] = None, - max_seqlen_q: Optional[int] = None, - max_seqlen_kv: Optional[int] = None, - cu_seqlens_q_padded: Optional[torch.Tensor] = None, - cu_seqlens_kv_padded: Optional[torch.Tensor] = None, - attn_mask_type: str = "causal", - window_size: Optional[Tuple[int, int]] = None, - cp_group: Optional[Union[dist_group_type, List[dist_group_type]]] = None, - cp_global_ranks: List[int] = None, - cp_stream: torch.cuda.Stream = None, - cp_comm_type: str = "p2p", - load_balancing_strategy: CPLoadBalancingStrategy = ( - CPLoadBalancingStrategy.DUAL_CHUNK_SWAP - ), - ) -> torch.Tensor: - """Forward pass. Routes through the CP ring when a cp_group is present.""" - assert self.attention_dropout == 0.0, "FrostAttention does not support dropout" - - # Same form as FlashAttention and FusedAttention above. cp_group is a list of two groups - # for cp_comm_type="a2a+p2p", and passing that list to get_distributed_world_size raises - # TypeError: unhashable type: 'list'. - cp_size = 1 - if isinstance(cp_group, dist_group_type): - cp_size = get_distributed_world_size(cp_group) - elif isinstance(cp_group, list): - for group in cp_group: - cp_size *= get_distributed_world_size(group) - context_parallel = cp_size > 1 - if context_parallel: - output = attn_forward_func_with_cp( - self.training, - query_layer, - key_layer, - value_layer, - cu_seqlens_q, - cu_seqlens_kv, - max_seqlen_q, - max_seqlen_kv, - cu_seqlens_q_padded, - cu_seqlens_kv_padded, - 0.0, - cp_group, - cp_global_ranks, - cp_stream, - cp_comm_type, - softmax_scale=self.softmax_scale, - qkv_format=qkv_format, - attn_mask_type=attn_mask_type, - attn_bias_type="no_bias", - attn_bias=None, - deterministic=self.deterministic, - use_fused_attention=False, - use_frost_attention=True, - window_size=window_size, - layer_number=self.layer_number, - load_balancing_strategy=load_balancing_strategy, - ) - # Same flattening the other backends apply after the CP call: the ring returns - # [b, s_local, h, d] but TE attention modules return heads in the last dimension. - return output.reshape(output.shape[0], output.shape[1], -1).contiguous() - - return FrostAttnFunc.apply( - query_layer, - key_layer, - value_layer, - self.softmax_scale, - attn_mask_type, - qkv_format, - self.training, - self.deterministic, - window_size, - ) - - class FusedAttention(torch.nn.Module): """Dot product attention using `cuDNN attention `_: @@ -2809,7 +2610,9 @@ def forward( if context_parallel: assert ( - fp8 or fused_attention_backend == FusedAttnBackend["F16_arbitrary_seqlen"] + fp8 + or fused_attention_backend + in (FusedAttnBackend["F16_arbitrary_seqlen"], FusedAttnBackend["FROST"]) ), f"{fused_attention_backend} does not work with context parallelism!" assert core_attention_bias_type not in [ "alibi" @@ -2841,6 +2644,7 @@ def forward( attn_bias=core_attention_bias, deterministic=self.deterministic, use_fused_attention=True, + fused_attention_backend=fused_attention_backend, window_size=window_size, fp8=fp8, fp8_meta=fp8_meta, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index a547532d5e7..a98c5538bb9 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1572,259 +1572,6 @@ def cp_p2p_bwd_flash_attn( return dq, dk, dv -def _frost_mask_for_section(attn_mask_type, section): - """Per-ring-step mask, mirroring cp_p2p_fwd_fused_attn. - - Only the diagonal tile keeps the causal mask; the off-diagonal tiles see a fully visible KV - block. This matches what was validated on B200: causal on the square diagonal, no_mask on the - rectangular off-diagonal tiles. - """ - if section in ("diagonal", "all"): - return attn_mask_type - if section in ("lower-triangle", "upper-triangle"): - return "no_mask" - raise ValueError(f"unknown CP section {section!r}") - - -def _frost_mask_for_window(window_size): - """Per-step mask for the all_gather path, derived from its adjusted window. - - get_kv_seq_info_after_all_gather trims KV and returns a window that is BOTTOM-RIGHT aligned: - (-1, 0) means causal relative to the trimmed KV, not top-left causal. Using top-left would be - wrong wherever the two differ, which is whenever the trim leaves SKV > SQ. - """ - if window_size is None or tuple(window_size) == (-1, -1): - return "no_mask", None - # A positive right bound is look-ahead, which none of the supported masks express. _mask_spec - # rejects it at selection time, but that is a different file, so assert the invariant here - # rather than quietly returning a causal mask that admits future keys. - assert window_size[1] in ( - -1, - 0, - ), f"all_gather produced a look-ahead window {window_size}" - # Anything with a bounded side is causal relative to the trimmed KV, and a bounded left side - # is a sliding window. Both are expressed as a band against the bottom-right diagonal, so the - # window travels with the mask type rather than needing a separate spelling per case. - return "causal_bottom_right", tuple(window_size) - - -def cp_ag_fwd_frost_attn( - softmax_scale, - qkv_format, - window_size, - q_part, - k_part, - v_part, -): - """Per-step forward for CP all_gather with the cuDNN FROST backend. - - Simpler than the p2p ring: KV is already gathered and trimmed, so each step is a single - attention call with no LSE correction. Returns (out, softmax_lse). - """ - from .frost_attention import ( # pylint: disable=import-outside-toplevel - frost_attn_fwd, - from_frost_layout, - to_frost_layout, - ) - - mask_type, window = _frost_mask_for_window(window_size) - out, softmax_lse = frost_attn_fwd( - to_frost_layout(q_part.contiguous(), qkv_format), - to_frost_layout(k_part.contiguous(), qkv_format), - to_frost_layout(v_part.contiguous(), qkv_format), - attn_scale=softmax_scale, - attn_mask_type=mask_type, - window_size=window, - ) - return from_frost_layout(out, qkv_format), softmax_lse - - -def cp_ag_bwd_frost_attn( - softmax_scale, - qkv_format, - window_size, - softmax_lse, - q_part, - k_part, - v_part, - out_part, - dout_part, - deterministic=False, -): - """Per-step backward for CP all_gather with the cuDNN FROST backend.""" - from .frost_attention import ( # pylint: disable=import-outside-toplevel - frost_attn_bwd, - from_frost_layout, - to_frost_layout, - ) - - mask_type, window = _frost_mask_for_window(window_size) - dq, dk, dv = frost_attn_bwd( - to_frost_layout(q_part.contiguous(), qkv_format), - to_frost_layout(k_part.contiguous(), qkv_format), - to_frost_layout(v_part.contiguous(), qkv_format), - to_frost_layout(out_part.contiguous(), qkv_format), - softmax_lse, - to_frost_layout(dout_part.contiguous(), qkv_format), - attn_scale=softmax_scale, - attn_mask_type=mask_type, - window_size=window, - deterministic=deterministic, - ) - return ( - from_frost_layout(dq, qkv_format), - from_frost_layout(dk, qkv_format), - from_frost_layout(dv, qkv_format), - ) - - -def cp_a2a_fwd_frost_attn(softmax_scale, attn_mask_type, qkv_format, q, k, v, window_size=None): - """Forward for CP a2a with the cuDNN FROST backend. - - The simplest of the three. After the all-to-all each rank holds the FULL sequence for a subset - of heads, so there is no ring, no KV trimming and no LSE correction: one ordinary attention - call with the caller mask type, top-left causal as usual. - """ - from .frost_attention import ( # pylint: disable=import-outside-toplevel - frost_attn_fwd, - from_frost_layout, - to_frost_layout, - ) - - out, softmax_lse = frost_attn_fwd( - to_frost_layout(q.contiguous(), qkv_format), - to_frost_layout(k.contiguous(), qkv_format), - to_frost_layout(v.contiguous(), qkv_format), - attn_scale=softmax_scale, - attn_mask_type=attn_mask_type, - window_size=window_size, - ) - return from_frost_layout(out, qkv_format), softmax_lse - - -def cp_a2a_bwd_frost_attn( - softmax_scale, - attn_mask_type, - qkv_format, - softmax_lse, - q, - k, - v, - out, - dout, - deterministic=False, - window_size=None, -): - """Backward for CP a2a with the cuDNN FROST backend.""" - from .frost_attention import ( # pylint: disable=import-outside-toplevel - frost_attn_bwd, - from_frost_layout, - to_frost_layout, - ) - - dq, dk, dv = frost_attn_bwd( - to_frost_layout(q.contiguous(), qkv_format), - to_frost_layout(k.contiguous(), qkv_format), - to_frost_layout(v.contiguous(), qkv_format), - to_frost_layout(out.contiguous(), qkv_format), - softmax_lse, - to_frost_layout(dout.contiguous(), qkv_format), - attn_scale=softmax_scale, - attn_mask_type=attn_mask_type, - window_size=window_size, - deterministic=deterministic, - ) - return ( - from_frost_layout(dq, qkv_format), - from_frost_layout(dk, qkv_format), - from_frost_layout(dv, qkv_format), - ) - - -def cp_p2p_fwd_frost_attn( - softmax_scale, - attn_mask_type, - qkv_format, - q_part, - k_part, - v_part, - cu_seqlens_q_per_step, - cu_seqlens_kv_per_step, - section, -): # pylint: disable=unused-argument - """Per-tile forward call of CP P2P with the cuDNN FROST backend. - - cu_seqlens_*_per_step are accepted but unused: they carry the thd offsets, and thd is - declined by the selector. They stay in the signature so the ring can call this and - cp_p2p_fwd_fused_attn with one argument list. - - Returns the same 5-tuple shape as cp_p2p_fwd_fused_attn so the ring code can consume it - unchanged. rng_state, attn_bias and max_logit are None: FROST supports neither dropout nor - bias, and the selector declines those configurations before we get here. - - softmax_lse comes back as [b, h, s] natural-log logsumexp in fp32, which is what the ring - correction in this file consumes. - """ - from .frost_attention import ( # pylint: disable=import-outside-toplevel - frost_attn_fwd, - from_frost_layout, - to_frost_layout, - ) - - out, softmax_lse = frost_attn_fwd( - to_frost_layout(q_part.contiguous(), qkv_format), - to_frost_layout(k_part.contiguous(), qkv_format), - to_frost_layout(v_part.contiguous(), qkv_format), - attn_scale=softmax_scale, - attn_mask_type=_frost_mask_for_section(attn_mask_type, section), - ) - return from_frost_layout(out, qkv_format), softmax_lse, None, None, None - - -def cp_p2p_bwd_frost_attn( - softmax_scale, - attn_mask_type, - qkv_format, - softmax_lse, - softmax_lse_, - q_part, - k_part, - v_part, - out_part, - dout_part, - section, - deterministic=False, -): - """Per-tile backward call of CP P2P with the cuDNN FROST backend. - - Returns (dq, dk, dv, dbias) to match cp_p2p_bwd_fused_attn; dbias is always None. - """ - from .frost_attention import ( # pylint: disable=import-outside-toplevel - frost_attn_bwd, - from_frost_layout, - to_frost_layout, - ) - - softmax_lse_part = softmax_lse_ if section == "upper-triangle" else softmax_lse - dq, dk, dv = frost_attn_bwd( - to_frost_layout(q_part.contiguous(), qkv_format), - to_frost_layout(k_part.contiguous(), qkv_format), - to_frost_layout(v_part.contiguous(), qkv_format), - to_frost_layout(out_part.contiguous(), qkv_format), - softmax_lse_part, - to_frost_layout(dout_part.contiguous(), qkv_format), - attn_scale=softmax_scale, - attn_mask_type=_frost_mask_for_section(attn_mask_type, section), - deterministic=deterministic, - ) - return ( - from_frost_layout(dq, qkv_format), - from_frost_layout(dk, qkv_format), - from_frost_layout(dv, qkv_format), - None, - ) - - class AttnFuncWithCPAndKVP2P(torch.autograd.Function): """ Attention implementation with context parallelism. Exchange KV between CP ranks @@ -1858,6 +1605,7 @@ def forward( attn_bias, deterministic, use_fused_attention, + fused_attention_backend, return_max_logit, softcap, fp8, @@ -1871,7 +1619,6 @@ def forward( use_flash_attn_4, fp8_output, layer_number, - use_frost_attention, ): # pylint: disable=missing-function-docstring @@ -2034,7 +1781,9 @@ def forward( # q, k, v: torch.Tensor, dtype=fwd_nominal_dtype q_f16 = q if use_fused_attention: - fused_attn_backend = FusedAttnBackend["F16_arbitrary_seqlen"] + fused_attn_backend = ( + fused_attention_backend or FusedAttnBackend["F16_arbitrary_seqlen"] + ) if return_max_logit: max_logit_per_step = [ torch.empty(q.shape[-2], dtype=q.dtype, device=q.device) for _ in range(2) @@ -2216,9 +1965,7 @@ def forward( i, cp_size, ] - if use_frost_attention: - frost_attn_inputs = [softmax_scale, attn_mask_type, qkv_format] - elif use_fused_attention: + if use_fused_attention: fused_attn_inputs = [ attn_bias, attn_bias_, @@ -2289,17 +2036,7 @@ def forward( cu_seqlens_kv_per_step[i], ) = prepare_outputs q_inputs[i % 2] = q_part - if use_frost_attention: - ( - out_per_step[i % 2], - softmax_lse_per_step[i % 2], - rng_states[i], - attn_biases[i], - max_logit_per_step[i % 2], - ) = cp_p2p_fwd_frost_attn( - *frost_attn_inputs, *prepare_outputs, section - ) - elif use_fused_attention: + if use_fused_attention: ( out_per_step[i % 2], softmax_lse_per_step[i % 2], @@ -2328,17 +2065,7 @@ def forward( cu_seqlens_kv_per_step[i], ) = prepare_outputs q_inputs[i % 2] = q_part - if use_frost_attention: - ( - out_per_step[i % 2], - softmax_lse_per_step[i % 2], - rng_states[i], - attn_biases[i], - max_logit_per_step[i % 2], - ) = cp_p2p_fwd_frost_attn( - *frost_attn_inputs, *prepare_outputs, section - ) - elif use_fused_attention: + if use_fused_attention: ( out_per_step[i % 2], softmax_lse_per_step[i % 2], @@ -2367,17 +2094,7 @@ def forward( cu_seqlens_kv_per_step[i], ) = prepare_outputs q_inputs[i % 2] = q_part - if use_frost_attention: - ( - out_per_step[i % 2], - softmax_lse_per_step[i % 2], - rng_states[i], - attn_biases[i], - max_logit_per_step[i % 2], - ) = cp_p2p_fwd_frost_attn( - *frost_attn_inputs, *prepare_outputs, section - ) - elif use_fused_attention: + if use_fused_attention: ( out_per_step[i % 2], softmax_lse_per_step[i % 2], @@ -2407,15 +2124,7 @@ def forward( cu_seqlens_kv_per_step[i], ) = prepare_outputs q_inputs[i % 2] = q_part - if use_frost_attention: - ( - out_per_step[i % 2], - softmax_lse_per_step[i % 2], - rng_states[i], - attn_biases[i], - max_logit_per_step[i % 2], - ) = cp_p2p_fwd_frost_attn(*frost_attn_inputs, *prepare_outputs, section) - elif use_fused_attention: + if use_fused_attention: ( out_per_step[i % 2], softmax_lse_per_step[i % 2], @@ -2680,7 +2389,7 @@ def forward( ctx.deterministic = deterministic ctx.softcap = softcap ctx.use_fused_attention = use_fused_attention - ctx.use_frost_attention = use_frost_attention + ctx.fused_attention_backend = fused_attention_backend ctx.pad_between_seqs = pad_between_seqs ctx.softmax_lse_in_packed_format = softmax_lse_in_packed_format ctx.second_half_lse_seqlen = second_half_lse_seqlen @@ -2925,7 +2634,9 @@ def backward(ctx, dout, *_args): ] p2p_comm_buffers[0][0].copy_(kv) if ctx.use_fused_attention: - fused_attn_backend = FusedAttnBackend["F16_arbitrary_seqlen"] + fused_attn_backend = ( + ctx.fused_attention_backend or FusedAttnBackend["F16_arbitrary_seqlen"] + ) # communicate for the 'a2a' part of 'a2a+p2p' dout = dout.view(*ctx.orig_o_shape) @@ -3042,15 +2753,7 @@ def backward(ctx, dout, *_args): cu_seqlens_q_padded, cu_seqlens_kv_padded, ] - if ctx.use_frost_attention: - frost_attn_inputs = [ - ctx.softmax_scale, - ctx.attn_mask_type, - ctx.qkv_format, - softmax_lse, - softmax_lse_, - ] - elif ctx.use_fused_attention: + if ctx.use_fused_attention: fused_attn_inputs = [ ctx.fp8, ctx.fp8_recipe, @@ -3122,14 +2825,7 @@ def backward(ctx, dout, *_args): if i == (cp_size - 1): section = "diagonal" prepare_outputs = cp_p2p_bwd_prepare_qkv(*prepare_inputs, section) - if ctx.use_frost_attention: - dq_, dk_, dv_, dbias_ = cp_p2p_bwd_frost_attn( - *frost_attn_inputs, - *prepare_outputs, - section, - deterministic=ctx.deterministic, - ) - elif ctx.use_fused_attention: + if ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( *fused_attn_inputs, *prepare_outputs, section ) @@ -3142,14 +2838,7 @@ def backward(ctx, dout, *_args): elif i >= (cp_size - rank - 1): section = "lower-triangle" prepare_outputs = cp_p2p_bwd_prepare_qkv(*prepare_inputs, section) - if ctx.use_frost_attention: - dq_, dk_, dv_, dbias_ = cp_p2p_bwd_frost_attn( - *frost_attn_inputs, - *prepare_outputs, - section, - deterministic=ctx.deterministic, - ) - elif ctx.use_fused_attention: + if ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( *fused_attn_inputs, *prepare_outputs, section ) @@ -3162,14 +2851,7 @@ def backward(ctx, dout, *_args): else: section = "upper-triangle" prepare_outputs = cp_p2p_bwd_prepare_qkv(*prepare_inputs, section) - if ctx.use_frost_attention: - dq_, dk_, dv_, dbias_ = cp_p2p_bwd_frost_attn( - *frost_attn_inputs, - *prepare_outputs, - section, - deterministic=ctx.deterministic, - ) - elif ctx.use_fused_attention: + if ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( *fused_attn_inputs, *prepare_outputs, section ) @@ -3182,14 +2864,7 @@ def backward(ctx, dout, *_args): else: section = "all" prepare_outputs = cp_p2p_bwd_prepare_qkv(*prepare_inputs, section) - if ctx.use_frost_attention: - dq_, dk_, dv_, dbias_ = cp_p2p_bwd_frost_attn( - *frost_attn_inputs, - *prepare_outputs, - section, - deterministic=ctx.deterministic, - ) - elif ctx.use_fused_attention: + if ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( *fused_attn_inputs, *prepare_outputs, section ) @@ -3530,7 +3205,7 @@ def backward(ctx, dout, *_args): None, None, None, - None, # use_frost_attention + None, ) @@ -3607,6 +3282,7 @@ def forward( attn_bias, deterministic, use_fused_attention, + fused_attention_backend, return_max_logit, softcap, window_size, @@ -3620,7 +3296,6 @@ def forward( quantizers, fp8_output, load_balancing_strategy, - use_frost_attention, ): # pylint: disable=missing-function-docstring nvtx_range_push("transformer_engine.AttnFuncWithCPAndKVAllGather.forward") @@ -3654,12 +3329,11 @@ def forward( or use_fused_attention or use_flash_attn_3 or use_flash_attn_4 - or use_frost_attention or fa_utils.v2_3_plus ), ( - "cp_comm_type='all_gather' only supports SWA through FusedAttention, FrostAttention" - f" or FlashAttention >= 2.3. Found {use_fused_attention=}, {use_flash_attn_3=}, " - f"{use_flash_attn_4=}, {use_frost_attention=}, " + "cp_comm_type='all_gather' only supports SWA through FusedAttention or FlashAttention" + f" >= 2.3. Found {use_fused_attention=}, {use_flash_attn_3=}, " + f"{use_flash_attn_4=}, " f"and {fa_utils.v2_3_plus=}." ) if load_balancing_strategy is CPLoadBalancingStrategy.DUAL_CHUNK_SWAP: @@ -3775,7 +3449,7 @@ def forward( fp8_meta_kwargs["s_quantizer"] = S_quantizer fp8_meta_kwargs["o_quantizer"] = O_quantizer elif use_fused_attention: - fused_attn_backend = FusedAttnBackend["F16_arbitrary_seqlen"] + fused_attn_backend = fused_attention_backend or FusedAttnBackend["F16_arbitrary_seqlen"] orig_q_shape, _, orig_v_shape = q.shape, k.shape, v.shape orig_o_shape = orig_q_shape[:-1] + orig_v_shape[-1:] @@ -4028,17 +3702,7 @@ def forward( Float8Tensor.make_like(x, data=y, dtype=fwd_nominal_dtype) for x, y in zip([q_fp8, k_fp8, v_fp8], [q_part, k_part, v_part]) ] - if use_frost_attention: - out_per_step[i], softmax_lse_per_step[i] = cp_ag_fwd_frost_attn( - softmax_scale, - qkv_format, - window_size_per_step[i], - q_part, - k_part, - v_part, - ) - rng_states[i] = None # FROST has no dropout, so no RNG state - elif use_fused_attention: + if use_fused_attention: # Set per-step parameters for THD vs bshd/sbhd if qkv_format == "thd": cu_seqlens_q_ = thd_cu_seqlens_q_per_step[i] @@ -4308,7 +3972,7 @@ def forward( ctx.deterministic = deterministic ctx.softcap = softcap ctx.use_fused_attention = use_fused_attention - ctx.use_frost_attention = use_frost_attention + ctx.fused_attention_backend = fused_attention_backend ctx.use_flash_attn_3 = use_flash_attn_3 ctx.use_flash_attn_4 = use_flash_attn_4 ctx.pad_between_seqs = pad_between_seqs @@ -4579,24 +4243,7 @@ def backward(ctx, dout, *_args): out_part = out.select(seq_dim_o, i).contiguous() dout_part = dout.select(seq_dim_o, i).contiguous() - if ctx.use_frost_attention: - ( - dq_per_step[i], - dk_per_step[i], - dv_per_step[i], - ) = cp_ag_bwd_frost_attn( - ctx.softmax_scale, - ctx.qkv_format, - window_size_per_step[i], - softmax_lse_per_step[i], - q_part, - k_part, - v_part, - out_part, - dout_part, - deterministic=ctx.deterministic, - ) - elif ctx.use_fused_attention: + if ctx.use_fused_attention: # Set per-step parameters for THD if ctx.qkv_format == "thd": cu_seqlens_q_ = thd_cu_seqlens_q_per_step[i] @@ -4611,7 +4258,9 @@ def backward(ctx, dout, *_args): softmax_lse_per_step[i], rng_states[i], ] - fused_attn_backend = FusedAttnBackend["F16_arbitrary_seqlen"] + fused_attn_backend = ( + ctx.fused_attention_backend or FusedAttnBackend["F16_arbitrary_seqlen"] + ) fp8_meta_kwargs = {} new_qkv_layout = ctx.qkv_layout do_format = ctx.o_format @@ -4923,7 +4572,7 @@ def backward(ctx, dout, *_args): None, None, None, - None, # use_frost_attention + None, ) @@ -4954,6 +4603,7 @@ def forward( attn_bias, deterministic, use_fused_attention, + fused_attention_backend, return_max_logit, softcap, window_size, @@ -4968,7 +4618,6 @@ def forward( softmax_type, softmax_offset, fp8_output, - use_frost_attention, ): # pylint: disable=missing-function-docstring nvtx_range_push("transformer_engine.AttnFuncWithCPAndQKVOA2A.forward") @@ -4998,12 +4647,10 @@ def forward( or use_fused_attention or use_flash_attn_3 or use_flash_attn_4 - or use_frost_attention or fa_utils.v2_3_plus ), ( - "cp_comm_type='a2a' only supports SWA through FusedAttention, FrostAttention or" - f" FlashAttention >= 2.3. Found {use_fused_attention=}, {use_flash_attn_3=}, " - f"{use_flash_attn_4=}, {use_frost_attention=}, " + "cp_comm_type='a2a' only supports SWA through FusedAttention or FlashAttention >= 2.3." + f" Found {use_fused_attention=}, {use_flash_attn_3=}, {use_flash_attn_4=}, " f"and {fa_utils.v2_3_plus=}." ) assert q.shape[seq_dim_qkv] % 2 == 0 and k.shape[seq_dim_qkv] % 2 == 0, ( @@ -5101,7 +4748,9 @@ def forward( fp8_meta_kwargs["o_quantizer"] = O_quantizer else: if use_fused_attention: - fused_attn_backend = FusedAttnBackend["F16_arbitrary_seqlen"] + fused_attn_backend = ( + fused_attention_backend or FusedAttnBackend["F16_arbitrary_seqlen"] + ) # q, k, v: # FP8DS/FP8CS: torch.uint8 @@ -5156,19 +4805,7 @@ def forward( ) ) qkv_scale_inv_format = None - if use_frost_attention: - out_, softmax_lse = cp_a2a_fwd_frost_attn( - softmax_scale, attn_mask_type, qkv_format, q, k, v, window_size=window_size - ) - # Only the LSE: FROST has no dropout, so there is no RNG state to carry, and a - # None in this list would have to survive the save/restore machinery. - aux_ctx_tensors = [softmax_lse] - # out_part is what gets saved for backward (f16_tensors below). Leaving it at its - # None initialisation makes `out` arrive as None in backward, which is not obvious - # from this branch alone: the fused path sets it inside its fp8 bookkeeping. - out_part = out_ - out_f16 = out_ - elif use_fused_attention: + if use_fused_attention: if fp8: if fp8_recipe.mxfp8(): q_fp8, k_fp8, v_fp8, qkv_layout, qkv_scale_inv_format = combine_and_quantize( @@ -5404,10 +5041,7 @@ def forward( ctx.softcap = softcap ctx.window_size = window_size ctx.use_fused_attention = use_fused_attention - ctx.use_frost_attention = use_frost_attention - # The a2a class never needed qkv_format in backward before: the fused and flash paths - # take a qkv_layout instead. FROST builds its graphs from the tensor layout, so it does. - ctx.qkv_format = qkv_format + ctx.fused_attention_backend = fused_attention_backend ctx.fp8_meta = fp8_meta ctx.is_input_fp8 = is_input_fp8 ctx.is_output_fp8 = is_output_fp8 @@ -5485,7 +5119,9 @@ def backward(ctx, dout, *_args): if isinstance(dout, QuantizedTensorStorage): dout = dout.dequantize(dtype=bwd_nominal_dtype) if ctx.use_fused_attention: - fused_attn_backend = FusedAttnBackend["F16_arbitrary_seqlen"] + fused_attn_backend = ( + ctx.fused_attention_backend or FusedAttnBackend["F16_arbitrary_seqlen"] + ) dout = dout.view(*ctx.orig_o_shape) # dout: @@ -5559,25 +5195,7 @@ def backward(ctx, dout, *_args): fa_backward_kwargs["softcap"] = ctx.softcap dq_fp8, dk_fp8, dv_fp8 = None, None, None - # Only the fused branch below binds this, and only the fused branch reads it further - # down -- but with three branches that binding no longer dominates the read, so give it - # a definition rather than rely on the conditions staying in step. - rest = [] - if ctx.use_frost_attention: - dq, dk, dv = cp_a2a_bwd_frost_attn( - ctx.softmax_scale, - ctx.attn_mask_type, - ctx.qkv_format, - aux_ctx_tensors[0], - q, - k, - v, - out, - dout, - deterministic=ctx.deterministic, - window_size=ctx.window_size, - ) - elif ctx.use_fused_attention: + if ctx.use_fused_attention: do_format = ctx.o_format do_scale_inv_format = None q_part, k_part, v_part, out_part, dout_part = q, k, v, out, dout @@ -5810,7 +5428,7 @@ def backward(ctx, dout, *_args): None, d_softmax_offset, None, - None, # use_frost_attention + None, ) @@ -5950,7 +5568,7 @@ def attn_forward_func_with_cp( attn_bias=None, deterministic=False, use_fused_attention=False, - use_frost_attention=False, + fused_attention_backend=None, window_size=None, softcap=0.0, fp8=False, @@ -6088,11 +5706,8 @@ def attn_forward_func_with_cp( assert cu_seqlens_q is cu_seqlens_kv and ( cu_seqlens_q_padded is cu_seqlens_kv_padded ), "No-load-balance THD self-attention requires shared Q/KV sequence metadata tensors." - # The restriction is FlashAttention-specific; the condition infers "not fused means flash", - # which predates FROST. FROST builds its cuDNN graphs from each tensor's actual strides, so - # sbhd is served directly. This matters because Megatron uses sbhd internally. assert ( - qkv_format != "sbhd" or use_fused_attention or use_frost_attention + qkv_format != "sbhd" or use_fused_attention ), "Context parallelism does not support FlashAttention backend with qkv_format = 'sbhd'!" assert attn_bias is None or (use_fused_attention and "padding" not in attn_mask_type), ( "Context parallelism only supports attention bias with FusedAttention backend and" @@ -6141,6 +5756,7 @@ def attn_forward_func_with_cp( attn_bias, deterministic, use_fused_attention, + fused_attention_backend, return_max_logit, softcap, ] @@ -6158,7 +5774,6 @@ def attn_forward_func_with_cp( use_flash_attn_4, fp8_output, layer_number, - use_frost_attention, ] out = AttnFuncWithCPAndKVP2P.apply(*args) elif cp_comm_type == "all_gather": @@ -6174,7 +5789,6 @@ def attn_forward_func_with_cp( quantizers, fp8_output, load_balancing_strategy, - use_frost_attention, ] out = AttnFuncWithCPAndKVAllGather.apply(*args) elif cp_comm_type == "a2a": @@ -6191,7 +5805,6 @@ def attn_forward_func_with_cp( softmax_type, softmax_offset, fp8_output, - use_frost_attention, ] out = AttnFuncWithCPAndQKVOA2A.apply(*args) else: diff --git a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py index 41f94352dff..658dab5d88d 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py @@ -65,7 +65,6 @@ UnfusedDotProductAttention, FusedAttention, FlashAttention, - FrostAttention, ) @@ -80,7 +79,6 @@ "use_fused_attention": None, "fused_attention_backend": None, "use_unfused_attention": None, - "use_frost_attention": None, "backend_selection_requires_update": False, } @@ -158,7 +156,6 @@ def _get_thd_policy_attention_backend( use_fused_attention, fused_attention_backend, use_unfused_attention, - use_frost_attention, _, ) = selection _attention_backends.update( @@ -169,7 +166,6 @@ def _get_thd_policy_attention_backend( "use_fused_attention": use_fused_attention, "fused_attention_backend": fused_attention_backend, "use_unfused_attention": use_unfused_attention, - "use_frost_attention": use_frost_attention, "backend_selection_requires_update": False, } ) @@ -1000,16 +996,6 @@ def __init__( return_max_logit=self.return_max_logit, ) - # Only selectable for symmetric head_dim in (256, 512] on SM100/SM103, where no other - # backend can run at all. Cheap to construct, so instantiate unconditionally like the rest. - self.frost_attention = FrostAttention( - softmax_scale, - attention_type=attention_type, - layer_number=layer_number, - deterministic=self.deterministic, - **attn_kwargs, - ) - self.unfused_attention = UnfusedDotProductAttention( softmax_scale, attention_type=attention_type, @@ -2858,9 +2844,6 @@ def forward( use_flash_attention = False use_fused_attention = False use_unfused_attention = True - # Bound here too: the availability check below reads all four flags at this - # scope, and this branch never calls get_attention_backend. - use_frost_attention = False else: if ( _attention_backends["attention_params"] is None @@ -2875,7 +2858,6 @@ def forward( use_fused_attention, fused_attention_backend, use_unfused_attention, - use_frost_attention, _, ) = dpa_utils.get_attention_backend(attention_params) # Set global _attention_backends var using return value @@ -2885,7 +2867,6 @@ def forward( _attention_backends["use_fused_attention"] = use_fused_attention _attention_backends["fused_attention_backend"] = fused_attention_backend _attention_backends["use_unfused_attention"] = use_unfused_attention - _attention_backends["use_frost_attention"] = use_frost_attention _attention_backends["backend_selection_requires_update"] = False # logging.Logger methods graph-break under torch.compile, so # selection is only logged in eager -- as in @@ -2904,8 +2885,6 @@ def forward( "Running with FusedAttention backend (sub-backend %s)", int(fused_attention_backend), ) - elif use_frost_attention: - logger.info("Running with FrostAttention backend (cuDNN FROST)") elif use_unfused_attention: logger.info("Running with UnfusedDotProductAttention backend") else: @@ -2914,20 +2893,9 @@ def forward( use_fused_attention = _attention_backends["use_fused_attention"] fused_attention_backend = _attention_backends["fused_attention_backend"] use_unfused_attention = _attention_backends["use_unfused_attention"] - use_frost_attention = _attention_backends["use_frost_attention"] # raise exception if no backend is available - if ( - sum( - [ - use_flash_attention, - use_fused_attention, - use_unfused_attention, - use_frost_attention, - ] - ) - == 0 - ): + if sum([use_flash_attention, use_fused_attention, use_unfused_attention]) == 0: raise ValueError( "No dot product attention backend is available for the provided inputs. Please" " run with NVTE_DEBUG=1 NVTE_DEBUG_LEVEL=2 to find out the reasons for" @@ -3090,27 +3058,6 @@ def forward( bf16_backward=bf16_backward, ) - if use_frost_attention: - return self.frost_attention( - query_layer, - key_layer, - value_layer, - qkv_format=qkv_format, - cu_seqlens_q=cu_seqlens_q, - cu_seqlens_kv=cu_seqlens_kv, - max_seqlen_q=max_seqlen_q, - max_seqlen_kv=max_seqlen_kv, - cu_seqlens_q_padded=cu_seqlens_q_padded, - cu_seqlens_kv_padded=cu_seqlens_kv_padded, - attn_mask_type=attn_mask_type, - window_size=window_size, - cp_group=self.cp_group, - cp_global_ranks=self.cp_global_ranks, - cp_stream=self.cp_stream, - cp_comm_type=self.cp_comm_type, - load_balancing_strategy=self.load_balancing_strategy, - ) - if use_unfused_attention: allow_emulation = ( os.getenv("NVTE_UnfusedDPA_Emulate_FP8", "0") == "1" or is_in_onnx_export_mode() diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index b300be40a78..f1956aa2aec 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -42,6 +42,8 @@ __all__ = [ "is_frost_attention_available", "is_frost_attention_supported", + "fused_attn_fwd", + "fused_attn_bwd", "frost_attn_fwd", "frost_attn_bwd", "to_frost_layout", @@ -235,50 +237,148 @@ def _mask_options(cudnn, spec): return cudnn_pygraph.diagonal_band_kwargs(cudnn, attn_mask_type, window) -def is_frost_attention_supported( - head_dim_qk: int, - head_dim_v: int, - qkv_dtype: torch.dtype, - attn_mask_type: str, - dropout: float = 0.0, - attn_bias_type: str = "no_bias", - window_size: Optional[Tuple[int, int]] = None, -) -> Tuple[bool, str]: - """Whether this specific attention configuration should route to FROST. - - Shape and dtype are checked before availability, and the ordering is deliberate rather than - stylistic. Probing availability imports cuDNN Frontend and sets - CUDNN_FRONTEND_ENABLE_FROST_ENGINES, which registers extra engines process-wide and so is - visible to every other cuDNN consumer in the process. This function runs for every attention - config on the machine, the vast majority of which are nowhere near head_dim 512, and none of - them should pay that cost or have their engine pool changed underneath them. +_SUPPORTED_QKV_FORMATS = ("bshd", "sbhd") + + +def _qkv_format_from_layout(qkv_layout: str) -> str: + """The single qkv_format a TE qkv_layout names, e.g. 'bshd_bshd_bshd' -> 'bshd'.""" + formats = { + "".join(c for c in part if c.isalpha()) + for part in qkv_layout.replace("paged_kv_", "").split("_") + } + if len(formats) != 1: + raise NotImplementedError( + f"FROST attention needs q, k and v in one format; got qkv_layout {qkv_layout!r}" + ) + return formats.pop() + + +def _te_mask_spec(attn_mask_type: str, window_size, bottom_right_diagonal: bool): + """Fold TE's (mask type, window, diagonal anchor) into the spec the plan is keyed on. + + TE carries the anchor in its own flag, so normalise it into the mask type before building the + band: diagonal_band_kwargs reads the anchor off the name, and taking it from the name alone + would quietly give a top-left band where the caller asked for bottom-right. """ + if "padding" in attn_mask_type: + raise NotImplementedError( + f"FROST attention does not support a padding mask; got {attn_mask_type!r}" + ) + left, right = _NO_WINDOW if window_size is None else tuple(window_size) + if "causal" in attn_mask_type and right == -1: + right = 0 + if right == 0: + attn_mask_type = "causal_bottom_right" if bottom_right_diagonal else "causal" + else: + attn_mask_type = "no_mask" + return _mask_spec(attn_mask_type, (left, right)) + + +def _name_for(table, value, default=None): + """Reverse a cpp_extensions str-to-enum table.""" + for name, enum_value in table.items(): + if enum_value == value: + return name + return default + + +def is_frost_attention_supported(params) -> Tuple[int, str]: + """Whether this fused-attention config should run on the FROST sub-backend. + + Takes a FusedAttentionParams and returns (sub-backend value, reject message), the same shape + as tex.get_fused_attn_backend, so get_attention_backend can fall through to it when the C++ + backends decline. + + Deliberately does not probe availability. That imports cuDNN Frontend with the FROST engines + enabled, which changes the engine pool for every cuDNN consumer in the process, and this runs + for every attention config on the machine. get_attention_backend checks availability once at + the end, the way it checks flash-attn versions. + """ + # pylint: disable-next=import-outside-toplevel + from ...cpp_extensions.fused_attn import ( + AttnBiasType, + AttnMaskType, + FusedAttnBackend, + QKVFormat, + QKVLayout, + SoftmaxType, + TORCH_DType, + ) + + no_backend = int(FusedAttnBackend.No_Backend) + + if int(os.environ.get("NVTE_FROST_ATTN", "1")) == 0: + return no_backend, "FROST is disabled by NVTE_FROST_ATTN=0" + + head_dim_qk, head_dim_v = params.head_dim_qk, params.head_dim_v if head_dim_qk != head_dim_v: - return False, f"FROST path requires symmetric head_dim; got {head_dim_qk}/{head_dim_v}" + return no_backend, f"FROST requires symmetric head_dim; got {head_dim_qk}/{head_dim_v}" if not _MIN_HEAD_DIM <= head_dim_qk <= _MAX_HEAD_DIM: - return False, f"FROST path covers head_dim in (256, 512]; got {head_dim_qk}" + return no_backend, f"FROST covers head_dim in (256, 512]; got {head_dim_qk}" if head_dim_qk % _HEAD_DIM_MULTIPLE != 0: return ( - False, - ( - f"FROST path needs head_dim to be a multiple of {_HEAD_DIM_MULTIPLE}; got" - f" {head_dim_qk}" - ), + no_backend, + f"FROST needs head_dim to be a multiple of {_HEAD_DIM_MULTIPLE}; got {head_dim_qk}", ) + + qkv_dtype = TORCH_DType.get(params.qkv_dtype) if qkv_dtype not in (torch.bfloat16, torch.float16): - return False, f"FROST path supports bf16/fp16; got {qkv_dtype}" - if dropout != 0.0: - return False, "FROST path does not support dropout" - if attn_bias_type != "no_bias": - return False, "FROST path does not support attention bias" + return no_backend, f"FROST supports bf16/fp16; got {params.qkv_dtype}" + if params.dropout != 0.0: + return no_backend, "FROST does not support dropout" + if _name_for(AttnBiasType, params.bias_type) != "no_bias": + return no_backend, "FROST does not support attention bias" + if _name_for(SoftmaxType, params.softmax_type) != "vanilla": + return no_backend, "FROST only supports vanilla softmax" + if params.num_pages_k != 0 or params.num_pages_v != 0: + return no_backend, "FROST does not support paged KV" + if params.return_max_logit: + return no_backend, "FROST does not return max_logit" + if params.cuda_graph: + return no_backend, "FROST graphs are built lazily and cannot be captured" + if params.deterministic and params.is_training: + # The backward uses an atomic dQ accumulation whose order is not fixed, so repeat runs + # differ in the last bits. Nothing selects a deterministic variant, so decline instead. + return no_backend, "FROST does not have a deterministic backward" + + qkv_layout = _name_for(QKVLayout, params.qkv_layout) + if qkv_layout is None: + return no_backend, f"FROST got an unrecognised qkv_layout {params.qkv_layout}" + try: + qkv_format = _qkv_format_from_layout(qkv_layout) + except NotImplementedError as exc: + return no_backend, str(exc) + if qkv_format not in _SUPPORTED_QKV_FORMATS: + return ( + no_backend, + f"FROST supports qkv_format in {_SUPPORTED_QKV_FORMATS}; got {qkv_format}", + ) + # The kernels write O and dQKV with q's strides, so any format that differs from the input + # would need a copy the fused path does not make. Nothing asks for one today. + for name, value in ( + ("o_format", _name_for(QKVFormat, params.o_format)), + ("do_format", _name_for(QKVFormat, params.do_format)), + ("dqkv_layout", _name_for(QKVLayout, params.dqkv_layout)), + ): + if value is None: + continue + value = _qkv_format_from_layout(value) if name == "dqkv_layout" else value + if value != qkv_format: + return no_backend, f"FROST needs {name} to match qkv_format; got {value}/{qkv_format}" + + attn_mask_type = _name_for(AttnMaskType, params.attn_mask_type) + if attn_mask_type is None: + return no_backend, f"FROST got an unrecognised attn_mask_type {params.attn_mask_type}" try: - _mask_spec(attn_mask_type, window_size) + _te_mask_spec( + attn_mask_type, + (params.window_size_left, params.window_size_right), + params.bottom_right_diagonal, + ) except NotImplementedError as exc: - return False, str(exc) - ok, reason = is_frost_attention_available() - if not ok: - return False, reason - return True, "" + return no_backend, str(exc) + + return int(FusedAttnBackend.FROST), "" def to_frost_layout(t: torch.Tensor, qkv_format: str) -> torch.Tensor: @@ -647,3 +747,156 @@ def _as(t, ref): handle=_handle_for(q.device), ) return dq, dk, dv + + +def _frost_only(**unsupported): + """Raise if any feature the selector should have declined reached the kernels anyway.""" + for name, value in unsupported.items(): + if value: + raise NotImplementedError(f"FROST attention does not support {name}") + + +def fused_attn_fwd( + is_training, + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q, + k, + v, + fake_dtype, + fused_attention_backend, + attn_bias=None, + cu_seqlens_q_padded=None, + cu_seqlens_kv_padded=None, + page_table_k=None, + page_table_v=None, + s_quantizer=None, + o_quantizer=None, + attn_scale=None, + dropout=0.0, + fast_zero_fill=True, + qkv_layout="sbh3d", + o_format="sbhd", + qkv_scale_inv_format=None, + attn_bias_type="no_bias", + attn_mask_type="padding", + softmax_type="vanilla", + window_size=(-1, -1), + bottom_right_diagonal=None, + rng_gen=None, + softmax_offset=None, + return_max_logit=False, + cuda_graph=False, +): # pylint: disable=unused-argument + """FROST forward behind the cpp_extensions.fused_attn_fwd signature. + + Mirrors that signature so FusedAttnFunc and the context-parallel ring reach these kernels + without knowing which sub-backend they got. Returns (out, aux_ctx_tensors) with + aux_ctx_tensors = [softmax_lse, rng_state]; softmax_lse is [b, h, s] fp32 natural-log + logsumexp, which is what the ring correction consumes. + + cu_seqlens and the padded variants are ignored: they carry thd offsets, and thd is declined + at selection. + """ + _frost_only( + dropout=dropout != 0.0, + attention_bias=attn_bias_type != "no_bias", + paged_kv=page_table_k is not None or page_table_v is not None, + fp8=s_quantizer is not None or o_quantizer is not None, + sink_attention=softmax_type != "vanilla", + max_logit=return_max_logit, + cuda_graph_capture=cuda_graph, + ) + qkv_format = _qkv_format_from_layout(qkv_layout) + if o_format != qkv_format: + raise NotImplementedError( + f"FROST attention needs o_format to match qkv_format; got {o_format}/{qkv_format}" + ) + mask_type, window = _te_mask_spec(attn_mask_type, window_size, bool(bottom_right_diagonal)) + + out, softmax_lse = frost_attn_fwd( + to_frost_layout(q.contiguous(), qkv_format), + to_frost_layout(k.contiguous(), qkv_format), + to_frost_layout(v.contiguous(), qkv_format), + attn_scale=attn_scale, + attn_mask_type=mask_type, + window_size=window, + ) + # A real tensor rather than None: it is saved for backward and handed to the activation + # offload hooks alongside softmax_lse, neither of which accepts None. FROST has no dropout, + # so nothing reads it. + rng_state = torch.empty(2, dtype=torch.int64, device=q.device) + return from_frost_layout(out, qkv_format), [softmax_lse, rng_state] + + +def fused_attn_bwd( + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q, + k, + v, + o, + d_o, + fake_dtype, + aux_ctx_tensors, + fused_attention_backend, + cu_seqlens_q_padded=None, + cu_seqlens_kv_padded=None, + s_quantizer=None, + dp_quantizer=None, + dqkv_quantizer=None, + attn_scale=None, + dropout=0.0, + fast_zero_fill=True, + qkv_layout="sbh3d", + o_format="sbhd", + do_format="sbhd", + dqkv_layout="sbh3d", + qkv_scale_inv_format=None, + do_scale_inv_format=None, + attn_bias_type="no_bias", + attn_mask_type="padding", + softmax_type="vanilla", + window_size=(-1, -1), + bottom_right_diagonal=None, + deterministic=False, + cuda_graph=False, +): # pylint: disable=unused-argument + """FROST backward behind the cpp_extensions.fused_attn_bwd signature. + + Returns (dq, dk, dv, dbias) with dbias always None, matching what the fused path returns for + a no_bias config. + """ + _frost_only( + dropout=dropout != 0.0, + attention_bias=attn_bias_type != "no_bias", + fp8=s_quantizer is not None or dqkv_quantizer is not None, + sink_attention=softmax_type != "vanilla", + cuda_graph_capture=cuda_graph, + ) + qkv_format = _qkv_format_from_layout(qkv_layout) + mask_type, window = _te_mask_spec(attn_mask_type, window_size, bool(bottom_right_diagonal)) + softmax_lse = aux_ctx_tensors[0] + + dq, dk, dv = frost_attn_bwd( + to_frost_layout(q.contiguous(), qkv_format), + to_frost_layout(k.contiguous(), qkv_format), + to_frost_layout(v.contiguous(), qkv_format), + to_frost_layout(o.contiguous(), o_format), + softmax_lse, + to_frost_layout(d_o.contiguous(), do_format), + attn_scale=attn_scale, + attn_mask_type=mask_type, + deterministic=deterministic, + window_size=window, + ) + return ( + from_frost_layout(dq, qkv_format), + from_frost_layout(dk, qkv_format), + from_frost_layout(dv, qkv_format), + None, + ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 1ee6cf94873..83e6b8e211f 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -449,9 +449,20 @@ def _get_fused_attn_backend(**fused_attn_kwargs): graph break because it is baked into the graph as a literal, while an enum member comes out of the reconstruction corrupted (see the cast at the call site, which restores the enum).""" - fused_attention_backend, reject_message = tex.get_fused_attn_backend( - FusedAttentionParams(**fused_attn_kwargs) - ) + params = FusedAttentionParams(**fused_attn_kwargs) + fused_attention_backend, reject_message = tex.get_fused_attn_backend(params) + if fused_attention_backend == FusedAttnBackend.No_Backend: + # FROST is a python sub-backend, so the C++ selector cannot see it. It serves symmetric + # head_dim in (256, 512] on SM100/SM103, which nothing above it covers. Availability is + # checked once at the end of get_attention_backend, the way flash-attn's version is. + from .frost_attention import ( # pylint: disable=import-outside-toplevel + is_frost_attention_supported, + ) + + frost_backend, frost_reject = is_frost_attention_supported(params) + if frost_backend != FusedAttnBackend.No_Backend: + return int(frost_backend), frost_reject + reject_message = f"{reject_message} {frost_reject}" return int(fused_attention_backend), reject_message @@ -481,8 +492,6 @@ def get_attention_backend( available_backends : List[bool] All available backends that could support the provided input. A list of Booleans in the form of [use_flash_attention, use_fused_attention, use_unfused_attention]. - FrostAttention is deliberately not a member: the list's length is relied on by - existing three-way unpacks. Use the `use_frost_attention` return value instead. """ # NOTE: As part of refactoring attention.py, populating the _attention_backends cache in attention # is no longer performed at the end of get_attention_backend(), but the responsibility of doing so @@ -615,7 +624,6 @@ def get_attention_backend( flash_attention_backend = None use_fused_attention = int(os.environ.get("NVTE_FUSED_ATTN", "1")) use_unfused_attention = int(os.environ.get("NVTE_UNFUSED_ATTN", "1")) - use_frost_attention = int(os.environ.get("NVTE_FROST_ATTN", "1")) if not use_flash_attention_2 and FlashAttentionUtils.is_installed: logger.debug("Disabling FlashAttention 2 due to NVTE_FLASH_ATTN=0 or NVTE_FLASH_ATTN_V2=0") if not use_flash_attention_3 and FlashAttentionUtils.v3_is_installed: @@ -1862,170 +1870,6 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt ), ) FlashAttentionUtils.warning_printed = True - # FROST serves symmetric head_dim in (256, 512]; it is the only backend that also does - # context parallelism there. Experimental, and declined per-shape below. - if use_frost_attention: - # Local import: keeps TE importable without cudnn-frontend installed. - from .frost_attention import ( # pylint: disable=import-outside-toplevel - is_frost_attention_supported, - ) - - frost_supported, frost_reason = is_frost_attention_supported( - head_dim_qk=head_dim_qk, - head_dim_v=head_dim_v, - qkv_dtype=qkv_dtype, - attn_mask_type=attn_mask_type, - dropout=attention_dropout, - attn_bias_type=core_attention_bias_type, - window_size=window_size, - ) - if not frost_supported: - logger.debug("Disabling FrostAttention: %s", frost_reason) - use_frost_attention = False - # Conservative guards for capabilities that exist in cuDNN but are not validated here yet. - # Each is a silent-wrong-answer risk rather than an error, so default to declining. - if use_frost_attention and softmax_type != "vanilla": - # CP asserts non-vanilla softmax needs FusedAttention; FROST implements plain softmax. - logger.debug("Disabling FrostAttention for softmax_type = %s", softmax_type) - use_frost_attention = False - if use_frost_attention and fp8: - logger.debug("Disabling FrostAttention for FP8") - use_frost_attention = False - if use_frost_attention and softcap is not None and softcap != 0.0: - logger.debug("Disabling FrostAttention for softcap") - use_frost_attention = False - if use_frost_attention and "thd" in qkv_layout: - # bshd and sbhd are served directly from their own strides; thd is packed/varlen, which - # needs cu_seqlens plumbing that is neither implemented nor validated here. - logger.debug("Disabling FrostAttention for qkv_layout = %s", qkv_layout) - use_frost_attention = False - if use_frost_attention and deterministic and is_training: - # Measured on B200 with cuDNN Frontend 1.29.0: requesting a deterministic backward is - # refused outright -- cudnnGraphNotSupportedError, no engine proposes a plan -- so unlike - # the C++ fused path there is nothing to opt into. Declining keeps - # NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 an honest guarantee instead of silently running the - # non-deterministic kernel. The graph still passes the flag, so this lifts on its own if - # cuDNN ships a deterministic d512 backward. - logger.debug("Disabling FrostAttention as its backward has no deterministic cuDNN plan") - use_frost_attention = False - if use_frost_attention and (has_score_mod or has_score_mod_bprop): - # The score_mod filter above disables flash, fused and unfused, and at head_dim 512 the - # fused path is unavailable anyway -- so without this FROST would be the sole survivor - # and would compute plain attention with the callback silently dropped. That includes the - # score_mod_bprop-without-score_mod case, which is meant to end in "no backend available". - logger.debug("Disabling FrostAttention for score_mod") - use_frost_attention = False - if use_frost_attention and attention_params.qkv_type is not torch.Tensor: - # Every other backend filters on the tensor class, not just the dtype: a quantized tensor - # can carry a nominal bf16 dtype outside an fp8 autocast, and the fp8 guard below keys on - # the autocast flag rather than the type. - # - # Read from attention_params, not the local: the fused-attention dtype spec rebinds - # qkv_type to an NVTE dtype enum well before this point, so the local compares unequal to - # torch.Tensor for every input and would decline FROST unconditionally. - logger.debug("Disabling FrostAttention for qkv_type = %s", attention_params.qkv_type) - use_frost_attention = False - if use_frost_attention and num_splits != 1: - # Declined for the same reason the fused and unfused paths are: silently ignoring it - # would change the computation the caller asked for. - logger.debug("Disabling FrostAttention for num_splits = %s", num_splits) - use_frost_attention = False - if use_frost_attention and checkpoint_core_attention: - # The backend FROST displaces at this head dim is unfused, which does honour activation - # recompute. Selecting FROST would silently remove it, which is a memory regression - # rather than a wrong answer, but not one the caller asked for. - logger.debug("Disabling FrostAttention for checkpoint_core_attention") - use_frost_attention = False - if use_frost_attention and cuda_graph: - # Plan lookup and lazy handle creation are host-side work on the first call, which is - # hazardous inside a capture. Not validated under capture, so decline rather than guess. - logger.debug("Disabling FrostAttention for CUDA graph capture") - use_frost_attention = False - if use_frost_attention and return_max_logit: - # FrostAttention returns the context layer alone, where UnfusedDotProductAttention returns - # (context, max_logit). Selecting it here would break the caller's unpack. - logger.debug("Disabling FrostAttention for max_logit") - use_frost_attention = False - if use_frost_attention and inference_params is not None: - # Unreachable today, since KV caching asserts a padding mask and FROST declines those. - # Explicit anyway: no page table reaches the backend, so a paged cache would be read raw. - logger.debug("Disabling FrostAttention for KV caching") - use_frost_attention = False - if ( - use_frost_attention - and window_size is not None - and window_size[0] != -1 - and "causal" not in attn_mask_type - and max_seqlen_q != max_seqlen_kv - ): - # FROST anchors the band from the mask type, so a windowed non-causal mask always lands - # top-left. TE's bottom_right_diagonal defaults to True and the C++ fused path honours it - # (fused_attn_f16_arbitrary_seqlen.cu picks the alignment from that flag), so for unequal - # q/kv lengths the two would disagree silently. Decline rather than guess the anchor. - logger.debug( - "Disabling FrostAttention for a windowed non-causal mask with max_seqlen_q != " - "max_seqlen_kv, where the diagonal anchor is ambiguous" - ) - use_frost_attention = False - has_sliding_window = window_size is not None and ( - window_size[0] != -1 or window_size[1] not in [-1, 0] - ) - if ( - use_frost_attention - and context_parallel - and has_sliding_window - and cp_comm_type in ["p2p", "a2a+p2p"] - ): - # Same rule FusedAttention carries, and for a reason visible in the ring itself: the p2p - # path hardcodes the per-step window to (-1, 0) or (-1, -1) at every kernel call, so a - # user window is discarded there for any backend. all_gather has real machinery for this - # (window_size_per_step, from get_kv_seq_info_after_all_gather) and a2a sees the whole - # sequence, so both can serve it. - logger.debug( - "Disabling FrostAttention as it does not support context parallelism with sliding" - " window attention and cp_comm_type = %s", - cp_comm_type, - ) - use_frost_attention = False - if use_frost_attention and context_parallel: - # Same two restrictions FlashAttention and FusedAttention carry above. Both are about - # where the causal diagonal sits: the ring shards q and kv independently, so a mask whose - # position depends on the q/kv lengths lands differently per step. no_mask is unaffected - # and stays allowed even when the lengths differ. - if "bottom_right" in attn_mask_type: - logger.debug( - "Disabling FrostAttention as it does not support context parallelism with" - " causal_bottom_right masking" - ) - use_frost_attention = False - elif "causal" in attn_mask_type and max_seqlen_q != max_seqlen_kv: - logger.debug( - "Disabling FrostAttention as it does not support context parallelism with causal" - " masking for cross-attention" - ) - use_frost_attention = False - if ( - use_frost_attention - and context_parallel - and cp_comm_type - not in ( - "p2p", - "all_gather", - "a2a", - "a2a+p2p", - ) - ): - # a2a+p2p needs no separate wiring: it dispatches to AttnFuncWithCPAndKVP2P, the same class - # as plain p2p, and its a2a stage is flash_attn_a2a_communicate -- a redistribution between - # sequence- and head-sharding that calls no attention kernel. The per-step calls are the - # ordinary p2p section calls with fewer heads per rank. - # Non-p2p types matter for Gemma-4: TE refuses sliding-window attention with p2p, and the - # model has sliding layers, so those layers need all_gather or a2a. - logger.debug( - "Disabling FrostAttention for context parallelism with cp_comm_type = %s", cp_comm_type - ) - use_frost_attention = False - # All available backends if use_flash_attention_2 and not FlashAttentionUtils.is_installed: use_flash_attention_2 = False @@ -2034,6 +1878,18 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt if use_flash_attention_4 and not FlashAttentionUtils.v4_is_installed: use_flash_attention_4 = False use_flash_attention = use_flash_attention_2 or use_flash_attention_3 or use_flash_attention_4 + if use_fused_attention and fused_attention_backend == FusedAttnBackend.FROST.value: + # Deferred to here because probing it imports cuDNN Frontend with the FROST engines + # enabled, which changes the engine pool for every cuDNN consumer in the process. + from .frost_attention import ( # pylint: disable=import-outside-toplevel + is_frost_attention_available, + ) + + frost_available, frost_reason = is_frost_attention_available() + if not frost_available: + logger.debug("Disabling FusedAttention: %s", frost_reason) + use_fused_attention = False + fused_attention_backend = None available_backends = [use_flash_attention, use_fused_attention, use_unfused_attention] if use_flash_attention_2: flash_attention_backend = FlashAttentionUtils.version @@ -2044,7 +1900,7 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt logger.debug( "Available backends = {FlashAttention=%s%s, FusedAttention=%s%s," - " UnfusedDotProductAttention=%s, FrostAttention=%s}", + " UnfusedDotProductAttention=%s}", bool(available_backends[0]), (f" ({str(flash_attention_backend)})" if flash_attention_backend is not None else ""), bool(available_backends[1]), @@ -2054,10 +1910,6 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt else "" ), bool(available_backends[2]), - # Read from the local flag rather than available_backends, which excludes FROST by - # design. Without this the log reports every backend as unavailable and then selects - # FrostAttention a few lines later, which reads as a contradiction. - bool(use_frost_attention), ) # Prefer FA2 for THD training with dropout on SM100/103, where FusedAttention has a known @@ -2087,20 +1939,13 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt if use_flash_attention: use_fused_attention = False use_unfused_attention = False - use_frost_attention = False elif use_fused_attention: use_unfused_attention = False - use_frost_attention = False - elif use_frost_attention: - # Preferred over the unfused path: same shape coverage, but fused and CP-capable. - use_unfused_attention = False selected_backend = "NoBackend" if use_flash_attention: selected_backend = f"FlashAttention ({str(flash_attention_backend)})" elif use_fused_attention: selected_backend = f"FusedAttention (sub-backend {int(fused_attention_backend)})" - elif use_frost_attention: - selected_backend = "FrostAttention (cuDNN FROST)" elif use_unfused_attention: selected_backend = "UnfusedDotProductAttention" logger.debug("Selected backend = %s.", selected_backend) @@ -2111,7 +1956,6 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt use_fused_attention, fused_attention_backend, use_unfused_attention, - use_frost_attention, available_backends, ) diff --git a/transformer_engine/pytorch/cpp_extensions/fused_attn.py b/transformer_engine/pytorch/cpp_extensions/fused_attn.py index 9a33df7634b..8e731f813e3 100644 --- a/transformer_engine/pytorch/cpp_extensions/fused_attn.py +++ b/transformer_engine/pytorch/cpp_extensions/fused_attn.py @@ -103,7 +103,7 @@ class FusedAttnBackend(IntEnum): """Fused attention sub-backends. This is the canonical fused-attention backend enum for - ``transformer_engine.pytorch``. It mirrors the backend + ``transformer_engine.pytorch``. It mirrors every member of the backend ``transformer_engine_torch.NVTE_Fused_Attn_Backend`` (pybind11) enum value-for-value, and instances of the two enums compare equal when they share the same integer value. Unlike the pybind enum, a plain-python @@ -119,6 +119,10 @@ class FusedAttnBackend(IntEnum): No_Backend = int(NVTE_Fused_Attn_Backend.NVTE_No_Backend) F16_arbitrary_seqlen = int(NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen) FP8 = int(NVTE_Fused_Attn_Backend.NVTE_FP8) + # Python-only: cuDNN FROST runs through the cuDNN Frontend python API rather than the C++ + # fused-attention path, so it has no NVTE_Fused_Attn_Backend counterpart. fused_attn_fwd/bwd + # route it to frost_attention.py before any C++ call, so this value never reaches pybind. + FROST = 3 @classmethod def cast( @@ -153,9 +157,14 @@ def __hash__(self) -> int: return int.__hash__(self) +# Members with no C++ counterpart; excluded from the sync check below. +_PYTHON_ONLY_FUSED_ATTN_BACKENDS = frozenset({FusedAttnBackend.FROST}) + # Fail fast at import time if a new enumerator is added on the C++ side # without being mirrored above. -assert {f"NVTE_{m.name}" for m in FusedAttnBackend} == set(NVTE_Fused_Attn_Backend.__members__), ( +assert { + f"NVTE_{m.name}" for m in FusedAttnBackend if m not in _PYTHON_ONLY_FUSED_ATTN_BACKENDS +} == set(NVTE_Fused_Attn_Backend.__members__), ( "FusedAttnBackend in python is out of sync with" " transformer_engine_torch.NVTE_Fused_Attn_Backend defined on the C++ side." " Please make sure TE C++ and python are in sync." @@ -345,6 +354,46 @@ def fused_attn_fwd( # Accept the pybind enum for backward compatibility. fused_attention_backend = FusedAttnBackend.cast(fused_attention_backend) + if fused_attention_backend == FusedAttnBackend["FROST"]: + # FROST runs through the cuDNN Frontend python API rather than the C++ fused path. + # 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 + + return frost_attention.fused_attn_fwd( + is_training, + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q, + k, + v, + fake_dtype, + fused_attention_backend, + attn_bias=attn_bias, + cu_seqlens_q_padded=cu_seqlens_q_padded, + cu_seqlens_kv_padded=cu_seqlens_kv_padded, + page_table_k=page_table_k, + page_table_v=page_table_v, + s_quantizer=s_quantizer, + o_quantizer=o_quantizer, + attn_scale=attn_scale, + dropout=dropout, + fast_zero_fill=fast_zero_fill, + qkv_layout=qkv_layout, + o_format=o_format, + qkv_scale_inv_format=qkv_scale_inv_format, + attn_bias_type=attn_bias_type, + attn_mask_type=attn_mask_type, + softmax_type=softmax_type, + window_size=window_size, + bottom_right_diagonal=bottom_right_diagonal, + rng_gen=rng_gen, + softmax_offset=softmax_offset, + return_max_logit=return_max_logit, + cuda_graph=cuda_graph, + ) if fused_attention_backend == FusedAttnBackend["No_Backend"]: raise ValueError( "Fused attention does not support this input combination:" @@ -600,6 +649,46 @@ def fused_attn_bwd( # Accept the pybind enum for backward compatibility. fused_attention_backend = FusedAttnBackend.cast(fused_attention_backend) + if fused_attention_backend == FusedAttnBackend["FROST"]: + # See the matching branch in fused_attn_fwd. + # pylint: disable-next=import-outside-toplevel + from ..attention.dot_product_attention import frost_attention + + return frost_attention.fused_attn_bwd( + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q, + k, + v, + o, + d_o, + fake_dtype, + aux_ctx_tensors, + fused_attention_backend, + cu_seqlens_q_padded=cu_seqlens_q_padded, + cu_seqlens_kv_padded=cu_seqlens_kv_padded, + s_quantizer=s_quantizer, + dp_quantizer=dp_quantizer, + dqkv_quantizer=dqkv_quantizer, + attn_scale=attn_scale, + dropout=dropout, + fast_zero_fill=fast_zero_fill, + qkv_layout=qkv_layout, + o_format=o_format, + do_format=do_format, + dqkv_layout=dqkv_layout, + qkv_scale_inv_format=qkv_scale_inv_format, + do_scale_inv_format=do_scale_inv_format, + attn_bias_type=attn_bias_type, + attn_mask_type=attn_mask_type, + softmax_type=softmax_type, + window_size=window_size, + bottom_right_diagonal=bottom_right_diagonal, + deterministic=deterministic, + cuda_graph=cuda_graph, + ) if fused_attention_backend == FusedAttnBackend["No_Backend"]: raise ValueError( "Fused attention backward does not support this input combination:" From 9eb866d81cb460b15eb3830c29a9b7cb3dcd4543 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Tue, 6 Oct 2026 01:00:38 +0000 Subject: [PATCH 50/97] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/pytorch/attention/run_attention_with_cp.py | 1 + .../pytorch/attention/dot_product_attention/backends.py | 7 +++---- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/tests/pytorch/attention/run_attention_with_cp.py b/tests/pytorch/attention/run_attention_with_cp.py index 342084474c3..af5eea21722 100644 --- a/tests/pytorch/attention/run_attention_with_cp.py +++ b/tests/pytorch/attention/run_attention_with_cp.py @@ -615,6 +615,7 @@ def run_dpa_with_cp( from transformer_engine.pytorch.attention.dot_product_attention.dot_product_attention import ( # pylint: disable=import-outside-toplevel _attention_backends, ) + # pylint: disable-next=import-outside-toplevel from transformer_engine.pytorch.cpp_extensions.fused_attn import FusedAttnBackend diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index 3a81b5142bc..6e5c331b902 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -2609,10 +2609,9 @@ def forward( ) if context_parallel: - assert ( - fp8 - or fused_attention_backend - in (FusedAttnBackend["F16_arbitrary_seqlen"], FusedAttnBackend["FROST"]) + assert fp8 or fused_attention_backend in ( + FusedAttnBackend["F16_arbitrary_seqlen"], + FusedAttnBackend["FROST"], ), f"{fused_attention_backend} does not work with context parallelism!" assert core_attention_bias_type not in [ "alibi" From 0d1c4bf5352562ef389c73ca5504e8a26069d521 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Mon, 5 Oct 2026 21:39:12 -0700 Subject: [PATCH 51/97] fix(attention): decline FROST where the diagonal anchor is ambiguous A right-bounded window on a non-causal mask takes its alignment only from bottom_right_diagonal, which defaults to top-left, while the all-gather ring trims KV and measures its window against the bottom-right diagonal. Those differ exactly when max_seqlen_q != max_seqlen_kv, so decline instead of guessing. The selector carried this decline before FROST became a sub-backend; it was dropped when the checks moved onto FusedAttentionParams. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_frost_attention.py | 31 +++++++++++++++++++ .../dot_product_attention/frost_attention.py | 15 ++++++++- 2 files changed, 45 insertions(+), 1 deletion(-) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 45c301f703d..76686de768d 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -330,6 +330,19 @@ def test_frost_declines_unsupported_configs(): # selector contract rather than an internal detail. (dict(window_size_right=5), "a right window past the diagonal"), (dict(window_size_left=-2, window_size_right=0), "a left window below -1"), + # A right-bounded window on a non-causal mask takes its anchor only from + # bottom_right_diagonal; the all-gather ring measures its window bottom-right. Those + # differ exactly when the lengths do. + ( + dict( + attn_mask_type=AttnMaskType["no_mask"], + window_size_left=128, + window_size_right=0, + max_seqlen_q=512, + max_seqlen_kv=640, + ), + "an ambiguous diagonal anchor", + ), # The engine pads head_dim to a multiple of 8, so an in-range but unpadded dim has to be # declined here rather than failing later at plan selection. (dict(head_dim_qk=260, head_dim_v=260), "head_dim not a multiple of 8"), @@ -339,6 +352,24 @@ def test_frost_declines_unsupported_configs(): assert reason, "a decline must explain itself" +@requires_frost +def test_frost_serves_an_unambiguous_window_on_a_non_causal_mask(): + """The anchor is only ambiguous when the q and kv lengths differ; equal lengths must serve.""" + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + is_frost_attention_supported, + ) + from transformer_engine.pytorch.cpp_extensions.fused_attn import AttnMaskType, FusedAttnBackend + + params = _frost_params( + attn_mask_type=AttnMaskType["no_mask"], + window_size_left=128, + window_size_right=0, + max_seqlen_q=4096, + max_seqlen_kv=4096, + ) + assert is_frost_attention_supported(params)[0] == FusedAttnBackend.FROST + + @requires_frost def test_frost_mask_spec_rejects_malformed_windows(): """_mask_spec is the only validation between a caller-supplied window and a built band.""" diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index f1956aa2aec..96cb07544cf 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -370,13 +370,26 @@ def is_frost_attention_supported(params) -> Tuple[int, str]: if attn_mask_type is None: return no_backend, f"FROST got an unrecognised attn_mask_type {params.attn_mask_type}" try: - _te_mask_spec( + mask_for_band, _ = _te_mask_spec( attn_mask_type, (params.window_size_left, params.window_size_right), params.bottom_right_diagonal, ) except NotImplementedError as exc: return no_backend, str(exc) + if ( + mask_for_band == "causal" + and "causal" not in attn_mask_type + and params.max_seqlen_q != params.max_seqlen_kv + ): + # A right-bounded window on a non-causal mask takes its anchor only from + # bottom_right_diagonal, which defaults to top-left, while the all-gather ring trims KV + # and measures its window against the bottom-right diagonal. Those differ exactly when + # the q and kv lengths do, so decline rather than guess which one was meant. + return no_backend, ( + "FROST declines a right-bounded window on a non-causal mask with max_seqlen_q !=" + " max_seqlen_kv, where the diagonal anchor is ambiguous" + ) return int(FusedAttnBackend.FROST), "" From 91a37f17023a76537bc70c06ff15bbf61a91c6c8 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Tue, 6 Oct 2026 04:41:10 +0000 Subject: [PATCH 52/97] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- .../attention/dot_product_attention/frost_attention.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 96cb07544cf..a4a51662ccd 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -386,9 +386,12 @@ def is_frost_attention_supported(params) -> Tuple[int, str]: # bottom_right_diagonal, which defaults to top-left, while the all-gather ring trims KV # and measures its window against the bottom-right diagonal. Those differ exactly when # the q and kv lengths do, so decline rather than guess which one was meant. - return no_backend, ( - "FROST declines a right-bounded window on a non-causal mask with max_seqlen_q !=" - " max_seqlen_kv, where the diagonal anchor is ambiguous" + return ( + no_backend, + ( + "FROST declines a right-bounded window on a non-causal mask with max_seqlen_q !=" + " max_seqlen_kv, where the diagonal anchor is ambiguous" + ), ) return int(FusedAttnBackend.FROST), "" From 521c9a8480ae2e6e75e3f5af102897c1930c4175 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Mon, 5 Oct 2026 21:49:24 -0700 Subject: [PATCH 53/97] refactor(attention): make frost_attention.py standalone again Defers the cudnn_pygraph.py extraction, as asked: get FROST working as a sub-backend first, then work out the duplication against flex_attention.py in a follow-up. flex_attention.py and its tests are back to main, cudnn_pygraph.py is removed, and the helpers frost needs (the cuDNN import, the per-device handle, the graph builder, the diagonal band and plan finalization) are private to frost_attention.py. finalize_plans drops exclude_plan_tokens, which only flex needed. The separate fix barring the FROST engines from flex's score_mod graphs moves to its own branch; it stands on its own whether or not FROST lands. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_flex_attention.py | 192 ------------ .../pytorch/attention/test_frost_attention.py | 50 +-- .../dot_product_attention/cudnn_pygraph.py | 284 ------------------ .../dot_product_attention/flex_attention.py | 172 ++++++----- .../dot_product_attention/frost_attention.py | 212 +++++++++++-- 5 files changed, 302 insertions(+), 608 deletions(-) delete mode 100644 transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index 42236812b28..beed4069917 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -705,195 +705,3 @@ def test_dot_product_attention_score_mod(dtype, qkv_format, score_mod_case, scal torch.testing.assert_close(q.grad, q_ref.grad, **tols) torch.testing.assert_close(k.grad, k_ref.grad, **tols) torch.testing.assert_close(v.grad, v_ref.grad, **tols) - - -@pytest.mark.parametrize("direction", ["fwd", "bwd"]) -def test_score_mod_graph_signatures_stay_aligned(direction): - """The cache key, the builder and the getter are splatted from one positional tuple. - - `_get_cudnn_score_mod_*_graph` passes the same `build_args` tuple to the cache key and to the - builder, so the three parameter lists have to stay in the same order. Nothing enforced that, - and the failure is quiet in the worst direction: a parameter inserted in one signature and not - another shifts the rest by one, and a shifted *cache key* is not a crash, it is two different - configurations sharing a cached graph. - - No GPU: this reads signatures only. - """ - import inspect - - names = [ - getattr(flex_attention, "_cudnn_score_mod_%s_cache_key" % direction), - getattr(flex_attention, "_build_cudnn_score_mod_%s_graph" % direction), - getattr(flex_attention, "_get_cudnn_score_mod_%s_graph" % direction), - ] - signatures = [list(inspect.signature(fn).parameters) for fn in names] - reference = signatures[0] - for fn, params in zip(names[1:], signatures[1:]): - assert params == reference, ( - "%s takes %s but _cudnn_score_mod_%s_cache_key takes %s; these are splatted from one" - " positional tuple and must stay in the same order" - % (fn.__name__, params, direction, reference) - ) - - -@pytest.mark.parametrize( - "mask_spec,expected", - [ - # Causal: top-left aligned, right bound pinned to the diagonal, no left bound. - (("causal", (-1, 0)), {"diagonal_alignment": "TOP_LEFT", "diagonal_band_right_bound": 0}), - # Bottom-right causal, which is what KV trimming produces whenever SKV > SQ. - ( - ("causal_bottom_right", (-1, 0)), - {"diagonal_alignment": "BOTTOM_RIGHT", "diagonal_band_right_bound": 0}, - ), - # Sliding window. cuDNN's left bound counts the diagonal itself and TE's window_size does - # not, so 511 must arrive as 512. Getting this wrong drops one token of context per layer - # and no shape-level test would notice. - ( - ("causal", (511, 0)), - { - "diagonal_alignment": "TOP_LEFT", - "diagonal_band_right_bound": 0, - "diagonal_band_left_bound": 512, - }, - ), - # No mask at all: no alignment, no bounds. - (("no_mask", (-1, -1)), {}), - ], -) -def test_mask_spec_translates_to_a_diagonal_band(mask_spec, expected): - """A mask_spec must become the cuDNN band kwargs, and never a score_mod. - - No GPU: this builds no graph, it checks the kwargs the graph would be given. The frontend is - an optional dependency, so skip rather than fail where it is absent. - """ - try: - cudnn = flex_attention._import_cudnn_frontend() - except ImportError: - pytest.skip("cuDNN frontend Python package is required for the diagonal-band kwargs.") - got = flex_attention._mask_or_score_mod_kwargs(mask_spec, None) - - assert "score_mod" not in got and "use_causal_mask" not in got - for key, want in expected.items(): - if key == "diagonal_alignment": - assert got[key] == getattr(cudnn.diagonal_alignment, want) - else: - assert got[key] == want - assert set(got) == set(expected) - - -def test_mask_spec_and_score_mod_cannot_be_combined(): - """cuDNN's backward refuses the pair, so flex must refuse it before building the graph.""" - with pytest.raises(ValueError, match="cannot be combined"): - flex_attention._mask_or_score_mod_kwargs(("causal", (-1, 0)), lambda *a, **k: None) - - -def test_no_mask_spec_still_takes_the_score_mod_path(): - """The default path must be byte-identical to what it was before mask_spec existed.""" - sentinel = object() - assert flex_attention._mask_or_score_mod_kwargs(None, sentinel) == { - "use_causal_mask": False, - "score_mod": sentinel, - } - - -def test_flex_bars_the_frost_engines(): - """flex must tell cuDNN not to use a FROST engine, not merely decline to ask for them. - - The switch that offers those engines is process-wide, so a FrostAttention call elsewhere in the - process, or a user setting CUDNN_FRONTEND_ENABLE_FROST_ENGINES, puts them ahead of the backend - engines for these graphs too. They accept a score_mod graph, pass check_support, build, and - then compute without the callback. - - No GPU: this checks the instruction is passed, not what cuDNN does with it. - """ - from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph - - seen = {} - - def fake_finalize(graph, **kwargs): - seen.update(kwargs) - return 4096, None - - original = cudnn_pygraph.finalize_plans - cudnn_pygraph.finalize_plans = fake_finalize - try: - assert flex_attention._finalize_cudnn_graph(object()) == 4096 - finally: - cudnn_pygraph.finalize_plans = original - - excluded = seen.get("exclude_plan_tokens") - assert excluded, "flex did not ask cuDNN to exclude any engine" - assert "sdpa_fwd_prefill_sm100" in excluded and "sdpa_bwd_sm100" in excluded, excluded - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required.") -def test_frost_switch_does_not_change_what_flex_computes(): - """Enabling the FROST engines must not change flex's output. - - This is the property the silent drop violated: with the engines on, an unpinned build selected - a FROST plan at every head dim measured on B200, and that plan returns plain attention with the - score_mod discarded. Comparing flex against itself across the switch needs no reference and no - knowledge of which plan ran; if the two differ, a different kernel answered. - """ - try: - flex_attention._import_cudnn_frontend() - except ImportError: - pytest.skip("cuDNN frontend Python package is required for score_mod attention.") - - # Without this the test is vacuous nearly everywhere: where the FROST engines are absent or - # decline on arch, both runs get a backend plan and agree no matter what flex does. The - # engines themselves are found lazily at planning time, so the switch works whenever it is - # set, but they still have to exist. - from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( - is_frost_attention_available, - ) - - frost_ok, frost_reason = is_frost_attention_available() - if not frost_ok: - pytest.skip( - "the FROST engines must be reachable for this to test anything: %s" % frost_reason - ) - - env = "CUDNN_FRONTEND_ENABLE_FROST_ENGINES" - saved = os.environ.get(env) - torch.manual_seed(0) - b, h, s, d = 2, 4, 512, 64 - dtype = torch.bfloat16 if is_bf16_available() else torch.float16 - q, k, v = (torch.randn(b, s, h, d, device="cuda", dtype=dtype) for _ in range(3)) - - def bias_score_mod(score_mod_graph, score_tensor, _tensors): - """score += (row - col). Self-contained, and large enough that dropping it is obvious.""" - cudnn = flex_attention._import_cudnn_frontend() - row = score_mod_graph.gen_index(input=score_tensor, axis=2) - row.set_data_type(cudnn.data_type.INT32) - col = score_mod_graph.gen_index(input=score_tensor, axis=3) - col.set_data_type(cudnn.data_type.INT32) - bias = score_mod_graph.sub(a=row, b=col, compute_data_type=cudnn.data_type.FLOAT) - bias.set_data_type(cudnn.data_type.FLOAT) - return score_mod_graph.add(a=score_tensor, b=bias, compute_data_type=cudnn.data_type.FLOAT) - - def run(): - flex_attention._cudnn_score_mod_graph_cache.clear() - return flex_attention.FusedAttentionWithScoreModFunc.apply( - False, q, k, v, "bshd", "bshd", d**-0.5, bias_score_mod, None, None, None, False - ) - - try: - os.environ.pop(env, None) - without = run() - os.environ[env] = "1" - with_engines = run() - finally: - flex_attention._cudnn_score_mod_graph_cache.clear() - if saved is None: - os.environ.pop(env, None) - else: - os.environ[env] = saved - - torch.testing.assert_close( - with_engines, - without, - msg=lambda m: "flex computed something different with the FROST engines enabled, which means a FROST plan answered and dropped the score_mod:\n" - + m, - ) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 76686de768d..395f3f6beb1 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -468,15 +468,16 @@ def test_frost_rejects_mismatched_kv(): frost_attn_fwd(q, k, k.to(torch.float32)) -def test_frost_engines_are_enabled_even_if_flex_imported_cudnn_first(): - """Enabling the FROST engines must not depend on which backend touched cuDNN first. +def test_frost_engines_are_enabled_even_if_cudnn_was_imported_without_them(): + """Enabling the FROST engines must not depend on who imported cuDNN first. - flex_attention and frost_attention share one cuDNN import in cudnn_pygraph. flex asks for the - import without the FROST engines and frost asks with them, so if the enabling sat inside the - "already imported?" memo, a process that ran a score_mod layer first would leave FROST with a - cuDNN that offers it no engine. That surfaces far from its cause, as "no cuDNN engine matching - 'sdpa_fwd_prefill_sm100' was offered" on the first head_dim 512 forward, with a hint pointing - at package versions that are in fact fine. + is_frost_attention_available imports cuDNN WITHOUT the engines, because enabling them reorders + plan selection for every cuDNN consumer in the process and the checks after it may still + decline. So the enabling cannot sit inside the "already imported?" memo: a process that + probed availability first would otherwise leave FROST with a cuDNN that offers it no engine. + That surfaces far from its cause, as "no cuDNN engine matching 'sdpa_fwd_prefill_sm100' was + offered" on the first head_dim 512 forward, with a hint pointing at package versions that are + in fact fine. No GPU and no real cuDNN: a stub stands in for the package, because what is under test is the order-dependence of our own wrapper. It also has to run in-process with the globals reset, @@ -485,12 +486,12 @@ def test_frost_engines_are_enabled_even_if_flex_imported_cudnn_first(): import sys import types - from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph + from transformer_engine.pytorch.attention.dot_product_attention import frost_attention env = "CUDNN_FRONTEND_ENABLE_FROST_ENGINES" saved = ( - cudnn_pygraph._cudnn, - cudnn_pygraph._frost_engines_enabled, + frost_attention._cudnn, + frost_attention._frost_engines_enabled, os.environ.get(env), sys.modules.get("cudnn"), sys.modules.get("cudnn.sdpa"), @@ -500,23 +501,24 @@ def test_frost_engines_are_enabled_even_if_flex_imported_cudnn_first(): stub.sdpa = types.ModuleType("cudnn.sdpa") sys.modules["cudnn"] = stub sys.modules["cudnn.sdpa"] = stub.sdpa - cudnn_pygraph._cudnn = None - cudnn_pygraph._frost_engines_enabled = False + frost_attention._cudnn = None + frost_attention._frost_engines_enabled = False os.environ.pop(env, None) - # flex first, which must not enable anything. - cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=False) + # The availability probe first, which must not enable anything. + frost_attention._import_cudnn_frontend(enable_frost_engines=False) assert env not in os.environ, "the non-FROST caller must not set the switch" - assert not cudnn_pygraph.frost_engines_enabled() + assert not frost_attention._frost_engines_enabled - # frost second, on an already-imported cuDNN. This is the case that used to be skipped. - cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=True) + # A use site second, on an already-imported cuDNN. This is the case that used to be + # skipped. + frost_attention._import_cudnn_frontend(enable_frost_engines=True) assert os.environ.get(env) == "1", "FROST was requested after the import and not enabled" - assert cudnn_pygraph.frost_engines_enabled() + assert frost_attention._frost_engines_enabled finally: ( - cudnn_pygraph._cudnn, - cudnn_pygraph._frost_engines_enabled, + frost_attention._cudnn, + frost_attention._frost_engines_enabled, prior_env, prior_cudnn, prior_sdpa, @@ -543,10 +545,10 @@ def test_pinned_plan_decline_reports_the_engine_reason(): No GPU: a stub graph stands in, raising the real cuDNN exception type. """ - from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph + from transformer_engine.pytorch.attention.dot_product_attention import frost_attention try: - cudnn = cudnn_pygraph.import_cudnn_frontend() + cudnn = frost_attention._import_cudnn_frontend() except ImportError: pytest.skip("cuDNN frontend Python package is required for the decline-reason path.") @@ -575,7 +577,7 @@ def check_support(self): raise cudnn.cudnnGraphNotSupportedError("head_dim 512 needs SM100; this is SM90") with pytest.raises(RuntimeError) as excinfo: - cudnn_pygraph.finalize_plans( + frost_attention._finalize_plans( _DeclinedGraph(), heuristics=[cudnn.heur_mode.A], require_plan_token="sdpa_fwd_prefill_sm100", diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py deleted file mode 100644 index 7a07a08484c..00000000000 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ /dev/null @@ -1,284 +0,0 @@ -# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# -# See LICENSE for license information. - -"""Shared cuDNN Frontend Python-graph plumbing. - -Two attention backends build cuDNN graphs from Python: flex_attention.py, for score_mod, and -frost_attention.py, for the CuTe-DSL SDPA kernels at head_dim in (256, 512]. They differ in the -SDPA node they build, which cannot be shared because cuDNN treats a score_mod and a diagonal band -as mutually exclusive, but everything around that node is the same work: importing the frontend, -holding one handle per device on PyTorch's current stream, describing an SBHD/BSHD tensor in the -BHSD form cuDNN wants, finalizing plans, and executing. - -This module is that common part: the plumbing, plus the one piece of shared attention -vocabulary, translating a TE mask type and window into cuDNN's diagonal band. -""" - -from typing import Any, Dict, Optional, Sequence, Tuple - -import os - -import torch - - -_cudnn = None -_frost_engines_enabled = False -_handles: Dict[torch.device, Any] = {} - - -def import_cudnn_frontend(enable_frost_engines: bool = False): - """Import cuDNN Frontend, enabling the FROST engines if this caller needs them. - - ``enable_frost_engines`` is not merely additive: the switch also ranks FROST ahead of the - backend engines everywhere, so a caller that does not want FROST must not ask for it. - - The enabling is deliberately outside the import memo. Both backends call this, and whichever - one reaches it first would otherwise decide for the process: with the flag inside the memo, a - flex call would cache the module with FROST off and every later FROST call would get a cuDNN - that offers no FROST engine, which surfaces much later as "no cuDNN engine matching ... was - offered". Enabling late is sound because the switch is read per graph rather than at import: - in cuDNN Frontend 1.29.0 ``engines/manifest.py`` consults the environment inside - ``offered_ids()``, reached from ``engines_for(graph)`` on every ``create_execution_plans``. - - Note the switch is process-wide and never unset, so enabling it for FROST also reorders the - candidates a concurrent score_mod graph sees. Callers that require a particular engine should - verify by plan name rather than rely on the switch, which is what - ``finalize_plans(require_plan_token=...)`` does. - """ - global _cudnn, _frost_engines_enabled # pylint: disable=global-statement - if _cudnn is None: - try: - import cudnn # pylint: disable=import-outside-toplevel - except ImportError as exc: - raise ImportError( - "cuDNN frontend Python package not found. " - "Install it with: pip install nvidia-cudnn-frontend" - ) from exc - - _cudnn = cudnn - - if enable_frost_engines and not _frost_engines_enabled: - os.environ.setdefault("CUDNN_FRONTEND_ENABLE_FROST_ENGINES", "1") - # pylint: disable=import-outside-toplevel,unused-import - import cudnn.sdpa # noqa: F401 - - _frost_engines_enabled = True - - return _cudnn - - -def frost_engines_enabled() -> bool: - """Whether this process has enabled the FROST engines through ``import_cudnn_frontend``.""" - return _frost_engines_enabled - - -def handle_for(device: torch.device, *, backend_name: str = "cuDNN attention"): - """A cuDNN handle for ``device``, rebound to PyTorch's current stream on every call. - - Without the rebinding, cuDNN runs on its handle's own stream while the tensors and workspace - are allocated on PyTorch's current stream, and nothing orders the two. That is not - hypothetical: the p2p context-parallel ring issues attention inside - ``with torch.cuda.stream(cp_stream)``, so on alternating ring steps the kernel and its buffers - would otherwise be on different streams. The same cached plan is executed from different - streams across steps, so this has to happen per call rather than once per handle. - """ - if device.type != "cuda": - raise ValueError(f"{backend_name} requires CUDA tensors; got device {device}") - cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() - if device.index is None: - device = torch.device("cuda", torch.cuda.current_device()) - with torch.cuda.device(device): - handle = _handles.get(device) - if handle is None: - handle = cudnn.create_handle() - _handles[device] = handle - cudnn.set_stream(handle=handle, stream=torch.cuda.current_stream(device).cuda_stream) - return handle - - -def io_data_type(cudnn, dtype: torch.dtype, *, backend_name: str = "cuDNN attention"): - """Map a torch dtype to the cuDNN frontend enum, for the dtypes these backends accept.""" - if dtype == torch.float16: - return cudnn.data_type.HALF - if dtype == torch.bfloat16: - return cudnn.data_type.BFLOAT16 - raise ValueError(f"{backend_name} only supports FP16/BF16 tensors, got {dtype}") - - -def build_pygraph( - dtype: torch.dtype, device: torch.device, *, backend_name: str = "cuDNN attention" -): - """A cuDNN frontend graph for F16/BF16 SDPA, bound to this device's stream-current handle.""" - cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() - return cudnn.pygraph( - io_data_type=io_data_type(cudnn, dtype, backend_name=backend_name), - intermediate_data_type=cudnn.data_type.FLOAT, - compute_data_type=cudnn.data_type.FLOAT, - handle=handle_for(device, backend_name=backend_name), - ) - - -def bhsd_dim_stride( - tensor: torch.Tensor, tensor_format: str -) -> Tuple[Tuple[int, ...], Tuple[int, ...]]: - """Describe an SBHD/BSHD tensor as cuDNN frontend's logical BHSD form. - - No copy and no permute: the strides are handed to cuDNN as they are, which is what lets both - layouts be served directly. sbhd matters because that is what Megatron uses internally. - """ - if tensor_format == "sbhd": - return ( - (tensor.shape[1], tensor.shape[2], tensor.shape[0], tensor.shape[3]), - (tensor.stride(1), tensor.stride(2), tensor.stride(0), tensor.stride(3)), - ) - if tensor_format == "bshd": - return ( - (tensor.shape[0], tensor.shape[2], tensor.shape[1], tensor.shape[3]), - (tensor.stride(0), tensor.stride(2), tensor.stride(1), tensor.stride(3)), - ) - raise ValueError(f"Only SBHD/BSHD tensor formats are supported, got {tensor_format}.") - - -def bhsd_graph_tensor(graph, tensor: torch.Tensor, tensor_format: str): - """Create a cuDNN graph tensor with BHSD dims and the tensor's own strides.""" - dim, stride = bhsd_dim_stride(tensor, tensor_format) - return graph.tensor(dim=dim, stride=stride, data_type=tensor.dtype) - - -def diagonal_band_kwargs(cudnn, attn_mask_type: str, window: Tuple[int, int]) -> Dict[str, Any]: - """cuDNN sdpa kwargs for a TE (mask type, window): a diagonal alignment plus a band. - - Note the off-by-one. cuDNN's left bound counts the diagonal itself and TE's window_size does - not, so a window of w becomes a left bound of w + 1. Passing it through unconverted silently - drops one token of context per layer, which no shape-level test would catch. - - These kwargs are mutually exclusive with score_mod. cuDNN enforces that in the backward node - only ("Attention score mod enabled and hence other subgraphs are disabled"); its forward node - composes the two without complaint. Callers must still refuse the pair on both sides, because - forward and backward have to carry the same mask or the gradients belong to a different - attention than the output does. - """ - left, right = window - opts: Dict[str, Any] = {} - if attn_mask_type in ("causal", "causal_bottom_right") or right == 0: - opts["diagonal_alignment"] = ( - cudnn.diagonal_alignment.BOTTOM_RIGHT - if attn_mask_type == "causal_bottom_right" - else cudnn.diagonal_alignment.TOP_LEFT - ) - opts["diagonal_band_right_bound"] = 0 - if left != -1: - opts["diagonal_band_left_bound"] = left + 1 - return opts - - -def finalize_plans( - graph, - *, - heuristics: Optional[Sequence[Any]] = None, - build_policy: Any = None, - require_plan_token: Optional[str] = None, - not_found_hint: Any = "", - exclude_plan_tokens: Optional[Sequence[str]] = None, -) -> Tuple[int, Optional[str]]: - """Create plans, optionally pin one by name, build, and return (workspace size, plan name). - - ``require_plan_token`` makes the choice strict: only a plan whose name contains the token is - acceptable, and anything else raises. That is not a stylistic preference. Without a pin, - ``build_plans`` walks the ranked list from index 0 and finalizes the first plan that builds, - logging each decline at INFO, so a graph that the intended engine declines runs on whatever - cuDNN ranked next with nothing in the return value to say so. At head_dim 512 that matters in - the forward, where an ordinary engine may well build and compute a different function from the - FROST kernel. The backward is self-limiting, since no non-FROST d512 backward exists, so an - unpinned backward would fail loudly on its own. - - The token is matched as a substring rather than by equality on purpose: cuDNN has already - collapsed per-head-dim engine names (``..._d512`` and friends) into a single row once, and the - substring test survived that. - - ``exclude_plan_tokens`` is the opposite instruction, for a caller that must NOT run on a - particular engine. It is needed because the FROST engine switch is process-wide: a caller that - declines to ask for those engines still gets them ranked first once anything else in the - process has enabled them. Measured on B200 at head_dim 64, 128, 256 and 512, a FROST plan - ranks at index 0 for a score_mod graph and an unpinned build selects it every time. - - Pinning also changes what ``check_support`` means. Selecting a plan sets cuDNN's internal - ``_plan_pinned``, and only then is a decline fatal; unpinned, cuDNN records the decline and - keeps walking. So the pin has to come first both because the check is scoped to the selected - plan and because it is what makes the check binding at all. - """ - cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() - - graph.validate() - graph.build_operation_graph() - - if heuristics is None: - heuristics = [cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK] - - if require_plan_token is None: - try: - graph.create_execution_plans(list(heuristics)) - if exclude_plan_tokens: - # Bar the named engines before the walk, so build_plans falls through to the first - # entry that is both unbarred and buildable. Inert when those engines are not on - # offer, which is every process that has not enabled them. - graph.deselect_engines(list(exclude_plan_tokens)) - graph.check_support() - except cudnn.cudnnGraphNotSupportedError as exc: - raise RuntimeError(f"cuDNN SDPA graph is not supported: {exc}") from exc - if build_policy is None: - build_policy = cudnn.build_plan_policy.HEURISTICS_CHOICE - graph.build_plans(build_policy) - return max(graph.get_workspace_size(), 1), None - - graph.create_execution_plans(list(heuristics)) - names = [graph.get_plan_name_at_index(i) for i in range(graph.get_execution_plan_count())] - hits = [i for i, n in enumerate(names) if require_plan_token in n] - if not hits: - # Callable hints are resolved only here: a caller may want to look up package versions to - # explain the failure, and that work should not happen on the success path. - hint = not_found_hint() if callable(not_found_hint) else not_found_hint - raise RuntimeError( - f"no cuDNN engine matching {require_plan_token!r} was offered." - f" Candidate plans: {names[:6]}.{(' ' + hint) if hint else ''}" - ) - graph.select_plan(hits[0]) - # The engine is pinned, so a decline here is the engine's own verdict on this graph and cuDNN - # puts its reason in the exception. Surface that rather than letting it escape bare: a plan - # that was offered and then refused is the harder failure to read, and the reason is the only - # thing that says which constraint was missed. - try: - graph.check_support() - graph.build_plans() - except cudnn.cudnnGraphNotSupportedError as exc: - hint = not_found_hint() if callable(not_found_hint) else not_found_hint - raise RuntimeError( - f"cuDNN engine {names[hits[0]]!r} was offered but declined this graph:" - f" {exc}{(' ' + hint) if hint else ''}" - ) from exc - return max(graph.get_workspace_size(), 1), names[hits[0]] - - -def selected_plan_name(graph, index: int = 0) -> str: - """Name of the plan at ``index``, for logging and for asserting which engine answered.""" - return graph.get_plan_name_at_index(index) - - -def execute_graph( - graph, - variant_pack: Dict[Any, torch.Tensor], - workspace_size: int, - device: torch.device, - *, - backend_name: str = "cuDNN attention", -): - """Execute a built graph on this device's stream-current handle.""" - if device.type == "cuda" and device.index is None: - device = torch.device("cuda", torch.cuda.current_device()) - workspace = torch.empty(workspace_size, device=device, dtype=torch.uint8) - graph.execute( - variant_pack, - workspace, - handle=handle_for(device, backend_name=backend_name), - ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index 6df8655f0a1..b9593b42d9b 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -5,42 +5,52 @@ """cuDNN-backed Flex Attention helpers.""" from dataclasses import dataclass +import importlib import inspect from typing import Any, Callable, Dict, Optional, Tuple import torch -from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph - -# The handle cache lives in cudnn_pygraph now; the alias keeps the old name working. -_cudnn_score_mod_handles = cudnn_pygraph._handles # pylint: disable=protected-access +_cudnn_score_mod_handles: Dict[torch.device, Any] = {} _cudnn_score_mod_graph_cache: Dict[Tuple[Any, ...], Any] = {} _SCORE_MOD_UNCACHEABLE = object() -_BACKEND = "Flex Attention" - def _import_cudnn_frontend(): """Import the cuDNN frontend Python package.""" - # This path does not ask for the FROST engines, but asking is all it controls: the switch is - # process-wide, so a FrostAttention call elsewhere in the process, or a user setting - # CUDNN_FRONTEND_ENABLE_FROST_ENGINES themselves, still ranks FROST ahead of the backend - # engines for the graphs built here. - return cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=False) + try: + return importlib.import_module("cudnn") + except ImportError as exc: + raise ImportError( + "cuDNN frontend Python package not found. " + "Install it with: pip install nvidia-cudnn-frontend" + ) from exc def _bhsd_dim_stride( tensor: torch.Tensor, tensor_format: str ) -> Tuple[Tuple[int, ...], Tuple[int, ...]]: """Describe an SBHD/BSHD tensor as cuDNN frontend's logical BHSD format.""" - return cudnn_pygraph.bhsd_dim_stride(tensor, tensor_format) + if tensor_format == "sbhd": + return ( + (tensor.shape[1], tensor.shape[2], tensor.shape[0], tensor.shape[3]), + (tensor.stride(1), tensor.stride(2), tensor.stride(0), tensor.stride(3)), + ) + if tensor_format == "bshd": + return ( + (tensor.shape[0], tensor.shape[2], tensor.shape[1], tensor.shape[3]), + (tensor.stride(0), tensor.stride(2), tensor.stride(1), tensor.stride(3)), + ) + raise ValueError(f"Flex Attention only supports SBHD/BSHD tensor formats, got {tensor_format}.") def _bhsd_graph_tensor(graph, tensor: torch.Tensor, tensor_format: str): """Create a cuDNN graph tensor with BHSD dims and TE-layout strides.""" - return cudnn_pygraph.bhsd_graph_tensor(graph, tensor, tensor_format) + dim, stride = _bhsd_dim_stride(tensor, tensor_format) + return graph.tensor(dim=dim, stride=stride, data_type=tensor.dtype) +# score_mod graph cache helpers. def _freeze_score_mod_cache_key(value: Any) -> Any: """Convert a user-provided score_mod graph key into a hashable structure.""" if isinstance(value, torch.Tensor): @@ -163,27 +173,6 @@ def _score_mod_bhsd_tensor_metadata(tensor: torch.Tensor, tensor_format: str) -> return (dim, stride, tensor.dtype, _score_mod_device_key(tensor.device)) -def _mask_or_score_mod_kwargs( - mask_spec: Optional[Tuple[str, Tuple[int, int]]], wrapped_score_mod -) -> Dict[str, Any]: - """SDPA kwargs for exactly one of a diagonal band or a score_mod. - - cuDNN rejects the pair in its backward node ("Attention score mod enabled and hence other - subgraphs are disabled") while its forward node composes both silently. Refusing it here on - both sides is deliberate: the backward has to carry the same mask as the forward, or the - gradients belong to a different attention than the output does. - """ - if mask_spec is None: - return {"use_causal_mask": False, "score_mod": wrapped_score_mod} - if wrapped_score_mod is not None: - raise ValueError( - "a diagonal-band mask and a score_mod cannot be combined in one cuDNN SDPA graph; " - f"got mask_spec={mask_spec!r} alongside a score_mod" - ) - cudnn = _import_cudnn_frontend() - return cudnn_pygraph.diagonal_band_kwargs(cudnn, mask_spec[0], mask_spec[1]) - - def _make_cudnn_graph_tensor_dict(graph, tensors: Optional[Dict[str, torch.Tensor]]): """Create cuDNN graph tensors matching runtime tensors.""" if tensors is None: @@ -205,13 +194,40 @@ def _wrapped_score_mod(sdpa_graph, score_tensor): def _get_cudnn_current_stream_handle(cudnn, device: torch.device): """Return a cuDNN handle for device, bound to PyTorch's current stream.""" - del cudnn # the shared helper resolves the module itself - return cudnn_pygraph.handle_for(device, backend_name=_BACKEND) + if device.type != "cuda": + raise ValueError(f"Flex Attention only supports CUDA tensors, got device {device}.") + if device.index is None: + device = torch.device("cuda", torch.cuda.current_device()) + + handle = _cudnn_score_mod_handles.get(device) + with torch.cuda.device(device): + if handle is None: + handle = cudnn.create_handle() + _cudnn_score_mod_handles[device] = handle + + stream = torch.cuda.current_stream(device).cuda_stream + cudnn.set_stream(handle=handle, stream=stream) + return handle def _build_cudnn_pygraph(dtype: torch.dtype, device: torch.device): """Create a cuDNN frontend Python graph for F16/BF16 SDPA.""" - return cudnn_pygraph.build_pygraph(dtype, device, backend_name=_BACKEND) + cudnn = _import_cudnn_frontend() + + if dtype == torch.float16: + io_data_type = cudnn.data_type.HALF + elif dtype == torch.bfloat16: + io_data_type = cudnn.data_type.BFLOAT16 + else: + raise ValueError(f"Flex Attention only supports FP16/BF16 tensors, got {dtype}.") + + graph = cudnn.pygraph( + io_data_type=io_data_type, + intermediate_data_type=cudnn.data_type.FLOAT, + compute_data_type=cudnn.data_type.FLOAT, + handle=_get_cudnn_current_stream_handle(cudnn, device), + ) + return graph @dataclass @@ -247,19 +263,19 @@ class _CudnnScoreModBwdGraphEntry: workspace_size: int -# cuDNN FROST SDPA engine names. These are barred here, not merely left unasked for: the switch -# that offers them is process-wide, so any FrostAttention call elsewhere in the process, or a user -# setting CUDNN_FRONTEND_ENABLE_FROST_ENGINES, puts them ahead of the backend engines for these -# graphs too. They accept a score_mod graph, pass check_support, build, and then compute without -# the callback. Measured on B200 with cuDNN Frontend 1.29.0: a FROST plan ranks at index 0 at -# head_dim 64, 128, 256 and 512, and an unpinned build selects it and returns plain attention. -_FROST_PLAN_TOKENS = ("sdpa_fwd_prefill_sm100", "sdpa_bwd_sm100") - - def _finalize_cudnn_graph(graph) -> int: """Build a cuDNN frontend Python graph and return its workspace size.""" - workspace_size, _ = cudnn_pygraph.finalize_plans(graph, exclude_plan_tokens=_FROST_PLAN_TOKENS) - return workspace_size + cudnn = _import_cudnn_frontend() + + graph.validate() + graph.build_operation_graph() + try: + graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) + graph.check_support() + except cudnn.cudnnGraphNotSupportedError as exc: + raise RuntimeError(f"cuDNN Flex Attention SDPA graph is not supported: {exc}") from exc + graph.build_plans(cudnn.build_plan_policy.HEURISTICS_CHOICE) + return max(graph.get_workspace_size(), 1) def _execute_cudnn_graph( @@ -269,7 +285,20 @@ def _execute_cudnn_graph( device: torch.device, ): """Execute a built cuDNN frontend Python graph.""" - cudnn_pygraph.execute_graph(graph, variant_pack, workspace_size, device, backend_name=_BACKEND) + cudnn = _import_cudnn_frontend() + + if device.type == "cuda" and device.index is None: + device = torch.device("cuda", torch.cuda.current_device()) + workspace = torch.empty( + workspace_size, + device=device, + dtype=torch.uint8, + ) + graph.execute( + variant_pack, + workspace, + handle=_get_cudnn_current_stream_handle(cudnn, device), + ) def _cudnn_score_mod_fwd_cache_key( @@ -284,7 +313,6 @@ def _cudnn_score_mod_fwd_cache_key( score_mod_tensors: Optional[Dict[str, torch.Tensor]], output_layer: torch.Tensor, stats: Optional[torch.Tensor], - mask_spec: Optional[Tuple[str, Tuple[int, int]]] = None, ) -> Optional[Tuple[Any, ...]]: """Pre-build cache key for score_mod fprop execution plans. @@ -307,9 +335,6 @@ def _cudnn_score_mod_fwd_cache_key( _score_mod_bhsd_tensor_metadata(output_layer, q_format), _score_mod_tensor_metadata(stats) if stats is not None else None, _score_mod_tensor_dict_metadata(score_mod_tensors), - # The mask belongs in the key. Without it two graphs differing only in mask type collide - # and the second silently reuses the first, which is a wrong answer rather than a miss. - mask_spec, ) @@ -328,7 +353,6 @@ def _cudnn_score_mod_bwd_cache_key( score_mod_tensors: Optional[Dict[str, torch.Tensor]], score_mod_bprop_tensors: Optional[Dict[str, torch.Tensor]], deterministic: bool, - mask_spec: Optional[Tuple[str, Tuple[int, int]]] = None, ) -> Optional[Tuple[Any, ...]]: """Pre-build cache key for score_mod bprop execution plans.""" score_mod_key = _score_mod_callback_cache_key(score_mod) @@ -351,7 +375,6 @@ def _cudnn_score_mod_bwd_cache_key( _score_mod_tensor_metadata(stats), _score_mod_tensor_dict_metadata(score_mod_tensors), _score_mod_tensor_dict_metadata(score_mod_bprop_tensors), - mask_spec, ) @@ -367,15 +390,8 @@ def _build_cudnn_score_mod_fwd_graph( score_mod_tensors: Optional[Dict[str, torch.Tensor]], output_layer: torch.Tensor, stats: Optional[torch.Tensor], - mask_spec: Optional[Tuple[str, Tuple[int, int]]] = None, ) -> _CudnnScoreModFwdGraphEntry: - """Build a cached cuDNN frontend graph for score_mod fprop. - - ``mask_spec`` is an optional (attn_mask_type, window) pair. When given, the SDPA node carries - cuDNN's diagonal band instead of the unmasked default, which is how a backend without a - score_mod expresses causal, bottom-right and sliding-window attention. The two are mutually - exclusive: cuDNN rejects a graph carrying both. - """ + """Build a cached cuDNN frontend graph for score_mod fprop.""" cudnn = _import_cudnn_frontend() graph = _build_cudnn_pygraph(query_layer.dtype, query_layer.device) @@ -387,7 +403,6 @@ def _build_cudnn_score_mod_fwd_graph( wrapped_score_mod = _wrap_score_mod(score_mod, score_mod_graph_tensors) output_dim, output_stride = _bhsd_dim_stride(output_layer, q_format) - sdpa_kwargs = _mask_or_score_mod_kwargs(mask_spec, wrapped_score_mod) output, stats_tensor = graph.sdpa( name="te_score_mod_sdpa", q=q, @@ -395,7 +410,8 @@ def _build_cudnn_score_mod_fwd_graph( v=v, generate_stats=is_training, attn_scale=attn_scale, - **sdpa_kwargs, + use_causal_mask=False, + score_mod=wrapped_score_mod, ) output.set_output(True).set_dim(output_dim).set_stride(output_stride) @@ -432,7 +448,6 @@ def _get_cudnn_score_mod_fwd_graph( score_mod_tensors: Optional[Dict[str, torch.Tensor]], output_layer: torch.Tensor, stats: Optional[torch.Tensor], - mask_spec: Optional[Tuple[str, Tuple[int, int]]] = None, ) -> _CudnnScoreModFwdGraphEntry: """Return a cached cuDNN frontend graph for score_mod fprop.""" build_args = ( @@ -448,15 +463,12 @@ def _get_cudnn_score_mod_fwd_graph( output_layer, stats, ) - # Only when set: an unconditional extra argument would change the call shape for every - # existing caller, including the tests that substitute their own builder. - extra = {} if mask_spec is None else {"mask_spec": mask_spec} - key = _cudnn_score_mod_fwd_cache_key(*build_args, **extra) + key = _cudnn_score_mod_fwd_cache_key(*build_args) if key is None: - return _build_cudnn_score_mod_fwd_graph(*build_args, **extra) + return _build_cudnn_score_mod_fwd_graph(*build_args) entry = _cudnn_score_mod_graph_cache.get(key) if entry is None: - entry = _build_cudnn_score_mod_fwd_graph(*build_args, **extra) + entry = _build_cudnn_score_mod_fwd_graph(*build_args) _cudnn_score_mod_graph_cache[key] = entry return entry @@ -476,11 +488,8 @@ def _build_cudnn_score_mod_bwd_graph( score_mod_tensors: Optional[Dict[str, torch.Tensor]], score_mod_bprop_tensors: Optional[Dict[str, torch.Tensor]], deterministic: bool, - mask_spec: Optional[Tuple[str, Tuple[int, int]]] = None, ) -> _CudnnScoreModBwdGraphEntry: - """Build a cached cuDNN frontend graph for score_mod bprop. See the fprop builder for - ``mask_spec``; the backward must carry the same mask as the forward or the gradients are - computed against a different attention.""" + """Build a cached cuDNN frontend graph for score_mod bprop.""" graph = _build_cudnn_pygraph(query_layer.dtype, query_layer.device) q = _bhsd_graph_tensor(graph, query_layer, q_format) k = _bhsd_graph_tensor(graph, key_layer, kv_format) @@ -513,7 +522,8 @@ def _build_cudnn_score_mod_bwd_graph( dO=d_output, stats=stats_tensor, attn_scale=attn_scale, - **_mask_or_score_mod_kwargs(mask_spec, wrapped_score_mod), + use_causal_mask=False, + score_mod=wrapped_score_mod, score_mod_bprop=wrapped_score_mod_bprop, use_deterministic_algorithm=deterministic, ) @@ -554,7 +564,6 @@ def _get_cudnn_score_mod_bwd_graph( score_mod_tensors: Optional[Dict[str, torch.Tensor]], score_mod_bprop_tensors: Optional[Dict[str, torch.Tensor]], deterministic: bool, - mask_spec: Optional[Tuple[str, Tuple[int, int]]] = None, ) -> _CudnnScoreModBwdGraphEntry: """Return a cached cuDNN frontend graph for score_mod bprop.""" build_args = ( @@ -573,13 +582,12 @@ def _get_cudnn_score_mod_bwd_graph( score_mod_bprop_tensors, deterministic, ) - extra = {} if mask_spec is None else {"mask_spec": mask_spec} - key = _cudnn_score_mod_bwd_cache_key(*build_args, **extra) + key = _cudnn_score_mod_bwd_cache_key(*build_args) if key is None: - return _build_cudnn_score_mod_bwd_graph(*build_args, **extra) + return _build_cudnn_score_mod_bwd_graph(*build_args) entry = _cudnn_score_mod_graph_cache.get(key) if entry is None: - entry = _build_cudnn_score_mod_bwd_graph(*build_args, **extra) + entry = _build_cudnn_score_mod_bwd_graph(*build_args) _cudnn_score_mod_graph_cache[key] = entry return entry diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index a4a51662ccd..4cb8a9ea4ee 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -32,13 +32,11 @@ import contextlib import os from importlib.metadata import PackageNotFoundError, version as get_pkg_version -from typing import Optional, Tuple +from typing import Any, Dict, Optional, Sequence, Tuple import torch from packaging.version import InvalidVersion, Version as PkgVersion -from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph - __all__ = [ "is_frost_attention_available", "is_frost_attention_supported", @@ -74,29 +72,191 @@ _HEAD_DIM_MULTIPLE = 8 _cudnn = None +_frost_engines_enabled = False _availability: Optional[Tuple[bool, str]] = None _PLAN_CACHE: dict = {} -_HANDLES = cudnn_pygraph._handles # pylint: disable=protected-access +_HANDLES: Dict[torch.device, Any] = {} + + +def _import_cudnn_frontend(enable_frost_engines: bool = True): + """Import cuDNN Frontend, enabling the FROST engines if this caller needs them. + ``enable_frost_engines`` is not merely additive: the switch also ranks FROST ahead of the + backend engines everywhere, so a caller that does not want FROST must not ask for it. -def _import_cudnn(enable_frost_engines: bool = True): - """Import cuDNN Frontend, registering the FROST engines unless told not to. + The enabling is deliberately outside the import memo. Both backends call this, and whichever + one reaches it first would otherwise decide for the process: with the flag inside the memo, a + flex call would cache the module with FROST off and every later FROST call would get a cuDNN + that offers no FROST engine, which surfaces much later as "no cuDNN engine matching ... was + offered". Enabling late is sound because the switch is read per graph rather than at import: + in cuDNN Frontend 1.29.0 ``engines/manifest.py`` consults the environment inside + ``offered_ids()``, reached from ``engines_for(graph)`` on every ``create_execution_plans``. - The switch is process-wide and ranks FROST ahead of the backend engines for every cuDNN Python - graph afterwards, including other backends’ graphs, so it is set only where FROST is actually - used. _select_frost_plan verifies the engine by plan name regardless, rather than trusting the - flag. + Note the switch is process-wide and never unset, so enabling it for FROST also reorders the + candidates a concurrent score_mod graph sees. Callers that require a particular engine should + verify by plan name rather than rely on the switch, which is what + ``_finalize_plans(require_plan_token=...)`` does. """ - global _cudnn # pylint: disable=global-statement - # Kept bound: _pkg_version falls back to the module's __version__ when distribution metadata - # is unavailable, which is how a source or vendored install avoids being misreported. - _cudnn = cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=enable_frost_engines) + global _cudnn, _frost_engines_enabled # pylint: disable=global-statement + if _cudnn is None: + try: + import cudnn # pylint: disable=import-outside-toplevel + except ImportError as exc: + raise ImportError( + "cuDNN frontend Python package not found. " + "Install it with: pip install nvidia-cudnn-frontend" + ) from exc + + _cudnn = cudnn + + if enable_frost_engines and not _frost_engines_enabled: + os.environ.setdefault("CUDNN_FRONTEND_ENABLE_FROST_ENGINES", "1") + # pylint: disable=import-outside-toplevel,unused-import + import cudnn.sdpa # noqa: F401 + + _frost_engines_enabled = True + return _cudnn -def _handle_for(device: torch.device): - """A cuDNN handle for `device`, bound to PyTorch's current stream on every call.""" - return cudnn_pygraph.handle_for(device, backend_name="FrostAttention") +def _handle_for(device: torch.device, *, backend_name: str = "FrostAttention"): + """A cuDNN handle for ``device``, rebound to PyTorch's current stream on every call. + + Without the rebinding, cuDNN runs on its handle's own stream while the tensors and workspace + are allocated on PyTorch's current stream, and nothing orders the two. That is not + hypothetical: the p2p context-parallel ring issues attention inside + ``with torch.cuda.stream(cp_stream)``, so on alternating ring steps the kernel and its buffers + would otherwise be on different streams. The same cached plan is executed from different + streams across steps, so this has to happen per call rather than once per handle. + """ + if device.type != "cuda": + raise ValueError(f"{backend_name} requires CUDA tensors; got device {device}") + cudnn = _cudnn if _cudnn is not None else _import_cudnn_frontend() + if device.index is None: + device = torch.device("cuda", torch.cuda.current_device()) + with torch.cuda.device(device): + handle = _HANDLES.get(device) + if handle is None: + handle = cudnn.create_handle() + _HANDLES[device] = handle + cudnn.set_stream(handle=handle, stream=torch.cuda.current_stream(device).cuda_stream) + return handle + + +def _build_pygraph( + dtype: torch.dtype, device: torch.device, *, backend_name: str = "FrostAttention" +): + """A cuDNN frontend graph for F16/BF16 SDPA, bound to this device's stream-current handle.""" + cudnn = _cudnn if _cudnn is not None else _import_cudnn_frontend() + return cudnn.pygraph( + io_data_type=_cudnn_dtype(dtype), + intermediate_data_type=cudnn.data_type.FLOAT, + compute_data_type=cudnn.data_type.FLOAT, + handle=_handle_for(device, backend_name=backend_name), + ) + + +def _diagonal_band_kwargs(cudnn, attn_mask_type: str, window: Tuple[int, int]) -> Dict[str, Any]: + """cuDNN sdpa kwargs for a TE (mask type, window): a diagonal alignment plus a band. + + Note the off-by-one. cuDNN's left bound counts the diagonal itself and TE's window_size does + not, so a window of w becomes a left bound of w + 1. Passing it through unconverted silently + drops one token of context per layer, which no shape-level test would catch. + + These kwargs are mutually exclusive with score_mod. cuDNN enforces that in the backward node + only ("Attention score mod enabled and hence other subgraphs are disabled"); its forward node + composes the two without complaint. Callers must still refuse the pair on both sides, because + forward and backward have to carry the same mask or the gradients belong to a different + attention than the output does. + """ + left, right = window + opts: Dict[str, Any] = {} + if attn_mask_type in ("causal", "causal_bottom_right") or right == 0: + opts["diagonal_alignment"] = ( + cudnn.diagonal_alignment.BOTTOM_RIGHT + if attn_mask_type == "causal_bottom_right" + else cudnn.diagonal_alignment.TOP_LEFT + ) + opts["diagonal_band_right_bound"] = 0 + if left != -1: + opts["diagonal_band_left_bound"] = left + 1 + return opts + + +def _finalize_plans( + graph, + *, + heuristics: Optional[Sequence[Any]] = None, + build_policy: Any = None, + require_plan_token: Optional[str] = None, + not_found_hint: Any = "", +) -> Tuple[int, Optional[str]]: + """Create plans, optionally pin one by name, build, and return (workspace size, plan name). + + ``require_plan_token`` makes the choice strict: only a plan whose name contains the token is + acceptable, and anything else raises. That is not a stylistic preference. Without a pin, + ``build_plans`` walks the ranked list from index 0 and finalizes the first plan that builds, + logging each decline at INFO, so a graph that the intended engine declines runs on whatever + cuDNN ranked next with nothing in the return value to say so. At head_dim 512 that matters in + the forward, where an ordinary engine may well build and compute a different function from the + FROST kernel. The backward is self-limiting, since no non-FROST d512 backward exists, so an + unpinned backward would fail loudly on its own. + + The token is matched as a substring rather than by equality on purpose: cuDNN has already + collapsed per-head-dim engine names (``..._d512`` and friends) into a single row once, and the + substring test survived that. + + + Pinning also changes what ``check_support`` means. Selecting a plan sets cuDNN's internal + ``_plan_pinned``, and only then is a decline fatal; unpinned, cuDNN records the decline and + keeps walking. So the pin has to come first both because the check is scoped to the selected + plan and because it is what makes the check binding at all. + """ + cudnn = _cudnn if _cudnn is not None else _import_cudnn_frontend() + + graph.validate() + graph.build_operation_graph() + + if heuristics is None: + heuristics = [cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK] + + if require_plan_token is None: + try: + graph.create_execution_plans(list(heuristics)) + graph.check_support() + except cudnn.cudnnGraphNotSupportedError as exc: + raise RuntimeError(f"cuDNN SDPA graph is not supported: {exc}") from exc + if build_policy is None: + build_policy = cudnn.build_plan_policy.HEURISTICS_CHOICE + graph.build_plans(build_policy) + return max(graph.get_workspace_size(), 1), None + + graph.create_execution_plans(list(heuristics)) + names = [graph.get_plan_name_at_index(i) for i in range(graph.get_execution_plan_count())] + hits = [i for i, n in enumerate(names) if require_plan_token in n] + if not hits: + # Callable hints are resolved only here: a caller may want to look up package versions to + # explain the failure, and that work should not happen on the success path. + hint = not_found_hint() if callable(not_found_hint) else not_found_hint + raise RuntimeError( + f"no cuDNN engine matching {require_plan_token!r} was offered." + f" Candidate plans: {names[:6]}.{(' ' + hint) if hint else ''}" + ) + graph.select_plan(hits[0]) + # The engine is pinned, so a decline here is the engine's own verdict on this graph and cuDNN + # puts its reason in the exception. Surface that rather than letting it escape bare: a plan + # that was offered and then refused is the harder failure to read, and the reason is the only + # thing that says which constraint was missed. + try: + graph.check_support() + graph.build_plans() + except cudnn.cudnnGraphNotSupportedError as exc: + hint = not_found_hint() if callable(not_found_hint) else not_found_hint + raise RuntimeError( + f"cuDNN engine {names[hits[0]]!r} was offered but declined this graph:" + f" {exc}{(' ' + hint) if hint else ''}" + ) from exc + return max(graph.get_workspace_size(), 1), names[hits[0]] def _device_from_key(device_key) -> torch.device: @@ -157,7 +317,7 @@ def _no(reason): # Without the engines: this only needs the module to read a version off it, and enabling # here would reorder plan selection for the whole process even when the checks below go on # to decline FROST, which is all cost and no benefit. The use sites enable it. - _import_cudnn(enable_frost_engines=False) + _import_cudnn_frontend(enable_frost_engines=False) except ImportError as exc: return _no(f"nvidia-cudnn-frontend not importable: {exc}") @@ -234,7 +394,7 @@ def _mask_spec(attn_mask_type: str, window_size=None): def _mask_options(cudnn, spec): """cuDNN sdpa kwargs for a (mask type, window) spec: a diagonal alignment plus a band.""" attn_mask_type, window = spec - return cudnn_pygraph.diagonal_band_kwargs(cudnn, attn_mask_type, window) + return _diagonal_band_kwargs(cudnn, attn_mask_type, window) _SUPPORTED_QKV_FORMATS = ("bshd", "sbhd") @@ -426,7 +586,7 @@ def from_frost_layout(t: torch.Tensor, qkv_format: str) -> torch.Tensor: def _cudnn_dtype(dtype: torch.dtype): - cudnn = _import_cudnn() + cudnn = _import_cudnn_frontend() return { torch.bfloat16: cudnn.data_type.BFLOAT16, torch.float16: cudnn.data_type.HALF, @@ -499,8 +659,8 @@ def hint(): f" (floor {_MIN_CUTLASS_DSL})." ) - cudnn = _import_cudnn() - _, name = cudnn_pygraph.finalize_plans( + cudnn = _import_cudnn_frontend() + _, name = _finalize_plans( graph, heuristics=[cudnn.heur_mode.A], require_plan_token=token, @@ -511,13 +671,13 @@ def hint(): def _build_fwd(key) -> dict: """Build (and JIT-compile) a forward graph. Expensive; always reached through the cache.""" - cudnn = _import_cudnn() + cudnn = _import_cudnn_frontend() # deterministic is unused here: it selects a backward algorithm. Callers pass False for the # forward so the two never split the forward cache. *_device, b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks, _deterministic = key shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] - graph = cudnn_pygraph.build_pygraph( + graph = _build_pygraph( dtype, _device_from_key(_device), backend_name="FrostAttention" ) tq = graph.tensor(name="q", dim=shq, stride=list(qs)) @@ -547,12 +707,12 @@ def _build_fwd(key) -> dict: def _build_bwd(key) -> dict: """Build (and JIT-compile) a backward graph. Expensive; always reached through the cache.""" - cudnn = _import_cudnn() + cudnn = _import_cudnn_frontend() *_device, b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks, deterministic = key io_dt = _cudnn_dtype(dtype) shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] - graph = cudnn_pygraph.build_pygraph( + graph = _build_pygraph( dtype, _device_from_key(_device), backend_name="FrostAttention" ) handles = {} From 76f59c7d0736fb564cd7c4650c2f1079e9914248 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Tue, 6 Oct 2026 04:50:46 +0000 Subject: [PATCH 54/97] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- .../attention/dot_product_attention/frost_attention.py | 8 ++------ 1 file changed, 2 insertions(+), 6 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 4cb8a9ea4ee..3afb9c939e0 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -677,9 +677,7 @@ def _build_fwd(key) -> dict: *_device, b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks, _deterministic = key shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] - graph = _build_pygraph( - dtype, _device_from_key(_device), backend_name="FrostAttention" - ) + graph = _build_pygraph(dtype, _device_from_key(_device), backend_name="FrostAttention") tq = graph.tensor(name="q", dim=shq, stride=list(qs)) tk = graph.tensor(name="k", dim=shkv, stride=list(ks)) tv = graph.tensor(name="v", dim=shkv, stride=list(ks)) @@ -712,9 +710,7 @@ def _build_bwd(key) -> dict: io_dt = _cudnn_dtype(dtype) shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] - graph = _build_pygraph( - dtype, _device_from_key(_device), backend_name="FrostAttention" - ) + graph = _build_pygraph(dtype, _device_from_key(_device), backend_name="FrostAttention") handles = {} # o and dO share q's layout; k, v and their grads share k's. for name, shape, stride in ( From 94577ea84a570634dc2cf9d1e4206836ca51348e Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Mon, 5 Oct 2026 21:59:52 -0700 Subject: [PATCH 55/97] fix(attention): return max_logit by index from the CP p2p fused step cp_p2p_fwd_fused_attn returned its max_logit with a starred tail, which is only unambiguous while no backend has a statically known return length. fused_attn_fwd now has one, the FROST branch, so pylint resolves that tail to empty and reports the five-label unpack at each of the four p2p call sites as unbalanced. Callers already unpack exactly five, and max_logit holds exactly one element whenever return_max_logit is set, so indexing it says what the function actually returns. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../attention/dot_product_attention/context_parallel.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index a98c5538bb9..abc244756f8 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1109,8 +1109,10 @@ def cp_p2p_fwd_fused_attn( softmax_lse_per_step, rng_states, *rest = aux_ctx_tensors attn_bias = rest[0] if len(rest) > 0 else None + # Indexed rather than starred: callers unpack exactly five, and a starred tail lets a + # backend whose return length is statically known collapse it to four. if return_max_logit: - return out_per_step, softmax_lse_per_step, rng_states, attn_bias, *max_logit + return out_per_step, softmax_lse_per_step, rng_states, attn_bias, max_logit[0] return out_per_step, softmax_lse_per_step, rng_states, attn_bias, None From 2cb9db1d36700e6c88dab04b0fa64685256efafb Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Mon, 5 Oct 2026 22:19:54 -0700 Subject: [PATCH 56/97] fix(attention): put the new CP grad slot at the right index The fused_attention_backend argument went in at index 18 of each context-parallel forward, but its gradient slot was appended at the end of each backward return. For AttnFuncWithCPAndKVP2P and AttnFuncWithCPAndKVAllGather that is the same tuple, since every grad from index 18 on is None. AttnFuncWithCPAndQKVOA2A is not: it returns d_softmax_offset, which sat at index 30 against softmax_offset. Shifting the parameters by one without moving the grad left d_softmax_offset on softmax_type and softmax_offset with no gradient at all, so sink attention under cp_comm_type='a2a' would have trained with a silently dropped offset gradient. Verified by comparing, for all three classes, which parameter NAME each non-None gradient lands on against origin/main. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../attention/dot_product_attention/context_parallel.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index abc244756f8..8cb10bcc48c 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -3194,7 +3194,7 @@ def backward(ctx, dout, *_args): attn_dbias, None, None, - None, + None, # fused_attention_backend None, None, None, @@ -4561,7 +4561,7 @@ def backward(ctx, dout, *_args): None, None, None, - None, + None, # fused_attention_backend None, None, None, @@ -5416,6 +5416,7 @@ def backward(ctx, dout, *_args): d_bias, None, None, + None, # fused_attention_backend None, None, None, @@ -5430,7 +5431,6 @@ def backward(ctx, dout, *_args): None, d_softmax_offset, None, - None, ) From d400eea92722166e0ba9395c60af6061ff186796 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 00:01:08 -0700 Subject: [PATCH 57/97] test(attention): count FROST as a comparable fused sub-backend get_available_attention_backends only counted sub-backends 1 and 2, so a FROST selection contributed nothing to the two-backend threshold in test_dot_product_attention. That silently cost coverage at head_dim 512. The test re-queries forward-only when FusedAttention cannot train a config, which is how base_5_0 and base_5_1 used to recover a second backend. FROST can train them, so that fallback no longer fires, and with FROST uncounted the pair fell back under the threshold: four tests went from passing to skipped. Counting it restores them, and as a stronger check than before -- FROST against UnfusedDotProductAttention, forward and backward, instead of a forward-only cuDNN query. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- tests/pytorch/utils.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/tests/pytorch/utils.py b/tests/pytorch/utils.py index 62917a5c8d0..5a52003dc38 100644 --- a/tests/pytorch/utils.py +++ b/tests/pytorch/utils.py @@ -468,7 +468,11 @@ def test(): _attention_backends["backend_selection_requires_update"] = False return available_backends, flash_attention_backend, fused_attention_backend - backends = {1: "F16_arbitrary_seqlen", 2: "FP8"} + # Every fused sub-backend this helper is willing to count as comparable. FROST has to be + # here: it serves head_dim in (256, 512], where it is the only backward-capable fused + # sub-backend, so omitting it makes those configs look like they have one backend and + # the caller skips instead of comparing. + backends = {1: "F16_arbitrary_seqlen", 2: "FP8", 3: "FROST"} if AttentionLogging._is_logging_setup is False: AttentionLogging.setup_logging() From dbc298eff4b403df8bb7e7c43ff7b950ae26d2be Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 00:59:53 -0700 Subject: [PATCH 58/97] test(attention): drop the comment on the sub-backend list Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- tests/pytorch/utils.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/tests/pytorch/utils.py b/tests/pytorch/utils.py index 7534b39a74e..2a7cb1c449d 100644 --- a/tests/pytorch/utils.py +++ b/tests/pytorch/utils.py @@ -511,10 +511,6 @@ def test(): _attention_backends["backend_selection_requires_update"] = False return available_backends, flash_attention_backend, fused_attention_backend - # Every fused sub-backend this helper is willing to count as comparable. FROST has to be - # here: it serves head_dim in (256, 512], where it is the only backward-capable fused - # sub-backend, so omitting it makes those configs look like they have one backend and - # the caller skips instead of comparing. backends = {1: "F16_arbitrary_seqlen", 2: "FP8", 3: "FROST"} if AttentionLogging._is_logging_setup is False: AttentionLogging.setup_logging() From aec8b4ba82b2160b04ff99528d3143908bc32d0c Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Mon, 5 Oct 2026 21:46:12 -0700 Subject: [PATCH 59/97] fix(attention): bar the cuDNN FROST engines from flex attention graphs cuDNN's FROST SDPA engines accept a score_mod graph, pass check_support, build, execute, and return output that is bit-identical to the same graph built with no score_mod at all. No error, no warning, no declined plan. The callback is discarded. flex_attention never asks for those engines, but the switch that offers them is process-wide, so anything else in the process that enables them puts them ahead of the backend engines for these graphs too. Measured on B200 at head_dim 64, 128, 256 and 512, a FROST plan ranks at index 0 for a score_mod graph and an unpinned build selects it every time. Declining to ask is therefore not enough; deselect_engines bars them by name. The call is inert where those engines are not on offer, which is every process that has not enabled them. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_flex_attention.py | 115 ++++++++++++++++++ .../dot_product_attention/flex_attention.py | 9 ++ 2 files changed, 124 insertions(+) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index beed4069917..29a15e2fead 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -705,3 +705,118 @@ def test_dot_product_attention_score_mod(dtype, qkv_format, score_mod_case, scal torch.testing.assert_close(q.grad, q_ref.grad, **tols) torch.testing.assert_close(k.grad, k_ref.grad, **tols) torch.testing.assert_close(v.grad, v_ref.grad, **tols) + + +def test_flex_bars_the_frost_engines(): + """flex must tell cuDNN not to use a FROST engine, not merely decline to ask for them. + + The switch that offers those engines is process-wide, so another caller in the process, or a + user setting CUDNN_FRONTEND_ENABLE_FROST_ENGINES, puts them ahead of the backend engines for + these graphs too. They accept a score_mod graph, pass check_support, build, and then compute + without the callback. + + No GPU: this checks the instruction is passed, not what cuDNN does with it. + """ + barred = [] + + class FakeGraph: + """Records the engines flex bars, and stops at the first call it cannot serve.""" + + def validate(self): + pass + + def build_operation_graph(self): + pass + + def create_execution_plans(self, _heuristics): + pass + + def deselect_engines(self, names): + barred.extend(names) + + def check_support(self): + pass + + def build_plans(self, _policy): + pass + + def get_workspace_size(self): + return 4096 + + try: + flex_attention._import_cudnn_frontend() + except ImportError: + pytest.skip("cuDNN frontend Python package is required for score_mod attention.") + + assert flex_attention._finalize_cudnn_graph(FakeGraph()) == 4096 + assert barred, "flex did not ask cuDNN to exclude any engine" + assert "sdpa_fwd_prefill_sm100" in barred and "sdpa_bwd_sm100" in barred, barred + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required.") +def test_frost_switch_does_not_change_what_flex_computes(): + """Enabling the FROST engines must not change flex's output. + + This is the property the silent drop violates: with the engines on, an unpinned build selects + a FROST plan at every head dim measured on B200, and that plan returns plain attention with + the score_mod discarded. Comparing flex against itself across the switch needs no reference + and no knowledge of which plan ran; if the two differ, a different kernel answered. + """ + try: + flex_attention._import_cudnn_frontend() + except ImportError: + pytest.skip("cuDNN frontend Python package is required for score_mod attention.") + # Without this the test is vacuous: if the engines are absent, decline on arch, or sit below + # their version floors, both runs get a backend plan and agree no matter what flex does. + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + is_frost_attention_available, + ) + + frost_ok, frost_reason = is_frost_attention_available() + if not frost_ok: + pytest.skip("the FROST engines must be reachable to test anything: %s" % frost_reason) + + env = "CUDNN_FRONTEND_ENABLE_FROST_ENGINES" + saved = os.environ.get(env) + torch.manual_seed(0) + b, h, s, d = 2, 4, 512, 64 + dtype = torch.bfloat16 if is_bf16_available() else torch.float16 + q, k, v = (torch.randn(b, s, h, d, device="cuda", dtype=dtype) for _ in range(3)) + + def bias_score_mod(score_mod_graph, score_tensor, _tensors): + """score += (row - col). Self-contained, and large enough that dropping it is obvious.""" + cudnn = flex_attention._import_cudnn_frontend() + row = score_mod_graph.gen_index(input=score_tensor, axis=2) + row.set_data_type(cudnn.data_type.INT32) + col = score_mod_graph.gen_index(input=score_tensor, axis=3) + col.set_data_type(cudnn.data_type.INT32) + bias = score_mod_graph.sub(a=row, b=col, compute_data_type=cudnn.data_type.FLOAT) + bias.set_data_type(cudnn.data_type.FLOAT) + return score_mod_graph.add(a=score_tensor, b=bias, compute_data_type=cudnn.data_type.FLOAT) + + def run(): + flex_attention._cudnn_score_mod_graph_cache.clear() + return flex_attention.FusedAttentionWithScoreModFunc.apply( + False, q, k, v, "bshd", "bshd", d**-0.5, bias_score_mod, None, None, None, False + ) + + try: + os.environ.pop(env, None) + without = run() + os.environ[env] = "1" + with_engines = run() + finally: + flex_attention._cudnn_score_mod_graph_cache.clear() + if saved is None: + os.environ.pop(env, None) + else: + os.environ[env] = saved + + torch.testing.assert_close( + with_engines, + without, + msg=lambda m: ( + "flex computed something different with the FROST engines enabled, which means a" + " FROST plan answered and dropped the score_mod:\n" + m + ), + ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index b9593b42d9b..e551d1c49d9 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -263,6 +263,12 @@ class _CudnnScoreModBwdGraphEntry: workspace_size: int +# cuDNN's FROST SDPA engines accept a score_mod graph, build, run, and then compute without the +# callback. The switch that offers them is process-wide, so they can be ranked ahead of the +# backend engines for these graphs even though this file never asks for them. +_FROST_PLAN_TOKENS = ("sdpa_fwd_prefill_sm100", "sdpa_bwd_sm100") + + def _finalize_cudnn_graph(graph) -> int: """Build a cuDNN frontend Python graph and return its workspace size.""" cudnn = _import_cudnn_frontend() @@ -271,6 +277,9 @@ def _finalize_cudnn_graph(graph) -> int: graph.build_operation_graph() try: graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) + # Bar them before the walk, so build_plans falls through to the first entry that is both + # unbarred and buildable. Inert when those engines are not on offer. + graph.deselect_engines(list(_FROST_PLAN_TOKENS)) graph.check_support() except cudnn.cudnnGraphNotSupportedError as exc: raise RuntimeError(f"cuDNN Flex Attention SDPA graph is not supported: {exc}") from exc From 5d506d88386a9f01a931d236168ae3d024985df1 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Tue, 6 Oct 2026 08:44:51 +0000 Subject: [PATCH 60/97] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/pytorch/attention/test_flex_attention.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index 29a15e2fead..174caa59490 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -817,6 +817,7 @@ def run(): without, msg=lambda m: ( "flex computed something different with the FROST engines enabled, which means a" - " FROST plan answered and dropped the score_mod:\n" + m + " FROST plan answered and dropped the score_mod:\n" + + m ), ) From df5ae6e8d756b05e4d870ef20308753ce8be4e4c Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 02:14:05 -0700 Subject: [PATCH 61/97] refactor(attention): trim the FROST comments to the repo's register The new code ran noticeably heavier on prose than the files it joins. The surrounding attention modules do carry long comment blocks, but almost all of them are support tables rather than rationale; prose is reserved for a non-obvious hazard. Applied that bar. Kept the measured hazards: the cutlass-dsl floor below which every FROST engine declines silently, the 1.29.0 backward floor, the head_dim padding, the process-wide engine switch, why the plan is checked by name, and why outputs use empty_strided. Cut the design rationale the commit messages and the PR description already carry. frost_attention.py drops from 9.2% comment lines to 7.0%, in line with the other attention modules, and its module docstring from 23 lines to 9; the three properties it listed are each already explained where they bite. No code changed: verified by comparing ASTs with docstrings stripped. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/flex_attention.py | 7 +- .../dot_product_attention/frost_attention.py | 95 ++++++------------- .../attention/dot_product_attention/utils.py | 9 +- .../pytorch/cpp_extensions/fused_attn.py | 6 +- 4 files changed, 37 insertions(+), 80 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index e551d1c49d9..b3e704d9f48 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -263,9 +263,8 @@ class _CudnnScoreModBwdGraphEntry: workspace_size: int -# cuDNN's FROST SDPA engines accept a score_mod graph, build, run, and then compute without the -# callback. The switch that offers them is process-wide, so they can be ranked ahead of the -# backend engines for these graphs even though this file never asks for them. +# These engines accept a score_mod graph and then compute without it, and the switch that offers +# them is process-wide, so declining to ask for them is not enough. _FROST_PLAN_TOKENS = ("sdpa_fwd_prefill_sm100", "sdpa_bwd_sm100") @@ -277,8 +276,6 @@ def _finalize_cudnn_graph(graph) -> int: graph.build_operation_graph() try: graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) - # Bar them before the walk, so build_plans falls through to the first entry that is both - # unbarred and buildable. Inert when those engines are not on offer. graph.deselect_engines(list(_FROST_PLAN_TOKENS)) graph.check_support() except cudnn.cudnnGraphNotSupportedError as exc: diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 3afb9c939e0..023aa7b5d7d 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -11,20 +11,6 @@ registered at Python import time behind CUDNN_FRONTEND_ENABLE_FROST_ENGINES and require the nvidia-cutlass-dsl Python package, while TE's C++ builds against cuDNN Frontend headers only. Reaching them requires a Python graph, which is what this module is. - -Three properties of these kernels were verified on Blackwell, and each constrains the code: - -1. cuDNN's causal masking is TOP_LEFT aligned unless bottom-right is requested. The two coincide - when SQ == SKV, so the distinction is invisible in square tests and decisive for all_gather, - which trims KV. Masking is built as a diagonal band so causal, bottom-right and sliding window - come from one mechanism. - -2. Plan building must be cached. It dominates an execute even after cuDNN has cached the JIT, so a - per-call build would leave training build-bound. Hence `_PLAN_CACHE`. - -3. The forward LSE is natural-log logsumexp in fp32, shaped [b, h, s, 1]. Squeezed to [b, h, s] it - is what the CP ring correction in context_parallel.py consumes, which is what makes ring - attention over these kernels valid at all. """ from __future__ import annotations @@ -49,26 +35,22 @@ ] -# FROST engines are opt-in inside cuDNN Frontend, and they additionally require a newer -# nvidia-cutlass-dsl than cudnn-frontend itself declares. cudnn-frontend requires >= 4.6.2 while -# FROST enforces >= 4.7.0 at plan-build time; with 4.6.2 installed every FROST engine silently -# declines and ordinary cuDNN backend plans are returned with no error at all. We therefore check -# the selected plan by NAME rather than trusting that the engine was used. +# cudnn-frontend declares cutlass-dsl >= 4.6.2 but FROST enforces >= 4.7.0 at plan-build time. +# Below that floor every FROST engine declines silently and backend plans come back instead, so +# the selected plan is checked by NAME rather than trusting that the engine was used. _FROST_FWD_PLAN_TOKEN = "sdpa_fwd_prefill_sm100" _FROST_BWD_PLAN_TOKEN = "sdpa_bwd_sm100" _MIN_CUTLASS_DSL = PkgVersion("4.7.0") -# 1.29.0 is the first release carrying the head_dim=512 BACKWARD (bprop_d512_f16_sm100). 1.28.0 -# ships the forward only, and the repo's own pin allows it, so without this check training would -# build a forward plan and then raise on the first backward. +# 1.29.0 is the first release carrying the head_dim=512 backward. 1.28.0 ships the forward only, +# so without this check training would build a forward plan and raise on the first backward. _MIN_CUDNN_FRONTEND = PkgVersion("1.29.0") _SUPPORTED_ARCHS = ((10, 0), (10, 3)) _MAX_HEAD_DIM = 512 _MIN_HEAD_DIM = 257 # below this the existing cuDNN/flash backends already serve the shape -# The engine pads head_dim to a multiple of 8, so 260 is not servable even though it is in range. -# Without this it passes the gate and then fails at plan selection with a message about missing -# engines, instead of declining cleanly here. +# The engine pads head_dim to a multiple of 8, so 260 is in range but not servable. Declined +# here rather than failing later at plan selection. _HEAD_DIM_MULTIPLE = 8 _cudnn = None @@ -243,10 +225,8 @@ def _finalize_plans( f" Candidate plans: {names[:6]}.{(' ' + hint) if hint else ''}" ) graph.select_plan(hits[0]) - # The engine is pinned, so a decline here is the engine's own verdict on this graph and cuDNN - # puts its reason in the exception. Surface that rather than letting it escape bare: a plan - # that was offered and then refused is the harder failure to read, and the reason is the only - # thing that says which constraint was missed. + # The engine is pinned, so a decline here is its own verdict and cuDNN puts the reason in the + # exception. Surface it: a plan offered and then refused is the harder failure to read. try: graph.check_support() graph.build_plans() @@ -314,16 +294,14 @@ def _no(reason): major, minor = torch.cuda.get_device_capability() return _no(f"cuDNN FROST head_dim>256 kernels are SM100/SM103 only; found sm{major}{minor}") try: - # Without the engines: this only needs the module to read a version off it, and enabling - # here would reorder plan selection for the whole process even when the checks below go on - # to decline FROST, which is all cost and no benefit. The use sites enable it. + # Without the engines: this only reads a version, and enabling reorders plan selection + # process-wide even when the checks below decline. The use sites enable it. _import_cudnn_frontend(enable_frost_engines=False) except ImportError as exc: return _no(f"nvidia-cudnn-frontend not importable: {exc}") - # Decline on positive evidence that FROST cannot work: a version below a floor, or a package - # that is absent outright. A version that is present but unparseable is NOT evidence, so it - # defers to _select_frost_plan, which checks the plan by name and reports both versions. + # Decline only on positive evidence: a version below a floor, or a package absent outright. + # An unparseable version defers to _select_frost_plan, which checks the plan by name. frontend, frontend_raw = _pkg_version("nvidia-cudnn-frontend", _cudnn) if frontend is not None and frontend < _MIN_CUDNN_FRONTEND: return _no( @@ -346,17 +324,9 @@ def _no(reason): return _availability -# TE mask types this backend serves. cuDNN expresses causal, bottom-right and sliding-window -# masking as ONE mechanism -- a diagonal alignment plus a two-sided band -- rather than three -# separate flags, so that is what _mask_options builds. The legacy spellings desugar into exactly -# that: pygraph/sdpa.cpp maps 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. Building the band directly is equivalent for those two and -# additionally expresses a left bound, which is what a sliding window is. -# -# Both alignments are needed. The p2p ring produces square diagonal tiles, where top-left and -# bottom-right coincide, while all_gather trims KV and relies on bottom-right alignment, where -# the two differ completely. +# cuDNN expresses causal, bottom-right and sliding-window masking as one mechanism, a diagonal +# alignment plus a two-sided band, which is what _mask_options builds. Both alignments are needed: +# the p2p ring produces square tiles where they coincide, while all_gather trims KV so they differ. _SUPPORTED_MASKS = ("no_mask", "causal", "causal_bottom_right") # Sliding window as TE spells it: (left, right), -1 meaning unbounded on that side. @@ -542,10 +512,9 @@ def is_frost_attention_supported(params) -> Tuple[int, str]: and "causal" not in attn_mask_type and params.max_seqlen_q != params.max_seqlen_kv ): - # A right-bounded window on a non-causal mask takes its anchor only from - # bottom_right_diagonal, which defaults to top-left, while the all-gather ring trims KV - # and measures its window against the bottom-right diagonal. Those differ exactly when - # the q and kv lengths do, so decline rather than guess which one was meant. + # Such a window takes its anchor only from bottom_right_diagonal, which defaults to + # top-left, while the all-gather ring measures its window bottom-right. Decline rather + # than guess which was meant. return ( no_backend, ( @@ -757,9 +726,8 @@ def _cached(kind: str, key): cache_key = (kind,) + key entry = _PLAN_CACHE.get(cache_key) if entry is None: - # Build under the device the key names, not merely with that device's handle: the plans - # are CuTe-DSL JIT-compiled, and a compile path is far more likely to read the ambient - # CUDA context than the handle. Free to do, and removes the question entirely. + # Build under the device the key names, not merely with its handle: the plans are + # JIT-compiled, and a compile path may read the ambient CUDA context rather than the handle. device = _device_from_key(key[:2]) with torch.cuda.device(device) if device.type == "cuda" else contextlib.nullcontext(): entry = _build_fwd(key) if kind == "fwd" else _build_bwd(key) @@ -769,9 +737,8 @@ def _cached(kind: str, key): def _key(q, k, mask, scale, deterministic=False): return ( - # The graph is built under whichever device was current, so it must not be reused on - # another one. Matches the C++ fused-attn cache, which keys on device_id for the same - # reason. Type is included too, so a CPU tensor cannot alias cuda:0. + # Built under whichever device was current, so it must not be reused on another. Matches + # the C++ fused-attn cache, which keys on device_id. Type too, so CPU cannot alias cuda:0. q.device.type, q.device.index, q.shape[0], @@ -828,10 +795,8 @@ def frost_attn_fwd( tq, tk, tv, tout, tlse = entry["handles"] b, hq, sq, _ = q.shape - # Allocate per call: the cache holds only the compiled plan, never output buffers, so that - # concurrent or nested uses cannot alias each other. empty_strided rather than empty_like: - # the latter does not preserve an arbitrary permuted stride, and the graph was built for - # q's exact strides. + # Allocated per call so concurrent uses cannot alias; the cache holds only the plan. + # empty_strided, not empty_like: the latter does not preserve an arbitrary permuted stride. out = torch.empty_strided(q.shape, q.stride(), device=q.device, dtype=q.dtype) lse = torch.empty(b, hq, sq, 1, device=q.device, dtype=torch.float32) workspace = torch.empty(entry["workspace"], device=q.device, dtype=torch.uint8) @@ -886,9 +851,8 @@ def frost_attn_bwd( softmax_lse = softmax_lse.unsqueeze(-1) softmax_lse = softmax_lse.contiguous() - # The graph expects o and dO in q's layout. A caller may hand us either with different - # strides (dO in particular comes from autograd), so restride rather than silently reading - # the wrong elements. + # The graph expects o and dO in q's layout, and dO comes from autograd with strides we do + # not control, so restride rather than silently reading the wrong elements. def _as(t, ref): if tuple(t.stride()) == tuple(ref.stride()): return t @@ -996,9 +960,8 @@ def fused_attn_fwd( attn_mask_type=mask_type, window_size=window, ) - # A real tensor rather than None: it is saved for backward and handed to the activation - # offload hooks alongside softmax_lse, neither of which accepts None. FROST has no dropout, - # so nothing reads it. + # A real tensor, not None: it is saved for backward and handed to the activation offload + # hooks, neither of which accepts None. FROST has no dropout, so nothing reads it. rng_state = torch.empty(2, dtype=torch.int64, device=q.device) return from_frost_layout(out, qkv_format), [softmax_lse, rng_state] diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 83e6b8e211f..b30429fc542 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -452,9 +452,8 @@ def _get_fused_attn_backend(**fused_attn_kwargs): params = FusedAttentionParams(**fused_attn_kwargs) fused_attention_backend, reject_message = tex.get_fused_attn_backend(params) if fused_attention_backend == FusedAttnBackend.No_Backend: - # FROST is a python sub-backend, so the C++ selector cannot see it. It serves symmetric - # head_dim in (256, 512] on SM100/SM103, which nothing above it covers. Availability is - # checked once at the end of get_attention_backend, the way flash-attn's version is. + # A python sub-backend, invisible to the C++ selector. Availability is checked at the + # end of get_attention_backend, the way flash-attn's version is. from .frost_attention import ( # pylint: disable=import-outside-toplevel is_frost_attention_supported, ) @@ -1879,8 +1878,8 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt use_flash_attention_4 = False use_flash_attention = use_flash_attention_2 or use_flash_attention_3 or use_flash_attention_4 if use_fused_attention and fused_attention_backend == FusedAttnBackend.FROST.value: - # Deferred to here because probing it imports cuDNN Frontend with the FROST engines - # enabled, which changes the engine pool for every cuDNN consumer in the process. + # Deferred to here: probing imports cuDNN Frontend with the engines enabled, which + # changes the engine pool for every cuDNN consumer in the process. from .frost_attention import ( # pylint: disable=import-outside-toplevel is_frost_attention_available, ) diff --git a/transformer_engine/pytorch/cpp_extensions/fused_attn.py b/transformer_engine/pytorch/cpp_extensions/fused_attn.py index 8e731f813e3..9dd93f3ce1a 100644 --- a/transformer_engine/pytorch/cpp_extensions/fused_attn.py +++ b/transformer_engine/pytorch/cpp_extensions/fused_attn.py @@ -119,9 +119,8 @@ class FusedAttnBackend(IntEnum): No_Backend = int(NVTE_Fused_Attn_Backend.NVTE_No_Backend) F16_arbitrary_seqlen = int(NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen) FP8 = int(NVTE_Fused_Attn_Backend.NVTE_FP8) - # Python-only: cuDNN FROST runs through the cuDNN Frontend python API rather than the C++ - # fused-attention path, so it has no NVTE_Fused_Attn_Backend counterpart. fused_attn_fwd/bwd - # route it to frost_attention.py before any C++ call, so this value never reaches pybind. + # Python-only: FROST runs through the cuDNN Frontend python API, so it has no + # NVTE_Fused_Attn_Backend counterpart and never reaches pybind. FROST = 3 @classmethod @@ -355,7 +354,6 @@ def fused_attn_fwd( # Accept the pybind enum for backward compatibility. fused_attention_backend = FusedAttnBackend.cast(fused_attention_backend) if fused_attention_backend == FusedAttnBackend["FROST"]: - # FROST runs through the cuDNN Frontend python API rather than the C++ fused path. # 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 From 597e5931c29a62e64a5ae208d8c096bdd34f16a3 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 18:38:48 -0700 Subject: [PATCH 62/97] refactor(attention): shorten the max_logit unpack comment State the invariant the indexed form relies on and drop the linter detail. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../attention/dot_product_attention/context_parallel.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index 8cb10bcc48c..f36bb46401c 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1109,8 +1109,7 @@ def cp_p2p_fwd_fused_attn( softmax_lse_per_step, rng_states, *rest = aux_ctx_tensors attn_bias = rest[0] if len(rest) > 0 else None - # Indexed rather than starred: callers unpack exactly five, and a starred tail lets a - # backend whose return length is statically known collapse it to four. + # Indexed rather than starred: every caller unpacks exactly five values. if return_max_logit: return out_per_step, softmax_lse_per_step, rng_states, attn_bias, max_logit[0] return out_per_step, softmax_lse_per_step, rng_states, attn_bias, None From 6b1a6d8d1b439e0dcef52115b7b8e0a1065a9e12 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 18:51:48 -0700 Subject: [PATCH 63/97] feat(attention): serve asymmetric head_dim on the FROST sub-backend Give v its own graph node and its own cache-key entry instead of declaring it with k's shape and stride, the way flex_attention keys each tensor separately. The symmetry requirement came from that shared declaration, not from the kernels. O, dO and the O-shaped grads follow q's layout with v's head_dim, so they get their strides from a helper the graph node and the allocation both call. Symmetric head dims keep q's exact strides, which leaves the existing path byte-identical. The selector now range-checks each head_dim on its own, so an asymmetric pair inside (256, 512] is accepted. The runtime guard keeps batch, heads and seqlen, which index the same KV positions as k by definition. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../attention/test_attention_with_cp.py | 2 +- .../pytorch/attention/test_frost_attention.py | 86 ++++++++---- .../dot_product_attention/frost_attention.py | 131 +++++++++++------- 3 files changed, 145 insertions(+), 74 deletions(-) diff --git a/tests/pytorch/attention/test_attention_with_cp.py b/tests/pytorch/attention/test_attention_with_cp.py index 321fbdb171f..ebe91da63fc 100644 --- a/tests/pytorch/attention/test_attention_with_cp.py +++ b/tests/pytorch/attention/test_attention_with_cp.py @@ -459,7 +459,7 @@ def test_cp_with_flash_attention_softcap(cp_pool, cp_comm_type): ) -# cuDNN FROST: symmetric head_dim in (256, 512] on SM100/SM103, the range no other backend +# cuDNN FROST: head_dim in (256, 512] on SM100/SM103, the range no other backend # serves together with context parallelism. Shapes are Gemma-4 global layers, which is what # motivated the backend. seqlen must stay divisible by cp_size * 2 for causal load balancing. model_configs_frost_attn = { diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 395f3f6beb1..3f578ebe023 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -59,14 +59,21 @@ def _frost_availability(): # head_dim 512 is the whole point of the backend; 320 checks the interior of the (256, 512] range # rather than only its endpoint. _SHAPES = [ - # b, hq, hkv, sq, skv, d - (2, 8, 4, 1024, 1024, 512), # Gemma-4 global layer, GQA - (2, 8, 8, 512, 512, 512), # MHA - (1, 4, 4, 256, 512, 512), # sq != skv, which is where mask alignment matters - (2, 4, 4, 512, 512, 320), # interior head_dim + # b, hq, hkv, sq, skv, d, d_v + (2, 8, 4, 1024, 1024, 512, 512), # Gemma-4 global layer, GQA + (2, 8, 8, 512, 512, 512, 512), # MHA + (1, 4, 4, 256, 512, 512, 512), # sq != skv, which is where mask alignment matters + (2, 4, 4, 512, 512, 320, 320), # interior head_dim + # d_v != d_qk. O and the O-shaped grads take q's layout with v's head_dim, so this is the + # case that catches a plan or an allocation still built from q's trailing dimension. + (2, 8, 4, 512, 512, 512, 320), ] +def _shape_id(s): + return "b%d_hq%d_hkv%d_sq%d_skv%d_d%d_dv%d" % s + + def _reference(q, k, v, scale, mask, window=None): """Attention in float64, computed independently of TE and of cuDNN. @@ -120,7 +127,7 @@ def _floor(q32, k32, v32, scale, mask, dtype, window=None): @requires_frost -@pytest.mark.parametrize("shape", _SHAPES, ids=lambda s: "b%d_hq%d_hkv%d_sq%d_skv%d_d%d" % s) +@pytest.mark.parametrize("shape", _SHAPES, ids=_shape_id) @pytest.mark.parametrize("mask", ["no_mask", "causal", "causal_bottom_right"]) @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) def test_frost_forward_matches_reference(shape, mask, dtype): @@ -129,15 +136,15 @@ def test_frost_forward_matches_reference(shape, mask, dtype): frost_attn_fwd, ) - b, hq, hkv, sq, skv, d = shape + b, hq, hkv, sq, skv, d, d_v = shape torch.manual_seed(0) # Generate in fp32 so there is a true high-precision original to measure against, then cast # for the kernel. [b, h, s, d] views over bshd-contiguous memory is what the backend consumes. # A bshd VIEW, which is what the backend receives: to_frost_layout permutes a bshd-contiguous # tensor and hands the result over without a copy. Materialising with .contiguous() here would # produce bhsd strides instead and leave the stride-keyed plan cache untested. - mk = lambda s_, h_: torch.randn(b, s_, h_, d, device="cuda").permute(0, 2, 1, 3) - q32, k32, v32 = mk(sq, hq), mk(skv, hkv), mk(skv, hkv) + mk = lambda s_, h_, d_: torch.randn(b, s_, h_, d_, device="cuda").permute(0, 2, 1, 3) + q32, k32, v32 = mk(sq, hq, d), mk(skv, hkv, d), mk(skv, hkv, d_v) q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) scale = 1.0 / math.sqrt(d) @@ -212,7 +219,7 @@ def test_frost_sliding_window_matches_reference(mask, window, sq, skv): @requires_frost -@pytest.mark.parametrize("shape", _SHAPES[:2], ids=lambda s: "b%d_hq%d_hkv%d_sq%d_skv%d_d%d" % s) +@pytest.mark.parametrize("shape", _SHAPES[:2] + _SHAPES[-1:], ids=_shape_id) @pytest.mark.parametrize("mask", ["no_mask", "causal"]) @pytest.mark.parametrize("window", [None, (128, 0)], ids=["nowin", "win128"]) @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) @@ -228,13 +235,13 @@ def test_frost_backward_matches_reference(shape, mask, window, dtype): frost_attn_fwd, ) - b, hq, hkv, sq, skv, d = shape + b, hq, hkv, sq, skv, d, d_v = shape torch.manual_seed(0) # A bshd VIEW, which is what the backend receives: to_frost_layout permutes a bshd-contiguous # tensor and hands the result over without a copy. Materialising with .contiguous() here would # produce bhsd strides instead and leave the stride-keyed plan cache untested. - mk = lambda s_, h_: torch.randn(b, s_, h_, d, device="cuda").permute(0, 2, 1, 3) - q32, k32, v32 = mk(sq, hq), mk(skv, hkv), mk(skv, hkv) + mk = lambda s_, h_, d_: torch.randn(b, s_, h_, d_, device="cuda").permute(0, 2, 1, 3) + q32, k32, v32 = mk(sq, hq, d), mk(skv, hkv, d), mk(skv, hkv, d_v) q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) scale = 1.0 / math.sqrt(d) @@ -312,10 +319,13 @@ def test_frost_declines_unsupported_configs(): assert ( is_frost_attention_supported(_frost_params())[0] == FusedAttnBackend.FROST ), "the supported case must be accepted" + assert ( + is_frost_attention_supported(_frost_params(head_dim_v=320))[0] == FusedAttnBackend.FROST + ), "an asymmetric head_dim pair inside the range must be accepted" for override, why in ( (dict(head_dim_qk=256, head_dim_v=256), "head_dim at the exclusive lower bound"), - (dict(head_dim_v=256), "asymmetric head_dim"), + (dict(head_dim_v=256), "head_dim_v below the range"), (dict(qkv_dtype=TE_DType[torch.float32]), "fp32"), (dict(dropout=0.1), "dropout"), (dict(bias_type=AttnBiasType["post_scale_bias"]), "attention bias"), @@ -346,6 +356,8 @@ def test_frost_declines_unsupported_configs(): # The engine pads head_dim to a multiple of 8, so an in-range but unpadded dim has to be # declined here rather than failing later at plan selection. (dict(head_dim_qk=260, head_dim_v=260), "head_dim not a multiple of 8"), + # Each dim is checked on its own, so v has to be covered as well as q. + (dict(head_dim_v=260), "head_dim_v not a multiple of 8"), ): backend, reason = is_frost_attention_supported(_frost_params(**override)) assert backend == FusedAttnBackend.No_Backend, "%s must be declined" % why @@ -444,7 +456,7 @@ def test_frost_sliding_window_selection_by_cp_comm_type(cp_comm_type, window, ex @requires_frost def test_frost_rejects_mismatched_kv(): - """k and v must agree: the graphs declare v with k's shape and stride.""" + """v must index the same KV positions as k. head_dim is free; the rest is not.""" from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( frost_attn_fwd, ) @@ -454,20 +466,46 @@ def test_frost_rejects_mismatched_kv(): mk = lambda hh: torch.randn(b, s, hh, d, device="cuda", dtype=dtype).permute(0, 2, 1, 3) q, k = mk(h).contiguous(), mk(h).contiguous() - with pytest.raises(ValueError, match="same shape"): + with pytest.raises(ValueError, match="batch, heads and seqlen"): frost_attn_fwd(q, k, mk(h * 2).contiguous()) - with pytest.raises(ValueError, match="same layout"): - # Same shape, different stride order: a cache hit would otherwise run a graph built for - # k's layout over v's memory and read the wrong elements silently. Build it as sbhd and - # permute, so the strides genuinely differ -- a [b, h, s, d] contiguous tensor would come - # out with exactly k's strides and prove nothing. - v_odd = torch.randn(s, b, h, d, device="cuda", dtype=dtype).permute(1, 2, 0, 3) - assert v_odd.shape == k.shape and v_odd.stride() != k.stride() - frost_attn_fwd(q, k, v_odd) with pytest.raises(ValueError, match="match q"): frost_attn_fwd(q, k, k.to(torch.float32)) +@requires_frost +def test_frost_serves_v_with_its_own_head_dim_and_layout(): + """v is keyed and declared separately, the way flex_attention keys each tensor. + + Checked against the float64 reference rather than against another FROST call: a v whose + head_dim AND stride order both differ from k's is exactly the case that a plan built from + k alone would compute wrongly without raising. + """ + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + frost_attn_fwd, + ) + + b, h, s, d, d_v = 2, 4, 512, 512, 320 + dtype = torch.bfloat16 + torch.manual_seed(0) + q32 = torch.randn(b, s, h, d, device="cuda").permute(0, 2, 1, 3) + k32 = torch.randn(b, s, h, d, device="cuda").permute(0, 2, 1, 3) + # sbhd rather than bshd, so v's stride order differs from k's as well as its head_dim. + v32 = torch.randn(s, b, h, d_v, device="cuda").permute(1, 2, 0, 3) + q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) + assert v.stride()[:3] != k.stride()[:3], "v must not share k's stride order here" + scale = 1.0 / math.sqrt(d) + + out, lse = frost_attn_fwd(q, k, v, attn_scale=scale, attn_mask_type="causal") + + floor_o, floor_l, ref_o, ref_lse = _floor(q32, k32, v32, scale, "causal", dtype) + assert out.shape == (b, h, s, d_v), "out takes v's head_dim; got %s" % (tuple(out.shape),) + assert out.stride(3) == 1, "out must stay head-contiguous; got stride %s" % (out.stride(),) + err_o = (out.double() - ref_o).abs().max().item() + err_l = (lse.double() - ref_lse).abs().max().item() + assert err_o <= 2 * floor_o + 1e-3, "out err %.3e exceeds 2x the floor %.3e" % (err_o, floor_o) + assert err_l <= 2 * floor_l + 1e-3, "lse err %.3e exceeds 2x the floor %.3e" % (err_l, floor_l) + + def test_frost_engines_are_enabled_even_if_cudnn_was_imported_without_them(): """Enabling the FROST engines must not depend on who imported cuDNN first. diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 023aa7b5d7d..9c3960d3a66 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -440,16 +440,19 @@ def is_frost_attention_supported(params) -> Tuple[int, str]: if int(os.environ.get("NVTE_FROST_ATTN", "1")) == 0: return no_backend, "FROST is disabled by NVTE_FROST_ATTN=0" - head_dim_qk, head_dim_v = params.head_dim_qk, params.head_dim_v - if head_dim_qk != head_dim_v: - return no_backend, f"FROST requires symmetric head_dim; got {head_dim_qk}/{head_dim_v}" - if not _MIN_HEAD_DIM <= head_dim_qk <= _MAX_HEAD_DIM: - return no_backend, f"FROST covers head_dim in (256, 512]; got {head_dim_qk}" - if head_dim_qk % _HEAD_DIM_MULTIPLE != 0: - return ( - no_backend, - f"FROST needs head_dim to be a multiple of {_HEAD_DIM_MULTIPLE}; got {head_dim_qk}", - ) + # Each head_dim is checked on its own: q/k and v get separate graph nodes, so an + # asymmetric pair is served as long as both dims land in the range. + for name, head_dim in ( + ("head_dim_qk", params.head_dim_qk), + ("head_dim_v", params.head_dim_v), + ): + if not _MIN_HEAD_DIM <= head_dim <= _MAX_HEAD_DIM: + return no_backend, f"FROST covers head_dim in (256, 512]; got {name}={head_dim}" + if head_dim % _HEAD_DIM_MULTIPLE != 0: + return ( + no_backend, + f"FROST needs {name} to be a multiple of {_HEAD_DIM_MULTIPLE}; got {head_dim}", + ) qkv_dtype = TORCH_DType.get(params.qkv_dtype) if qkv_dtype not in (torch.bfloat16, torch.float16): @@ -590,22 +593,43 @@ def _check_dtype(name: str, t: torch.Tensor, expected: torch.dtype) -> None: def _check_kv_match(k: torch.Tensor, v: torch.Tensor) -> None: - """Require v to match k in both shape and layout. + """Require v to agree with k on batch, heads and sequence length. - Both graphs declare v with k's shape and stride, and _key records only q's and k's, so a v - that differs would hit a cached plan built for k's layout and read the wrong elements with no - error at all. Callers in TE always split k and v from one QKV tensor, so this costs nothing - and is purely a guard against a silent wrong answer. + head_dim is free: v has its own graph node and its own cache-key entry, so an asymmetric + pair builds its own plan. The other three index the same KV positions as k by definition, + and a mismatch would bind a differently shaped buffer with no error at all. """ - if k.shape != v.shape: - raise ValueError(f"k and v must have the same shape; got {k.shape} and {v.shape}") - if k.stride() != v.stride(): + if tuple(k.shape[:3]) != tuple(v.shape[:3]): raise ValueError( - f"k and v must have the same layout; got strides {tuple(k.stride())} and" - f" {tuple(v.stride())}" + f"k and v must agree on batch, heads and seqlen; got {k.shape} and {v.shape}" ) +def _head_dim_strides(shape: Sequence[int], ref_strides: Sequence[int]) -> list: + """Dense strides for ``shape`` in the memory order ``ref_strides`` describes. + + O, dO and the O-shaped grads follow q's layout but carry v's head_dim, so when the two head + dims differ they cannot reuse q's strides. The graph node and the allocation both go through + here so they cannot drift apart. + """ + order = sorted(range(len(shape)), key=lambda i: ref_strides[i], reverse=True) + strides = [0] * len(shape) + acc = 1 + for i in reversed(order): + strides[i] = acc + acc *= shape[i] + return strides + + +def _o_shape_stride(b, hq, sq, d, d_v, qs): + """Shape and strides for an O-shaped tensor: q's layout, v's head_dim. + + Symmetric head dims keep q's exact strides, which preserves a caller's non-dense view. + """ + shape = [b, hq, sq, d_v] + return shape, (list(qs) if d_v == d else _head_dim_strides(shape, qs)) + + def _select_frost_plan(graph, token: str, what: str): """Select a plan whose name proves a FROST engine was chosen. @@ -643,13 +667,14 @@ def _build_fwd(key) -> dict: cudnn = _import_cudnn_frontend() # deterministic is unused here: it selects a backward algorithm. Callers pass False for the # forward so the two never split the forward cache. - *_device, b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks, _deterministic = key - shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] + *_device, b, hq, hkv, sq, skv, d, d_v, dtype, mask, scale, qs, ks, vs, _deterministic = key + shq, shk, shv = [b, hq, sq, d], [b, hkv, skv, d], [b, hkv, skv, d_v] + sho, o_stride = _o_shape_stride(b, hq, sq, d, d_v, qs) graph = _build_pygraph(dtype, _device_from_key(_device), backend_name="FrostAttention") tq = graph.tensor(name="q", dim=shq, stride=list(qs)) - tk = graph.tensor(name="k", dim=shkv, stride=list(ks)) - tv = graph.tensor(name="v", dim=shkv, stride=list(ks)) + tk = graph.tensor(name="k", dim=shk, stride=list(ks)) + tv = graph.tensor(name="v", dim=shv, stride=list(vs)) tout, tlse = graph.sdpa( name="frost_fwd", q=tq, @@ -659,7 +684,7 @@ def _build_fwd(key) -> dict: attn_scale=scale, **_mask_options(cudnn, mask), ) - tout.set_output(True).set_dim(shq).set_stride(list(qs)) # out mirrors q + tout.set_output(True).set_dim(sho).set_stride(list(o_stride)) # out: q's layout, v's head_dim tlse.set_output(True).set_dim([b, hq, sq, 1]).set_stride([hq * sq, sq, 1, 1]).set_data_type( cudnn.data_type.FLOAT ) @@ -675,19 +700,20 @@ def _build_fwd(key) -> dict: def _build_bwd(key) -> dict: """Build (and JIT-compile) a backward graph. Expensive; always reached through the cache.""" cudnn = _import_cudnn_frontend() - *_device, b, hq, hkv, sq, skv, d, dtype, mask, scale, qs, ks, deterministic = key + *_device, b, hq, hkv, sq, skv, d, d_v, dtype, mask, scale, qs, ks, vs, deterministic = key io_dt = _cudnn_dtype(dtype) - shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] + shq, shk, shv = [b, hq, sq, d], [b, hkv, skv, d], [b, hkv, skv, d_v] + sho, o_stride = _o_shape_stride(b, hq, sq, d, d_v, qs) graph = _build_pygraph(dtype, _device_from_key(_device), backend_name="FrostAttention") handles = {} - # o and dO share q's layout; k, v and their grads share k's. + # Each grad is declared with the layout of the tensor it differentiates. for name, shape, stride in ( ("q", shq, qs), - ("k", shkv, ks), - ("v", shkv, ks), - ("o", shq, qs), - ("do", shq, qs), + ("k", shk, ks), + ("v", shv, vs), + ("o", sho, o_stride), + ("do", sho, o_stride), ): handles[name] = graph.tensor(name=name, dim=shape, stride=list(stride)) handles["stats"] = graph.tensor( @@ -708,7 +734,7 @@ def _build_bwd(key) -> dict: use_deterministic_algorithm=deterministic, **_mask_options(cudnn, mask), ) - for tensor, stride in ((tdq, qs), (tdk, ks), (tdv, ks)): + for tensor, stride in ((tdq, qs), (tdk, ks), (tdv, vs)): tensor.set_output(True).set_data_type(io_dt).set_stride(list(stride)) plan = _select_frost_plan(graph, _FROST_BWD_PLAN_TOKEN, "backward") handles["dq"], handles["dk"], handles["dv"] = tdq, tdk, tdv @@ -735,7 +761,7 @@ def _cached(kind: str, key): return entry -def _key(q, k, mask, scale, deterministic=False): +def _key(q, k, v, mask, scale, deterministic=False): return ( # Built under whichever device was current, so it must not be reused on another. Matches # the C++ fused-attn cache, which keys on device_id. Type too, so CPU cannot alias cuda:0. @@ -747,6 +773,9 @@ def _key(q, k, mask, scale, deterministic=False): q.shape[2], k.shape[2], q.shape[3], + # v carries its own head_dim and strides, the way flex_attention keys each tensor + # separately. Without them an asymmetric v would reuse a plan built for k's shape. + v.shape[3], q.dtype, mask, float(scale), @@ -754,6 +783,7 @@ def _key(q, k, mask, scale, deterministic=False): # lets bshd and sbhd both run without a transpose. tuple(q.stride()), tuple(k.stride()), + tuple(v.stride()), # The deterministic backward is a different algorithm, not a flag on the same one, so a # plan built either way must not be handed to a call that asked for the other. bool(deterministic), @@ -771,8 +801,9 @@ def frost_attn_fwd( """Forward attention via cuDNN FROST. q, k, v are [b, h, s, d] views; bshd and sbhd are both served, since the graph is built from - each tensor's actual strides. GQA is supported directly (h_kv may differ from h_q) and SQ - need not equal SKV, which is what lets a CP ring step use this. Returns (out, softmax_lse) + each tensor's actual strides. GQA is supported directly (h_kv may differ from h_q), SQ need + not equal SKV, which is what lets a CP ring step use this, and v may carry its own head_dim, + in which case out follows q's layout with v's head_dim. Returns (out, softmax_lse) with softmax_lse as [b, h, s] fp32 natural-log logsumexp, the layout and convention the CP ring correction expects. """ @@ -791,13 +822,14 @@ def frost_attn_fwd( mask = _mask_spec(attn_mask_type, window_size) scale = attn_scale if attn_scale is not None else q.shape[-1] ** -0.5 - entry = _cached("fwd", _key(q, k, mask, scale)) + entry = _cached("fwd", _key(q, k, v, mask, scale)) tq, tk, tv, tout, tlse = entry["handles"] - b, hq, sq, _ = q.shape + b, hq, sq, d = q.shape + out_shape, out_stride = _o_shape_stride(b, hq, sq, d, v.shape[3], q.stride()) # Allocated per call so concurrent uses cannot alias; the cache holds only the plan. # empty_strided, not empty_like: the latter does not preserve an arbitrary permuted stride. - out = torch.empty_strided(q.shape, q.stride(), device=q.device, dtype=q.dtype) + out = torch.empty_strided(out_shape, out_stride, device=q.device, dtype=q.dtype) lse = torch.empty(b, hq, sq, 1, device=q.device, dtype=torch.float32) workspace = torch.empty(entry["workspace"], device=q.device, dtype=torch.uint8) entry["graph"].execute( @@ -831,9 +863,10 @@ def frost_attn_bwd( raise ValueError( f"num_heads must be divisible by num_gqa_groups; got {q.shape[1]} and {k.shape[1]}" ) + o_shape, o_stride = _o_shape_stride(*q.shape, v.shape[3], q.stride()) for name, tensor in (("out", out), ("dout", dout)): - if tensor.shape != q.shape: - raise ValueError(f"{name} must have q's shape; got {tensor.shape} and {q.shape}") + if list(tensor.shape) != o_shape: + raise ValueError(f"{name} must be shaped {o_shape}; got {list(tensor.shape)}") if softmax_lse.dtype != torch.float32: raise ValueError(f"softmax_lse must be fp32; got {softmax_lse.dtype}") if tuple(softmax_lse.shape[:3]) != tuple(q.shape[:3]): @@ -844,24 +877,24 @@ def frost_attn_bwd( mask = _mask_spec(attn_mask_type, window_size) scale = attn_scale if attn_scale is not None else q.shape[-1] ** -0.5 - entry = _cached("bwd", _key(q, k, mask, scale, deterministic)) + entry = _cached("bwd", _key(q, k, v, mask, scale, deterministic)) h = entry["handles"] if softmax_lse.dim() == 3: softmax_lse = softmax_lse.unsqueeze(-1) softmax_lse = softmax_lse.contiguous() - # The graph expects o and dO in q's layout, and dO comes from autograd with strides we do - # not control, so restride rather than silently reading the wrong elements. - def _as(t, ref): - if tuple(t.stride()) == tuple(ref.stride()): + # The graph expects o and dO in the layout the forward wrote, and dO comes from autograd + # with strides we do not control, so restride rather than silently reading the wrong elements. + def _as(t, stride): + if list(t.stride()) == list(stride): return t - buf = torch.empty_strided(t.shape, ref.stride(), device=t.device, dtype=t.dtype) + buf = torch.empty_strided(t.shape, stride, device=t.device, dtype=t.dtype) buf.copy_(t) return buf - out = _as(out, q) - dout = _as(dout, q) + out = _as(out, o_stride) + dout = _as(dout, o_stride) dq = torch.empty_strided(q.shape, q.stride(), device=q.device, dtype=q.dtype) dk = torch.empty_strided(k.shape, k.stride(), device=k.device, dtype=k.dtype) From 2fca15da5b392263d1cd16fc637f3541dc1d0517 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 19:01:33 -0700 Subject: [PATCH 64/97] revert(attention): move the flex FROST engine bar out of this change The bar belongs with flex attention, not with the FROST backend, so it is reviewed on its own rather than inside this one. It stays worth doing. frost_attention sets CUDNN_FRONTEND_ENABLE_FROST_ENGINES process-wide and never unsets it, and cuDNN reads that switch per graph, so a process that uses this backend also reorders the engines a later score_mod graph sees. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_flex_attention.py | 116 ------------------ .../dot_product_attention/flex_attention.py | 6 - 2 files changed, 122 deletions(-) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index 174caa59490..beed4069917 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -705,119 +705,3 @@ def test_dot_product_attention_score_mod(dtype, qkv_format, score_mod_case, scal torch.testing.assert_close(q.grad, q_ref.grad, **tols) torch.testing.assert_close(k.grad, k_ref.grad, **tols) torch.testing.assert_close(v.grad, v_ref.grad, **tols) - - -def test_flex_bars_the_frost_engines(): - """flex must tell cuDNN not to use a FROST engine, not merely decline to ask for them. - - The switch that offers those engines is process-wide, so another caller in the process, or a - user setting CUDNN_FRONTEND_ENABLE_FROST_ENGINES, puts them ahead of the backend engines for - these graphs too. They accept a score_mod graph, pass check_support, build, and then compute - without the callback. - - No GPU: this checks the instruction is passed, not what cuDNN does with it. - """ - barred = [] - - class FakeGraph: - """Records the engines flex bars, and stops at the first call it cannot serve.""" - - def validate(self): - pass - - def build_operation_graph(self): - pass - - def create_execution_plans(self, _heuristics): - pass - - def deselect_engines(self, names): - barred.extend(names) - - def check_support(self): - pass - - def build_plans(self, _policy): - pass - - def get_workspace_size(self): - return 4096 - - try: - flex_attention._import_cudnn_frontend() - except ImportError: - pytest.skip("cuDNN frontend Python package is required for score_mod attention.") - - assert flex_attention._finalize_cudnn_graph(FakeGraph()) == 4096 - assert barred, "flex did not ask cuDNN to exclude any engine" - assert "sdpa_fwd_prefill_sm100" in barred and "sdpa_bwd_sm100" in barred, barred - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required.") -def test_frost_switch_does_not_change_what_flex_computes(): - """Enabling the FROST engines must not change flex's output. - - This is the property the silent drop violates: with the engines on, an unpinned build selects - a FROST plan at every head dim measured on B200, and that plan returns plain attention with - the score_mod discarded. Comparing flex against itself across the switch needs no reference - and no knowledge of which plan ran; if the two differ, a different kernel answered. - """ - try: - flex_attention._import_cudnn_frontend() - except ImportError: - pytest.skip("cuDNN frontend Python package is required for score_mod attention.") - # Without this the test is vacuous: if the engines are absent, decline on arch, or sit below - # their version floors, both runs get a backend plan and agree no matter what flex does. - from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( - is_frost_attention_available, - ) - - frost_ok, frost_reason = is_frost_attention_available() - if not frost_ok: - pytest.skip("the FROST engines must be reachable to test anything: %s" % frost_reason) - - env = "CUDNN_FRONTEND_ENABLE_FROST_ENGINES" - saved = os.environ.get(env) - torch.manual_seed(0) - b, h, s, d = 2, 4, 512, 64 - dtype = torch.bfloat16 if is_bf16_available() else torch.float16 - q, k, v = (torch.randn(b, s, h, d, device="cuda", dtype=dtype) for _ in range(3)) - - def bias_score_mod(score_mod_graph, score_tensor, _tensors): - """score += (row - col). Self-contained, and large enough that dropping it is obvious.""" - cudnn = flex_attention._import_cudnn_frontend() - row = score_mod_graph.gen_index(input=score_tensor, axis=2) - row.set_data_type(cudnn.data_type.INT32) - col = score_mod_graph.gen_index(input=score_tensor, axis=3) - col.set_data_type(cudnn.data_type.INT32) - bias = score_mod_graph.sub(a=row, b=col, compute_data_type=cudnn.data_type.FLOAT) - bias.set_data_type(cudnn.data_type.FLOAT) - return score_mod_graph.add(a=score_tensor, b=bias, compute_data_type=cudnn.data_type.FLOAT) - - def run(): - flex_attention._cudnn_score_mod_graph_cache.clear() - return flex_attention.FusedAttentionWithScoreModFunc.apply( - False, q, k, v, "bshd", "bshd", d**-0.5, bias_score_mod, None, None, None, False - ) - - try: - os.environ.pop(env, None) - without = run() - os.environ[env] = "1" - with_engines = run() - finally: - flex_attention._cudnn_score_mod_graph_cache.clear() - if saved is None: - os.environ.pop(env, None) - else: - os.environ[env] = saved - - torch.testing.assert_close( - with_engines, - without, - msg=lambda m: ( - "flex computed something different with the FROST engines enabled, which means a" - " FROST plan answered and dropped the score_mod:\n" - + m - ), - ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index b3e704d9f48..b9593b42d9b 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -263,11 +263,6 @@ class _CudnnScoreModBwdGraphEntry: workspace_size: int -# These engines accept a score_mod graph and then compute without it, and the switch that offers -# them is process-wide, so declining to ask for them is not enough. -_FROST_PLAN_TOKENS = ("sdpa_fwd_prefill_sm100", "sdpa_bwd_sm100") - - def _finalize_cudnn_graph(graph) -> int: """Build a cuDNN frontend Python graph and return its workspace size.""" cudnn = _import_cudnn_frontend() @@ -276,7 +271,6 @@ def _finalize_cudnn_graph(graph) -> int: graph.build_operation_graph() try: graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) - graph.deselect_engines(list(_FROST_PLAN_TOKENS)) graph.check_support() except cudnn.cudnnGraphNotSupportedError as exc: raise RuntimeError(f"cuDNN Flex Attention SDPA graph is not supported: {exc}") from exc From d04ece93f04341a87baa434609bf611b68f1be27 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 19:40:04 -0700 Subject: [PATCH 65/97] docs(attention): drop the stale symmetric head_dim claim for FROST cc3f6016 removed the symmetry requirement; two lines still asserted it. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- docs/envvars.rst | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index 1df2ed5bcca..7353b717175 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -179,7 +179,7 @@ In PyTorch, the broad preference order is ``FlashAttention > FusedAttention > UnfusedDotProductAttention`` on supported pre-Hopper GPUs such as Ampere/Ada, and ``FusedAttention > FlashAttention > UnfusedDotProductAttention`` on Hopper and newer GPUs, including Blackwell. On Blackwell SM100/SM103, FusedAttention has an extra sub-backend, FROST, -which is selected only for symmetric ``head_dim`` in (256, 512] and only when the cuDNN +which is selected only for ``head_dim`` in (256, 512] and only when the cuDNN sub-backends decline; it does not change the order above. In JAX, Transformer Engine uses cuDNN fused attention when ``NVTE_FUSED_ATTN=1`` and an eligible cuDNN kernel is available; otherwise it falls back to the JAX-native implementation. See :doc:`examples/attention/attention` for a @@ -219,7 +219,7 @@ longer backend-selection overview. :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. + :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 ``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. .. envvar:: NVTE_UNFUSED_ATTN From 0c4601a25bb201077745d72a7d1816427bef88b7 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 21:43:17 -0700 Subject: [PATCH 66/97] fix(attention): resolve the FROST diagonal anchor the way the dispatcher does bottom_right_diagonal defaults to None, meaning "read it off the mask name". cpp_extensions.fused_attn resolves that before dispatching, so the shim only saw a bool and applied bool(). A direct call left None, and a caller asking for causal_bottom_right got a top-left band with nothing raised. Only the None case changes; an explicit True or False behaves as before. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/frost_attention.py | 20 +++++++++++++++++-- 1 file changed, 18 insertions(+), 2 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 9c3960d3a66..c51543c6c8b 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -404,6 +404,18 @@ def _te_mask_spec(attn_mask_type: str, window_size, bottom_right_diagonal: bool) return _mask_spec(attn_mask_type, (left, right)) +def _bottom_right_diagonal(attn_mask_type: str, bottom_right_diagonal) -> bool: + """Resolve the anchor flag the same way cpp_extensions.fused_attn does. + + ``None`` means "read it off the mask name". The dispatcher resolves it before calling in, + so this only matters for a direct call, where ``bool(None)`` would quietly give a top-left + band to a caller that asked for bottom-right. + """ + if bottom_right_diagonal is None: + return attn_mask_type in {"causal_bottom_right", "padding_causal_bottom_right"} + return bool(bottom_right_diagonal) + + def _name_for(table, value, default=None): """Reverse a cpp_extensions str-to-enum table.""" for name, enum_value in table.items(): @@ -983,7 +995,9 @@ def fused_attn_fwd( raise NotImplementedError( f"FROST attention needs o_format to match qkv_format; got {o_format}/{qkv_format}" ) - mask_type, window = _te_mask_spec(attn_mask_type, window_size, bool(bottom_right_diagonal)) + mask_type, window = _te_mask_spec( + attn_mask_type, window_size, _bottom_right_diagonal(attn_mask_type, bottom_right_diagonal) + ) out, softmax_lse = frost_attn_fwd( to_frost_layout(q.contiguous(), qkv_format), @@ -1047,7 +1061,9 @@ def fused_attn_bwd( cuda_graph_capture=cuda_graph, ) qkv_format = _qkv_format_from_layout(qkv_layout) - mask_type, window = _te_mask_spec(attn_mask_type, window_size, bool(bottom_right_diagonal)) + mask_type, window = _te_mask_spec( + attn_mask_type, window_size, _bottom_right_diagonal(attn_mask_type, bottom_right_diagonal) + ) softmax_lse = aux_ctx_tensors[0] dq, dk, dv = frost_attn_bwd( From c1c3a863c42db4c197e37a9a1ed4de3f2e7fbfbb Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 22:05:16 -0700 Subject: [PATCH 67/97] refactor(attention): give the cuDNN pygraph state one owner flex_attention.py and frost_attention.py both drive cuDNN Frontend's Python graph API, and both kept their own copy of the import memo, the per-device handle cache, the graph constructor and plan finalization. The point of sharing them is not line count -- the shared surface is small, and this change is net +57 lines -- but that the state is process-global. There is one cudnn module, one CUDNN_FRONTEND_ENABLE_FROST_ENGINES switch, one engine ranking and one handle per device, and two owners of those is how this code has produced bugs before. Three things worth calling out: - The enable switch defaults to off in the shared module and frost asks for it explicitly, so a caller that does not want the FROST engines cannot turn them on for the process by accident. The enabling stays outside the import memo for the reason its docstring gives. - Handles are keyed on (backend_name, device), not device. One handle shared by two backends would widen an existing within-backend race, since a cuDNN handle is not thread-safe and set_stream mutates it. - io_data_type takes the cudnn module as a parameter, so a dtype lookup can no longer import the frontend or flip the switch as a side effect. flex keeps its own function names as delegations, so its call sites are untouched and its error strings are byte-identical; verified by comparing the outputs and messages against the originals. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_frost_attention.py | 31 ++- .../dot_product_attention/cudnn_pygraph.py | 239 ++++++++++++++++++ .../dot_product_attention/flex_attention.py | 80 ++---- .../dot_product_attention/frost_attention.py | 189 ++------------ 4 files changed, 298 insertions(+), 241 deletions(-) create mode 100644 transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 3f578ebe023..2d2763b4062 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -524,12 +524,17 @@ def test_frost_engines_are_enabled_even_if_cudnn_was_imported_without_them(): import sys import types - from transformer_engine.pytorch.attention.dot_product_attention import frost_attention + from transformer_engine.pytorch.attention.dot_product_attention import ( + cudnn_pygraph, + frost_attention, + ) + # The import memo and the switch live in the shared module; frost's wrapper only supplies the + # default. Drive it through the wrapper, which is the real entry, and assert on the owner. env = "CUDNN_FRONTEND_ENABLE_FROST_ENGINES" saved = ( - frost_attention._cudnn, - frost_attention._frost_engines_enabled, + cudnn_pygraph._cudnn, + cudnn_pygraph._frost_engines_enabled, os.environ.get(env), sys.modules.get("cudnn"), sys.modules.get("cudnn.sdpa"), @@ -539,24 +544,24 @@ def test_frost_engines_are_enabled_even_if_cudnn_was_imported_without_them(): stub.sdpa = types.ModuleType("cudnn.sdpa") sys.modules["cudnn"] = stub sys.modules["cudnn.sdpa"] = stub.sdpa - frost_attention._cudnn = None - frost_attention._frost_engines_enabled = False + cudnn_pygraph._cudnn = None + cudnn_pygraph._frost_engines_enabled = False os.environ.pop(env, None) # The availability probe first, which must not enable anything. frost_attention._import_cudnn_frontend(enable_frost_engines=False) assert env not in os.environ, "the non-FROST caller must not set the switch" - assert not frost_attention._frost_engines_enabled + assert not cudnn_pygraph.frost_engines_enabled() # A use site second, on an already-imported cuDNN. This is the case that used to be # skipped. frost_attention._import_cudnn_frontend(enable_frost_engines=True) assert os.environ.get(env) == "1", "FROST was requested after the import and not enabled" - assert frost_attention._frost_engines_enabled + assert cudnn_pygraph.frost_engines_enabled() finally: ( - frost_attention._cudnn, - frost_attention._frost_engines_enabled, + cudnn_pygraph._cudnn, + cudnn_pygraph._frost_engines_enabled, prior_env, prior_cudnn, prior_sdpa, @@ -583,7 +588,10 @@ def test_pinned_plan_decline_reports_the_engine_reason(): No GPU: a stub graph stands in, raising the real cuDNN exception type. """ - from transformer_engine.pytorch.attention.dot_product_attention import frost_attention + from transformer_engine.pytorch.attention.dot_product_attention import ( + cudnn_pygraph, + frost_attention, + ) try: cudnn = frost_attention._import_cudnn_frontend() @@ -615,8 +623,9 @@ def check_support(self): raise cudnn.cudnnGraphNotSupportedError("head_dim 512 needs SM100; this is SM90") with pytest.raises(RuntimeError) as excinfo: - frost_attention._finalize_plans( + cudnn_pygraph.finalize_plans( _DeclinedGraph(), + backend_name="FrostAttention", heuristics=[cudnn.heur_mode.A], require_plan_token="sdpa_fwd_prefill_sm100", not_found_hint="nvidia-cutlass-dsl=4.8.0.", diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py new file mode 100644 index 00000000000..4ce7d178991 --- /dev/null +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -0,0 +1,239 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Mechanics of driving cuDNN Frontend's Python graph API from PyTorch. + +Importing the frontend, holding one stream-current handle per device, describing TE tensors in +cuDNN's logical BHSD form, and creating, selecting and building plans. No attention semantics, and +no knowledge of any backend's cache-key layout. + +``flex_attention.py`` and ``frost_attention.py`` both drive cuDNN through this API. They share +this module for ownership rather than for line count: the state below is process-global -- one +``cudnn`` module, one ``CUDNN_FRONTEND_ENABLE_FROST_ENGINES`` switch, one engine ranking, one +handle per device -- and giving it two owners is how this code has produced bugs before. + +``backend_name`` is threaded through purely so a failure still says which backend was driving. +""" + +from __future__ import annotations + +import importlib +import os +from typing import Any, Dict, Optional, Sequence, Tuple + +import torch + +_cudnn = None +_frost_engines_enabled = False +_HANDLES: Dict[Tuple[str, torch.device], Any] = {} + + +def import_cudnn_frontend(enable_frost_engines: bool = False): + """Import cuDNN Frontend, enabling the FROST engines if this caller needs them. + + ``enable_frost_engines`` is not merely additive: the switch also ranks FROST ahead of the + backend engines everywhere, so a caller that does not want FROST must not ask for it. Hence + the default is off, and FROST asks explicitly. + + The enabling is deliberately outside the import memo. Both backends call this, and whichever + one reaches it first would otherwise decide for the process: with the flag inside the memo, a + flex call would cache the module with FROST off and every later FROST call would get a cuDNN + that offers no FROST engine, which surfaces much later as "no cuDNN engine matching ... was + offered". Enabling late is sound because the switch is read per graph rather than at import: + in cuDNN Frontend 1.29.0 ``engines/manifest.py`` consults the environment inside + ``offered_ids()``, reached from ``engines_for(graph)`` on every ``create_execution_plans``. + + Note the switch is process-wide and never unset, so enabling it for FROST also reorders the + candidates a concurrent score_mod graph sees. Callers that require a particular engine should + verify by plan name rather than rely on the switch, which is what + ``finalize_plans(require_plan_token=...)`` does. + """ + global _cudnn, _frost_engines_enabled # pylint: disable=global-statement + if _cudnn is None: + try: + _cudnn = importlib.import_module("cudnn") + except ImportError as exc: + raise ImportError( + "cuDNN frontend Python package not found. " + "Install it with: pip install nvidia-cudnn-frontend" + ) from exc + + if enable_frost_engines and not _frost_engines_enabled: + os.environ.setdefault("CUDNN_FRONTEND_ENABLE_FROST_ENGINES", "1") + importlib.import_module("cudnn.sdpa") + _frost_engines_enabled = True + + return _cudnn + + +def cudnn_module(): + """The imported frontend, or None if nothing has imported it yet. + + For callers that want to inspect the module without triggering an import, such as a version + probe that must not enable anything as a side effect. + """ + return _cudnn + + +def frost_engines_enabled() -> bool: + """Whether this process has switched the FROST engines on.""" + return _frost_engines_enabled + + +def handle_for(device: torch.device, *, backend_name: str = "cuDNN attention"): + """A cuDNN handle for ``device``, rebound to PyTorch's current stream on every call. + + Without the rebinding, cuDNN runs on its handle's own stream while the tensors and workspace + are allocated on PyTorch's current stream, and nothing orders the two. That is not + hypothetical: the p2p context-parallel ring issues attention inside + ``with torch.cuda.stream(cp_stream)``, so on alternating ring steps the kernel and its buffers + would otherwise be on different streams. The same cached plan is executed from different + streams across steps, so this has to happen per call rather than once per handle. + + Keyed on ``(backend_name, device)`` rather than on the device alone. A cuDNN handle is not + thread-safe and ``set_stream`` mutates it, so one handle shared by two backends widens an + existing within-backend race into a cross-backend one for no benefit. + """ + if device.type != "cuda": + raise ValueError(f"{backend_name} only supports CUDA tensors, got device {device}.") + cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() + if device.index is None: + device = torch.device("cuda", torch.cuda.current_device()) + key = (backend_name, device) + with torch.cuda.device(device): + handle = _HANDLES.get(key) + if handle is None: + handle = cudnn.create_handle() + _HANDLES[key] = handle + cudnn.set_stream(handle=handle, stream=torch.cuda.current_stream(device).cuda_stream) + return handle + + +def io_data_type(cudnn, dtype: torch.dtype, *, backend_name: str = "cuDNN attention"): + """Map a torch dtype to the cuDNN enum these SDPA graphs are declared with. + + Takes ``cudnn`` rather than importing it, so a dtype lookup cannot import the frontend or + flip the FROST switch as a side effect. + """ + if dtype == torch.float16: + return cudnn.data_type.HALF + if dtype == torch.bfloat16: + return cudnn.data_type.BFLOAT16 + raise ValueError(f"{backend_name} only supports FP16/BF16 tensors, got {dtype}.") + + +def build_pygraph( + dtype: torch.dtype, device: torch.device, *, backend_name: str = "cuDNN attention" +): + """A cuDNN frontend graph for F16/BF16 SDPA, bound to this device's stream-current handle.""" + cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() + return cudnn.pygraph( + io_data_type=io_data_type(cudnn, dtype, backend_name=backend_name), + intermediate_data_type=cudnn.data_type.FLOAT, + compute_data_type=cudnn.data_type.FLOAT, + handle=handle_for(device, backend_name=backend_name), + ) + + +def bhsd_dim_stride( + tensor: torch.Tensor, tensor_format: str, *, backend_name: str = "cuDNN attention" +) -> Tuple[Tuple[int, ...], Tuple[int, ...]]: + """Describe an SBHD/BSHD tensor as cuDNN frontend's logical BHSD format. + + The tensor is never permuted. cuDNN takes dims and strides, so reordering the descriptors + says the same thing as permuting the tensor and costs nothing. + """ + if tensor_format == "sbhd": + return ( + (tensor.shape[1], tensor.shape[2], tensor.shape[0], tensor.shape[3]), + (tensor.stride(1), tensor.stride(2), tensor.stride(0), tensor.stride(3)), + ) + if tensor_format == "bshd": + return ( + (tensor.shape[0], tensor.shape[2], tensor.shape[1], tensor.shape[3]), + (tensor.stride(0), tensor.stride(2), tensor.stride(1), tensor.stride(3)), + ) + raise ValueError(f"{backend_name} only supports SBHD/BSHD tensor formats, got {tensor_format}.") + + +def bhsd_graph_tensor( + graph, tensor: torch.Tensor, tensor_format: str, *, backend_name: str = "cuDNN attention" +): + """Create a cuDNN graph tensor with BHSD dims and TE-layout strides.""" + dim, stride = bhsd_dim_stride(tensor, tensor_format, backend_name=backend_name) + return graph.tensor(dim=dim, stride=stride, data_type=tensor.dtype) + + +def finalize_plans( + graph, + *, + backend_name: str = "cuDNN attention", + heuristics: Optional[Sequence[Any]] = None, + build_policy: Any = None, + require_plan_token: Optional[str] = None, + not_found_hint: Any = "", +) -> Tuple[int, Optional[str]]: + """Create plans, optionally pin one by name, build, and return (workspace size, plan name). + + ``require_plan_token`` makes the choice strict: only a plan whose name contains the token is + acceptable, and anything else raises. That is not a stylistic preference. Without a pin, + ``build_plans`` walks the ranked list from index 0 and finalizes the first plan that builds, + logging each decline at INFO, so a graph that the intended engine declines runs on whatever + cuDNN ranked next with nothing in the return value to say so. At head_dim 512 that matters in + the forward, where an ordinary engine may well build and compute a different function from the + FROST kernel. The backward is self-limiting, since no non-FROST d512 backward exists, so an + unpinned backward would fail loudly on its own. + + The token is matched as a substring rather than by equality on purpose: cuDNN has already + collapsed per-head-dim engine names (``..._d512`` and friends) into a single row once, and the + substring test survived that. + + Pinning also changes what ``check_support`` means. Selecting a plan sets cuDNN's internal + ``_plan_pinned``, and only then is a decline fatal; unpinned, cuDNN records the decline and + keeps walking. So the pin has to come first both because the check is scoped to the selected + plan and because it is what makes the check binding at all. + """ + cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() + + graph.validate() + graph.build_operation_graph() + + if heuristics is None: + heuristics = [cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK] + + if require_plan_token is None: + try: + graph.create_execution_plans(list(heuristics)) + graph.check_support() + except cudnn.cudnnGraphNotSupportedError as exc: + raise RuntimeError(f"cuDNN {backend_name} SDPA graph is not supported: {exc}") from exc + if build_policy is None: + build_policy = cudnn.build_plan_policy.HEURISTICS_CHOICE + graph.build_plans(build_policy) + return max(graph.get_workspace_size(), 1), None + + graph.create_execution_plans(list(heuristics)) + names = [graph.get_plan_name_at_index(i) for i in range(graph.get_execution_plan_count())] + hits = [i for i, n in enumerate(names) if require_plan_token in n] + if not hits: + # Callable hints are resolved only here: a caller may want to look up package versions to + # explain the failure, and that work should not happen on the success path. + hint = not_found_hint() if callable(not_found_hint) else not_found_hint + raise RuntimeError( + f"no cuDNN engine matching {require_plan_token!r} was offered." + f" Candidate plans: {names[:6]}.{(' ' + hint) if hint else ''}" + ) + graph.select_plan(hits[0]) + # The engine is pinned, so a decline here is its own verdict and cuDNN puts the reason in the + # exception. Surface it: a plan offered and then refused is the harder failure to read. + try: + graph.check_support() + graph.build_plans() + except cudnn.cudnnGraphNotSupportedError as exc: + hint = not_found_hint() if callable(not_found_hint) else not_found_hint + raise RuntimeError( + f"cuDNN engine {names[hits[0]]!r} was offered but declined this graph:" + f" {exc}{(' ' + hint) if hint else ''}" + ) from exc + return max(graph.get_workspace_size(), 1), names[hits[0]] diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index b9593b42d9b..38caf5b740b 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -5,49 +5,37 @@ """cuDNN-backed Flex Attention helpers.""" from dataclasses import dataclass -import importlib import inspect from typing import Any, Callable, Dict, Optional, Tuple import torch -_cudnn_score_mod_handles: Dict[torch.device, Any] = {} +from . import cudnn_pygraph + +_BACKEND_NAME = "Flex Attention" _cudnn_score_mod_graph_cache: Dict[Tuple[Any, ...], Any] = {} _SCORE_MOD_UNCACHEABLE = object() def _import_cudnn_frontend(): - """Import the cuDNN frontend Python package.""" - try: - return importlib.import_module("cudnn") - except ImportError as exc: - raise ImportError( - "cuDNN frontend Python package not found. " - "Install it with: pip install nvidia-cudnn-frontend" - ) from exc + """Import the cuDNN frontend Python package. + + Never enables the FROST engines: that switch is process-wide, so a backend that does not want + them must not ask. See ``cudnn_pygraph.import_cudnn_frontend``. + """ + return cudnn_pygraph.import_cudnn_frontend() def _bhsd_dim_stride( tensor: torch.Tensor, tensor_format: str ) -> Tuple[Tuple[int, ...], Tuple[int, ...]]: """Describe an SBHD/BSHD tensor as cuDNN frontend's logical BHSD format.""" - if tensor_format == "sbhd": - return ( - (tensor.shape[1], tensor.shape[2], tensor.shape[0], tensor.shape[3]), - (tensor.stride(1), tensor.stride(2), tensor.stride(0), tensor.stride(3)), - ) - if tensor_format == "bshd": - return ( - (tensor.shape[0], tensor.shape[2], tensor.shape[1], tensor.shape[3]), - (tensor.stride(0), tensor.stride(2), tensor.stride(1), tensor.stride(3)), - ) - raise ValueError(f"Flex Attention only supports SBHD/BSHD tensor formats, got {tensor_format}.") + return cudnn_pygraph.bhsd_dim_stride(tensor, tensor_format, backend_name=_BACKEND_NAME) def _bhsd_graph_tensor(graph, tensor: torch.Tensor, tensor_format: str): """Create a cuDNN graph tensor with BHSD dims and TE-layout strides.""" - dim, stride = _bhsd_dim_stride(tensor, tensor_format) - return graph.tensor(dim=dim, stride=stride, data_type=tensor.dtype) + return cudnn_pygraph.bhsd_graph_tensor(graph, tensor, tensor_format, backend_name=_BACKEND_NAME) # score_mod graph cache helpers. @@ -194,40 +182,13 @@ def _wrapped_score_mod(sdpa_graph, score_tensor): def _get_cudnn_current_stream_handle(cudnn, device: torch.device): """Return a cuDNN handle for device, bound to PyTorch's current stream.""" - if device.type != "cuda": - raise ValueError(f"Flex Attention only supports CUDA tensors, got device {device}.") - if device.index is None: - device = torch.device("cuda", torch.cuda.current_device()) - - handle = _cudnn_score_mod_handles.get(device) - with torch.cuda.device(device): - if handle is None: - handle = cudnn.create_handle() - _cudnn_score_mod_handles[device] = handle - - stream = torch.cuda.current_stream(device).cuda_stream - cudnn.set_stream(handle=handle, stream=stream) - return handle + del cudnn # the shared module owns the import + return cudnn_pygraph.handle_for(device, backend_name=_BACKEND_NAME) def _build_cudnn_pygraph(dtype: torch.dtype, device: torch.device): """Create a cuDNN frontend Python graph for F16/BF16 SDPA.""" - cudnn = _import_cudnn_frontend() - - if dtype == torch.float16: - io_data_type = cudnn.data_type.HALF - elif dtype == torch.bfloat16: - io_data_type = cudnn.data_type.BFLOAT16 - else: - raise ValueError(f"Flex Attention only supports FP16/BF16 tensors, got {dtype}.") - - graph = cudnn.pygraph( - io_data_type=io_data_type, - intermediate_data_type=cudnn.data_type.FLOAT, - compute_data_type=cudnn.data_type.FLOAT, - handle=_get_cudnn_current_stream_handle(cudnn, device), - ) - return graph + return cudnn_pygraph.build_pygraph(dtype, device, backend_name=_BACKEND_NAME) @dataclass @@ -265,17 +226,8 @@ class _CudnnScoreModBwdGraphEntry: def _finalize_cudnn_graph(graph) -> int: """Build a cuDNN frontend Python graph and return its workspace size.""" - cudnn = _import_cudnn_frontend() - - graph.validate() - graph.build_operation_graph() - try: - graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) - graph.check_support() - except cudnn.cudnnGraphNotSupportedError as exc: - raise RuntimeError(f"cuDNN Flex Attention SDPA graph is not supported: {exc}") from exc - graph.build_plans(cudnn.build_plan_policy.HEURISTICS_CHOICE) - return max(graph.get_workspace_size(), 1) + workspace_size, _ = cudnn_pygraph.finalize_plans(graph, backend_name=_BACKEND_NAME) + return workspace_size def _execute_cudnn_graph( diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index c51543c6c8b..15f60a76477 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -23,6 +23,8 @@ import torch from packaging.version import InvalidVersion, Version as PkgVersion +from . import cudnn_pygraph + __all__ = [ "is_frost_attention_available", "is_frost_attention_supported", @@ -53,89 +55,19 @@ # here rather than failing later at plan selection. _HEAD_DIM_MULTIPLE = 8 -_cudnn = None -_frost_engines_enabled = False +_BACKEND_NAME = "FrostAttention" _availability: Optional[Tuple[bool, str]] = None _PLAN_CACHE: dict = {} -_HANDLES: Dict[torch.device, Any] = {} def _import_cudnn_frontend(enable_frost_engines: bool = True): - """Import cuDNN Frontend, enabling the FROST engines if this caller needs them. - - ``enable_frost_engines`` is not merely additive: the switch also ranks FROST ahead of the - backend engines everywhere, so a caller that does not want FROST must not ask for it. - - The enabling is deliberately outside the import memo. Both backends call this, and whichever - one reaches it first would otherwise decide for the process: with the flag inside the memo, a - flex call would cache the module with FROST off and every later FROST call would get a cuDNN - that offers no FROST engine, which surfaces much later as "no cuDNN engine matching ... was - offered". Enabling late is sound because the switch is read per graph rather than at import: - in cuDNN Frontend 1.29.0 ``engines/manifest.py`` consults the environment inside - ``offered_ids()``, reached from ``engines_for(graph)`` on every ``create_execution_plans``. - - Note the switch is process-wide and never unset, so enabling it for FROST also reorders the - candidates a concurrent score_mod graph sees. Callers that require a particular engine should - verify by plan name rather than rely on the switch, which is what - ``_finalize_plans(require_plan_token=...)`` does. - """ - global _cudnn, _frost_engines_enabled # pylint: disable=global-statement - if _cudnn is None: - try: - import cudnn # pylint: disable=import-outside-toplevel - except ImportError as exc: - raise ImportError( - "cuDNN frontend Python package not found. " - "Install it with: pip install nvidia-cudnn-frontend" - ) from exc - - _cudnn = cudnn - - if enable_frost_engines and not _frost_engines_enabled: - os.environ.setdefault("CUDNN_FRONTEND_ENABLE_FROST_ENGINES", "1") - # pylint: disable=import-outside-toplevel,unused-import - import cudnn.sdpa # noqa: F401 - - _frost_engines_enabled = True - - return _cudnn + """Import cuDNN Frontend with the FROST engines on, which is what this backend needs. - -def _handle_for(device: torch.device, *, backend_name: str = "FrostAttention"): - """A cuDNN handle for ``device``, rebound to PyTorch's current stream on every call. - - Without the rebinding, cuDNN runs on its handle's own stream while the tensors and workspace - are allocated on PyTorch's current stream, and nothing orders the two. That is not - hypothetical: the p2p context-parallel ring issues attention inside - ``with torch.cuda.stream(cp_stream)``, so on alternating ring steps the kernel and its buffers - would otherwise be on different streams. The same cached plan is executed from different - streams across steps, so this has to happen per call rather than once per handle. + The default differs from the shared module's, where it is off. Every use site here wants the + engines; a caller that does not must not ask for them, because the switch is process-wide. + See ``cudnn_pygraph.import_cudnn_frontend`` for why the enabling sits outside the import memo. """ - if device.type != "cuda": - raise ValueError(f"{backend_name} requires CUDA tensors; got device {device}") - cudnn = _cudnn if _cudnn is not None else _import_cudnn_frontend() - if device.index is None: - device = torch.device("cuda", torch.cuda.current_device()) - with torch.cuda.device(device): - handle = _HANDLES.get(device) - if handle is None: - handle = cudnn.create_handle() - _HANDLES[device] = handle - cudnn.set_stream(handle=handle, stream=torch.cuda.current_stream(device).cuda_stream) - return handle - - -def _build_pygraph( - dtype: torch.dtype, device: torch.device, *, backend_name: str = "FrostAttention" -): - """A cuDNN frontend graph for F16/BF16 SDPA, bound to this device's stream-current handle.""" - cudnn = _cudnn if _cudnn is not None else _import_cudnn_frontend() - return cudnn.pygraph( - io_data_type=_cudnn_dtype(dtype), - intermediate_data_type=cudnn.data_type.FLOAT, - compute_data_type=cudnn.data_type.FLOAT, - handle=_handle_for(device, backend_name=backend_name), - ) + return cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=enable_frost_engines) def _diagonal_band_kwargs(cudnn, attn_mask_type: str, window: Tuple[int, int]) -> Dict[str, Any]: @@ -165,80 +97,6 @@ def _diagonal_band_kwargs(cudnn, attn_mask_type: str, window: Tuple[int, int]) - return opts -def _finalize_plans( - graph, - *, - heuristics: Optional[Sequence[Any]] = None, - build_policy: Any = None, - require_plan_token: Optional[str] = None, - not_found_hint: Any = "", -) -> Tuple[int, Optional[str]]: - """Create plans, optionally pin one by name, build, and return (workspace size, plan name). - - ``require_plan_token`` makes the choice strict: only a plan whose name contains the token is - acceptable, and anything else raises. That is not a stylistic preference. Without a pin, - ``build_plans`` walks the ranked list from index 0 and finalizes the first plan that builds, - logging each decline at INFO, so a graph that the intended engine declines runs on whatever - cuDNN ranked next with nothing in the return value to say so. At head_dim 512 that matters in - the forward, where an ordinary engine may well build and compute a different function from the - FROST kernel. The backward is self-limiting, since no non-FROST d512 backward exists, so an - unpinned backward would fail loudly on its own. - - The token is matched as a substring rather than by equality on purpose: cuDNN has already - collapsed per-head-dim engine names (``..._d512`` and friends) into a single row once, and the - substring test survived that. - - - Pinning also changes what ``check_support`` means. Selecting a plan sets cuDNN's internal - ``_plan_pinned``, and only then is a decline fatal; unpinned, cuDNN records the decline and - keeps walking. So the pin has to come first both because the check is scoped to the selected - plan and because it is what makes the check binding at all. - """ - cudnn = _cudnn if _cudnn is not None else _import_cudnn_frontend() - - graph.validate() - graph.build_operation_graph() - - if heuristics is None: - heuristics = [cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK] - - if require_plan_token is None: - try: - graph.create_execution_plans(list(heuristics)) - graph.check_support() - except cudnn.cudnnGraphNotSupportedError as exc: - raise RuntimeError(f"cuDNN SDPA graph is not supported: {exc}") from exc - if build_policy is None: - build_policy = cudnn.build_plan_policy.HEURISTICS_CHOICE - graph.build_plans(build_policy) - return max(graph.get_workspace_size(), 1), None - - graph.create_execution_plans(list(heuristics)) - names = [graph.get_plan_name_at_index(i) for i in range(graph.get_execution_plan_count())] - hits = [i for i, n in enumerate(names) if require_plan_token in n] - if not hits: - # Callable hints are resolved only here: a caller may want to look up package versions to - # explain the failure, and that work should not happen on the success path. - hint = not_found_hint() if callable(not_found_hint) else not_found_hint - raise RuntimeError( - f"no cuDNN engine matching {require_plan_token!r} was offered." - f" Candidate plans: {names[:6]}.{(' ' + hint) if hint else ''}" - ) - graph.select_plan(hits[0]) - # The engine is pinned, so a decline here is its own verdict and cuDNN puts the reason in the - # exception. Surface it: a plan offered and then refused is the harder failure to read. - try: - graph.check_support() - graph.build_plans() - except cudnn.cudnnGraphNotSupportedError as exc: - hint = not_found_hint() if callable(not_found_hint) else not_found_hint - raise RuntimeError( - f"cuDNN engine {names[hits[0]]!r} was offered but declined this graph:" - f" {exc}{(' ' + hint) if hint else ''}" - ) from exc - return max(graph.get_workspace_size(), 1), names[hits[0]] - - def _device_from_key(device_key) -> torch.device: """Rebuild the torch.device that _key recorded, for building under the right device.""" kind, index = device_key @@ -302,7 +160,7 @@ def _no(reason): # Decline only on positive evidence: a version below a floor, or a package absent outright. # An unparseable version defers to _select_frost_plan, which checks the plan by name. - frontend, frontend_raw = _pkg_version("nvidia-cudnn-frontend", _cudnn) + frontend, frontend_raw = _pkg_version("nvidia-cudnn-frontend", cudnn_pygraph.cudnn_module()) if frontend is not None and frontend < _MIN_CUDNN_FRONTEND: return _no( f"nvidia-cudnn-frontend {frontend_raw} registers no sm100 backward engine; >=" @@ -569,14 +427,6 @@ def from_frost_layout(t: torch.Tensor, qkv_format: str) -> torch.Tensor: ) -def _cudnn_dtype(dtype: torch.dtype): - cudnn = _import_cudnn_frontend() - return { - torch.bfloat16: cudnn.data_type.BFLOAT16, - torch.float16: cudnn.data_type.HALF, - }[dtype] - - def _check_layout(name: str, t: torch.Tensor) -> None: """Validate a [b, h, s, d] view. @@ -658,15 +508,16 @@ def hint(): return ( f"Wanted the FROST {what} engine." " nvidia-cudnn-frontend=" - f"{_pkg_version('nvidia-cudnn-frontend', _cudnn)[1] or 'unknown'}" + f"{_pkg_version('nvidia-cudnn-frontend', cudnn_pygraph.cudnn_module())[1] or 'unknown'}" f" (floor {_MIN_CUDNN_FRONTEND})," f" nvidia-cutlass-dsl={_pkg_version('nvidia-cutlass-dsl')[1] or 'unknown'}" f" (floor {_MIN_CUTLASS_DSL})." ) cudnn = _import_cudnn_frontend() - _, name = _finalize_plans( + _, name = cudnn_pygraph.finalize_plans( graph, + backend_name=_BACKEND_NAME, heuristics=[cudnn.heur_mode.A], require_plan_token=token, not_found_hint=hint, @@ -683,7 +534,9 @@ def _build_fwd(key) -> dict: shq, shk, shv = [b, hq, sq, d], [b, hkv, skv, d], [b, hkv, skv, d_v] sho, o_stride = _o_shape_stride(b, hq, sq, d, d_v, qs) - graph = _build_pygraph(dtype, _device_from_key(_device), backend_name="FrostAttention") + graph = cudnn_pygraph.build_pygraph( + dtype, _device_from_key(_device), backend_name=_BACKEND_NAME + ) tq = graph.tensor(name="q", dim=shq, stride=list(qs)) tk = graph.tensor(name="k", dim=shk, stride=list(ks)) tv = graph.tensor(name="v", dim=shv, stride=list(vs)) @@ -713,11 +566,13 @@ def _build_bwd(key) -> dict: """Build (and JIT-compile) a backward graph. Expensive; always reached through the cache.""" cudnn = _import_cudnn_frontend() *_device, b, hq, hkv, sq, skv, d, d_v, dtype, mask, scale, qs, ks, vs, deterministic = key - io_dt = _cudnn_dtype(dtype) + io_dt = cudnn_pygraph.io_data_type(cudnn, dtype, backend_name=_BACKEND_NAME) shq, shk, shv = [b, hq, sq, d], [b, hkv, skv, d], [b, hkv, skv, d_v] sho, o_stride = _o_shape_stride(b, hq, sq, d, d_v, qs) - graph = _build_pygraph(dtype, _device_from_key(_device), backend_name="FrostAttention") + graph = cudnn_pygraph.build_pygraph( + dtype, _device_from_key(_device), backend_name=_BACKEND_NAME + ) handles = {} # Each grad is declared with the layout of the tensor it differentiates. for name, shape, stride in ( @@ -845,7 +700,9 @@ def frost_attn_fwd( lse = torch.empty(b, hq, sq, 1, device=q.device, dtype=torch.float32) workspace = torch.empty(entry["workspace"], device=q.device, dtype=torch.uint8) entry["graph"].execute( - {tq: q, tk: k, tv: v, tout: out, tlse: lse}, workspace, handle=_handle_for(q.device) + {tq: q, tk: k, tv: v, tout: out, tlse: lse}, + workspace, + handle=cudnn_pygraph.handle_for(q.device, backend_name=_BACKEND_NAME), ) return out, lse.squeeze(-1) @@ -925,7 +782,7 @@ def _as(t, stride): h["dv"]: dv, }, workspace, - handle=_handle_for(q.device), + handle=cudnn_pygraph.handle_for(q.device, backend_name=_BACKEND_NAME), ) return dq, dk, dv From 2e6787d981e03d4ae70bd79db931c1e2dccd59b5 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 22:14:47 -0700 Subject: [PATCH 68/97] refactor(attention): describe the FROST tensors instead of permuting them cuDNN takes dims and strides, so a bshd or sbhd tensor can be described in its logical BHSD order without being permuted. flex_attention.py already did this; frost permuted into [b, h, s, d] on the way in and back out again. The permute was always a no-op view, but it meant two layout conventions in one file and a pair of helpers to convert between them. frost_attn_fwd/bwd now take TE's qkv_format and hand the descriptors to cudnn_pygraph.bhsd_dim_stride; to_frost_layout and from_frost_layout are gone, along with the round trip through them in the fused shims. Outputs and gradients are allocated in the caller's format, so nothing is converted on the way out either. Verified equivalent rather than assumed: the cache key is byte-identical to the pre-change derivation for both formats, and the graph node's BHSD descriptor equals the BHSD view of the tensor actually allocated. Those were the two places a mistake would have been silent. Two things that had to move with the helpers: - to_frost_layout carried the thd rejection, which now sits in _qkv_format_from_layout with its reason intact. - the backward shim described o and dO with their own o_format/do_format. They now share qkv_format, so a divergence would silently mis-describe them; it is checked, matching what the forward shim already did. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_frost_attention.py | 97 +++++---- .../dot_product_attention/frost_attention.py | 190 +++++++++--------- 2 files changed, 154 insertions(+), 133 deletions(-) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 2d2763b4062..548cc585f74 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -74,6 +74,12 @@ def _shape_id(s): return "b%d_hq%d_hkv%d_sq%d_skv%d_d%d_dv%d" % s +def _bhsd(t): + """A [b, h, s, d] view of a bshd tensor. The reference works in that order; the kernel does + not, since it takes TE's format and reorders the cuDNN descriptors instead.""" + return t.permute(0, 2, 1, 3) + + def _reference(q, k, v, scale, mask, window=None): """Attention in float64, computed independently of TE and of cuDNN. @@ -139,19 +145,18 @@ def test_frost_forward_matches_reference(shape, mask, dtype): b, hq, hkv, sq, skv, d, d_v = shape torch.manual_seed(0) # Generate in fp32 so there is a true high-precision original to measure against, then cast - # for the kernel. [b, h, s, d] views over bshd-contiguous memory is what the backend consumes. - # A bshd VIEW, which is what the backend receives: to_frost_layout permutes a bshd-contiguous - # tensor and hands the result over without a copy. Materialising with .contiguous() here would - # produce bhsd strides instead and leave the stride-keyed plan cache untested. - mk = lambda s_, h_, d_: torch.randn(b, s_, h_, d_, device="cuda").permute(0, 2, 1, 3) + # for the kernel. bshd is what the backend takes now: it is never permuted, only described. + mk = lambda s_, h_, d_: torch.randn(b, s_, h_, d_, device="cuda") q32, k32, v32 = mk(sq, hq, d), mk(skv, hkv, d), mk(skv, hkv, d_v) q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) scale = 1.0 / math.sqrt(d) - out, lse = frost_attn_fwd(q, k, v, attn_scale=scale, attn_mask_type=mask) + out, lse = frost_attn_fwd(q, k, v, "bshd", attn_scale=scale, attn_mask_type=mask) - floor_o, floor_l, ref_o, ref_lse = _floor(q32, k32, v32, scale, mask, dtype) - err_o = (out.double() - ref_o).abs().max().item() + floor_o, floor_l, ref_o, ref_lse = _floor( + _bhsd(q32), _bhsd(k32), _bhsd(v32), scale, mask, dtype + ) + err_o = (_bhsd(out).double() - ref_o).abs().max().item() err_l = (lse.double() - ref_lse).abs().max().item() assert torch.isfinite(out).all(), "forward produced non-finite values" @@ -193,18 +198,17 @@ def test_frost_sliding_window_matches_reference(mask, window, sq, skv): b, hq, hkv, d = 2, 8, 4, 512 dtype = torch.bfloat16 torch.manual_seed(0) - # A bshd VIEW, which is what the backend receives: to_frost_layout permutes a bshd-contiguous - # tensor and hands the result over without a copy. Materialising with .contiguous() here would - # produce bhsd strides instead and leave the stride-keyed plan cache untested. - mk = lambda s_, h_: torch.randn(b, s_, h_, d, device="cuda").permute(0, 2, 1, 3) + mk = lambda s_, h_: torch.randn(b, s_, h_, d, device="cuda") q32, k32, v32 = mk(sq, hq), mk(skv, hkv), mk(skv, hkv) q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) scale = 1.0 / math.sqrt(d) - out, _ = frost_attn_fwd(q, k, v, attn_scale=scale, attn_mask_type=mask, window_size=window) + out, _ = frost_attn_fwd( + q, k, v, "bshd", attn_scale=scale, attn_mask_type=mask, window_size=window + ) - floor_o, _, ref_o, _ = _floor(q32, k32, v32, scale, mask, dtype, window) - err = (out.double() - ref_o).abs().max().item() + floor_o, _, ref_o, _ = _floor(_bhsd(q32), _bhsd(k32), _bhsd(v32), scale, mask, dtype, window) + err = (_bhsd(out).double() - ref_o).abs().max().item() assert torch.isfinite(out).all(), "sliding-window forward produced non-finite values" assert err <= 2 * floor_o + 1e-3, "out err %.3e exceeds 2x the floor %.3e for window %s" % ( err, @@ -214,7 +218,7 @@ def test_frost_sliding_window_matches_reference(mask, window, sq, skv): # A window must actually change the result; if the bound were dropped this would match the # unwindowed output and the check above would still pass. - full, _ = frost_attn_fwd(q, k, v, attn_scale=scale, attn_mask_type=mask) + full, _ = frost_attn_fwd(q, k, v, "bshd", attn_scale=scale, attn_mask_type=mask) assert not torch.equal(out, full), "window %s produced the same output as no window" % (window,) @@ -237,27 +241,31 @@ def test_frost_backward_matches_reference(shape, mask, window, dtype): b, hq, hkv, sq, skv, d, d_v = shape torch.manual_seed(0) - # A bshd VIEW, which is what the backend receives: to_frost_layout permutes a bshd-contiguous - # tensor and hands the result over without a copy. Materialising with .contiguous() here would - # produce bhsd strides instead and leave the stride-keyed plan cache untested. - mk = lambda s_, h_, d_: torch.randn(b, s_, h_, d_, device="cuda").permute(0, 2, 1, 3) + mk = lambda s_, h_, d_: torch.randn(b, s_, h_, d_, device="cuda") q32, k32, v32 = mk(sq, hq, d), mk(skv, hkv, d), mk(skv, hkv, d_v) q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) scale = 1.0 / math.sqrt(d) - out, lse = frost_attn_fwd(q, k, v, attn_scale=scale, attn_mask_type=mask, window_size=window) + out, lse = frost_attn_fwd( + q, k, v, "bshd", attn_scale=scale, attn_mask_type=mask, window_size=window + ) dout = torch.randn_like(out) dq, dk, dv = frost_attn_bwd( - q, k, v, out, lse, dout, attn_scale=scale, attn_mask_type=mask, window_size=window + q, k, v, out, lse, dout, "bshd", attn_scale=scale, attn_mask_type=mask, window_size=window ) - qr = q32.detach().clone().requires_grad_(True) - kr = k32.detach().clone().requires_grad_(True) - vr = v32.detach().clone().requires_grad_(True) + # The reference works in [b, h, s, d], so it takes views and returns grads in that order. + qr = _bhsd(q32).detach().clone().requires_grad_(True) + kr = _bhsd(k32).detach().clone().requires_grad_(True) + vr = _bhsd(v32).detach().clone().requires_grad_(True) ref_o, _ = _reference(qr, kr, vr, scale, mask, window) - ref_o.backward(dout.double()) + ref_o.backward(_bhsd(dout).double()) - for name, got, want in (("dq", dq, qr.grad), ("dk", dk, kr.grad), ("dv", dv, vr.grad)): + for name, got, want in ( + ("dq", _bhsd(dq), qr.grad), + ("dk", _bhsd(dk), kr.grad), + ("dv", _bhsd(dv), vr.grad), + ): assert torch.isfinite(got).all(), "%s has non-finite values" % name assert got.shape == want.shape, "%s shape %s != %s" % (name, got.shape, want.shape) err = (got.double() - want).abs().max().item() @@ -463,22 +471,25 @@ def test_frost_rejects_mismatched_kv(): b, h, s, d = 2, 4, 512, 512 dtype = torch.bfloat16 - mk = lambda hh: torch.randn(b, s, hh, d, device="cuda", dtype=dtype).permute(0, 2, 1, 3) - q, k = mk(h).contiguous(), mk(h).contiguous() + mk = lambda hh: torch.randn(b, s, hh, d, device="cuda", dtype=dtype) + q, k = mk(h), mk(h) with pytest.raises(ValueError, match="batch, heads and seqlen"): - frost_attn_fwd(q, k, mk(h * 2).contiguous()) + frost_attn_fwd(q, k, mk(h * 2), "bshd") with pytest.raises(ValueError, match="match q"): - frost_attn_fwd(q, k, k.to(torch.float32)) + frost_attn_fwd(q, k, k.to(torch.float32), "bshd") @requires_frost def test_frost_serves_v_with_its_own_head_dim_and_layout(): """v is keyed and declared separately, the way flex_attention keys each tensor. - Checked against the float64 reference rather than against another FROST call: a v whose - head_dim AND stride order both differ from k's is exactly the case that a plan built from - k alone would compute wrongly without raising. + Both halves of that are exercised: v carries its own head_dim, and its strides differ from + k's because it is a non-contiguous slice of a wider buffer rather than a fresh allocation. + A plan built from k alone would compute either case wrongly without raising. + + v cannot differ from k in qkv_format: one format describes all three, which is what the + fused path produces and what the selector enforces. """ from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( frost_attn_fwd, @@ -487,17 +498,21 @@ def test_frost_serves_v_with_its_own_head_dim_and_layout(): b, h, s, d, d_v = 2, 4, 512, 512, 320 dtype = torch.bfloat16 torch.manual_seed(0) - q32 = torch.randn(b, s, h, d, device="cuda").permute(0, 2, 1, 3) - k32 = torch.randn(b, s, h, d, device="cuda").permute(0, 2, 1, 3) - # sbhd rather than bshd, so v's stride order differs from k's as well as its head_dim. - v32 = torch.randn(s, b, h, d_v, device="cuda").permute(1, 2, 0, 3) + q32 = torch.randn(b, s, h, d, device="cuda") + k32 = torch.randn(b, s, h, d, device="cuda") + # A slice of a wider buffer, so v's strides are its own rather than k's shape re-derived. + v32 = torch.randn(b, s, h, d_v + 64, device="cuda")[..., :d_v] q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) - assert v.stride()[:3] != k.stride()[:3], "v must not share k's stride order here" + assert v.stride()[:3] != k.stride()[:3], "v must not share k's strides here" + assert v.stride(3) == 1, "the head dim must stay contiguous" scale = 1.0 / math.sqrt(d) - out, lse = frost_attn_fwd(q, k, v, attn_scale=scale, attn_mask_type="causal") + out, lse = frost_attn_fwd(q, k, v, "bshd", attn_scale=scale, attn_mask_type="causal") + out = _bhsd(out) - floor_o, floor_l, ref_o, ref_lse = _floor(q32, k32, v32, scale, "causal", dtype) + floor_o, floor_l, ref_o, ref_lse = _floor( + _bhsd(q32), _bhsd(k32), _bhsd(v32), scale, "causal", dtype + ) assert out.shape == (b, h, s, d_v), "out takes v's head_dim; got %s" % (tuple(out.shape),) assert out.stride(3) == 1, "out must stay head-contiguous; got stride %s" % (out.stride(),) err_o = (out.double() - ref_o).abs().max().item() diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 15f60a76477..f71d5ef8c8d 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -32,8 +32,6 @@ "fused_attn_bwd", "frost_attn_fwd", "frost_attn_bwd", - "to_frost_layout", - "from_frost_layout", ] @@ -238,7 +236,14 @@ def _qkv_format_from_layout(qkv_layout: str) -> str: raise NotImplementedError( f"FROST attention needs q, k and v in one format; got qkv_layout {qkv_layout!r}" ) - return formats.pop() + qkv_format = formats.pop() + # Carried over from the permute helper this replaced, so the thd decline keeps its reason. + if qkv_format not in _SUPPORTED_QKV_FORMATS: + raise NotImplementedError( + f"FROST attention supports qkv_format in {_SUPPORTED_QKV_FORMATS}; got" + f" {qkv_format!r}. thd needs varlen support that is not implemented here." + ) + return qkv_format def _te_mask_spec(attn_mask_type: str, window_size, bottom_right_diagonal: bool): @@ -399,34 +404,6 @@ def is_frost_attention_supported(params) -> Tuple[int, str]: return int(FusedAttnBackend.FROST), "" -def to_frost_layout(t: torch.Tensor, qkv_format: str) -> torch.Tensor: - """View a tensor in TE's qkv_format as [b, h, s, d]. - - No copy: the cuDNN graphs are built from each tensor's actual strides, so both bshd and - sbhd are served directly. sbhd matters because that is what Megatron uses internally, and - transposing into bshd on every call would copy the whole tensor. - """ - if qkv_format == "bshd": # [b, s, h, d] -> [b, h, s, d] - return t.permute(0, 2, 1, 3) - if qkv_format == "sbhd": # [s, b, h, d] -> [b, h, s, d] - return t.permute(1, 2, 0, 3) - raise NotImplementedError( - f"FROST attention supports qkv_format 'bshd' and 'sbhd'; got {qkv_format!r}. thd needs" - " varlen support that is not implemented here." - ) - - -def from_frost_layout(t: torch.Tensor, qkv_format: str) -> torch.Tensor: - """Inverse of to_frost_layout.""" - if qkv_format == "bshd": # [b, h, s, d] -> [b, s, h, d] - return t.permute(0, 2, 1, 3) - if qkv_format == "sbhd": # [b, h, s, d] -> [s, b, h, d] - return t.permute(2, 0, 1, 3) - raise NotImplementedError( - f"FROST attention supports qkv_format 'bshd' and 'sbhd'; got {qkv_format!r}." - ) - - def _check_layout(name: str, t: torch.Tensor) -> None: """Validate a [b, h, s, d] view. @@ -483,13 +460,16 @@ def _head_dim_strides(shape: Sequence[int], ref_strides: Sequence[int]) -> list: return strides -def _o_shape_stride(b, hq, sq, d, d_v, qs): - """Shape and strides for an O-shaped tensor: q's layout, v's head_dim. +def _o_shape_stride(shape, d_v, ref_strides): + """Shape and strides for an O-shaped tensor: ``shape``'s layout carrying v's head_dim. - Symmetric head dims keep q's exact strides, which preserves a caller's non-dense view. + Works in either space. The head dim is last in both TE's bshd/sbhd and cuDNN's BHSD, and the + rule only reorders by stride magnitude, so the graph node and the allocation can each apply it + in their own space and still agree. Equal head dims keep the reference strides untouched, + which preserves a caller's non-dense view. """ - shape = [b, hq, sq, d_v] - return shape, (list(qs) if d_v == d else _head_dim_strides(shape, qs)) + out = list(shape[:3]) + [d_v] + return out, (list(ref_strides) if d_v == shape[3] else _head_dim_strides(out, ref_strides)) def _select_frost_plan(graph, token: str, what: str): @@ -532,7 +512,7 @@ def _build_fwd(key) -> dict: # forward so the two never split the forward cache. *_device, b, hq, hkv, sq, skv, d, d_v, dtype, mask, scale, qs, ks, vs, _deterministic = key shq, shk, shv = [b, hq, sq, d], [b, hkv, skv, d], [b, hkv, skv, d_v] - sho, o_stride = _o_shape_stride(b, hq, sq, d, d_v, qs) + sho, o_stride = _o_shape_stride([b, hq, sq, d], d_v, qs) graph = cudnn_pygraph.build_pygraph( dtype, _device_from_key(_device), backend_name=_BACKEND_NAME @@ -568,7 +548,7 @@ def _build_bwd(key) -> dict: *_device, b, hq, hkv, sq, skv, d, d_v, dtype, mask, scale, qs, ks, vs, deterministic = key io_dt = cudnn_pygraph.io_data_type(cudnn, dtype, backend_name=_BACKEND_NAME) shq, shk, shv = [b, hq, sq, d], [b, hkv, skv, d], [b, hkv, skv, d_v] - sho, o_stride = _o_shape_stride(b, hq, sq, d, d_v, qs) + sho, o_stride = _o_shape_stride([b, hq, sq, d], d_v, qs) graph = cudnn_pygraph.build_pygraph( dtype, _device_from_key(_device), backend_name=_BACKEND_NAME @@ -628,29 +608,38 @@ def _cached(kind: str, key): return entry -def _key(q, k, v, mask, scale, deterministic=False): +def _bhsd(t: torch.Tensor, qkv_format: str): + """``t`` described in cuDNN's logical BHSD, without permuting it.""" + return cudnn_pygraph.bhsd_dim_stride(t, qkv_format, backend_name=_BACKEND_NAME) + + +def _key(q, k, v, qkv_format, mask, scale, deterministic=False): + qd, qs = _bhsd(q, qkv_format) + kd, ks = _bhsd(k, qkv_format) + vd, vs = _bhsd(v, qkv_format) return ( # Built under whichever device was current, so it must not be reused on another. Matches # the C++ fused-attn cache, which keys on device_id. Type too, so CPU cannot alias cuda:0. q.device.type, q.device.index, - q.shape[0], - q.shape[1], - k.shape[1], - q.shape[2], - k.shape[2], - q.shape[3], + qd[0], + qd[1], + kd[1], + qd[2], + kd[2], + qd[3], # v carries its own head_dim and strides, the way flex_attention keys each tensor # separately. Without them an asymmetric v would reuse a plan built for k's shape. - v.shape[3], + vd[3], q.dtype, mask, float(scale), # Strides are part of the plan: the graph is built for this exact layout, which is what - # lets bshd and sbhd both run without a transpose. - tuple(q.stride()), - tuple(k.stride()), - tuple(v.stride()), + # lets bshd and sbhd both run without a transpose. qkv_format does not need its own key + # entry, since two formats producing the same BHSD description are the same graph. + tuple(qs), + tuple(ks), + tuple(vs), # The deterministic backward is a different algorithm, not a flag on the same one, so a # plan built either way must not be handed to a call that asked for the other. bool(deterministic), @@ -661,39 +650,43 @@ def frost_attn_fwd( q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, + qkv_format: str = "bshd", attn_scale: Optional[float] = None, attn_mask_type: str = "causal", window_size: Optional[Tuple[int, int]] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: """Forward attention via cuDNN FROST. - q, k, v are [b, h, s, d] views; bshd and sbhd are both served, since the graph is built from - each tensor's actual strides. GQA is supported directly (h_kv may differ from h_q), SQ need + q, k, v are in TE's ``qkv_format`` and are never permuted: cuDNN takes dims and strides, so + the descriptors are reordered into its logical BHSD instead. That is what serves bshd and + sbhd alike without a transpose. GQA is supported directly (h_kv may differ from h_q), SQ need not equal SKV, which is what lets a CP ring step use this, and v may carry its own head_dim, - in which case out follows q's layout with v's head_dim. Returns (out, softmax_lse) - with softmax_lse as [b, h, s] fp32 natural-log logsumexp, the layout and convention the CP - ring correction expects. + in which case out follows q's layout with v's head_dim. ``out`` comes back in ``qkv_format``; + softmax_lse is [b, h, s] fp32 natural-log logsumexp, the layout and convention the CP ring + correction expects, and is BHSD regardless of the input format. """ for name, tensor in (("q", q), ("k", k), ("v", v)): _check_layout(name, tensor) _check_dtype(name, tensor, q.dtype) _check_kv_match(k, v) - if k.shape[0] != q.shape[0] or k.shape[3] != q.shape[3]: + qd, _ = _bhsd(q, qkv_format) + kd, _ = _bhsd(k, qkv_format) + vd, _ = _bhsd(v, qkv_format) + if kd[0] != qd[0] or kd[3] != qd[3]: # The graph declares k and v with q's batch and head_dim, so a mismatch would bind a # differently shaped buffer to that node and read the wrong elements silently. - raise ValueError(f"k must match q in batch and head_dim; got q {q.shape} and k {k.shape}") - if q.shape[1] % k.shape[1] != 0: - raise ValueError( - f"num_heads must be divisible by num_gqa_groups; got {q.shape[1]} and {k.shape[1]}" - ) + raise ValueError(f"k must match q in batch and head_dim; got q {qd} and k {kd} in BHSD") + if qd[1] % kd[1] != 0: + raise ValueError(f"num_heads must be divisible by num_gqa_groups; got {qd[1]} and {kd[1]}") mask = _mask_spec(attn_mask_type, window_size) - scale = attn_scale if attn_scale is not None else q.shape[-1] ** -0.5 - entry = _cached("fwd", _key(q, k, v, mask, scale)) + scale = attn_scale if attn_scale is not None else qd[3] ** -0.5 + entry = _cached("fwd", _key(q, k, v, qkv_format, mask, scale)) tq, tk, tv, tout, tlse = entry["handles"] - b, hq, sq, d = q.shape - out_shape, out_stride = _o_shape_stride(b, hq, sq, d, v.shape[3], q.stride()) + b, hq, sq = qd[0], qd[1], qd[2] + # Allocated in the caller's format, so no permute is needed on the way out either. + out_shape, out_stride = _o_shape_stride(q.shape, vd[3], q.stride()) # Allocated per call so concurrent uses cannot alias; the cache holds only the plan. # empty_strided, not empty_like: the latter does not preserve an arbitrary permuted stride. out = torch.empty_strided(out_shape, out_stride, device=q.device, dtype=q.dtype) @@ -714,39 +707,47 @@ def frost_attn_bwd( out: torch.Tensor, softmax_lse: torch.Tensor, dout: torch.Tensor, + qkv_format: str = "bshd", attn_scale: Optional[float] = None, attn_mask_type: str = "causal", deterministic: bool = False, window_size: Optional[Tuple[int, int]] = None, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Backward attention via cuDNN FROST. `softmax_lse` is [b, h, s] as returned by the forward.""" + """Backward attention via cuDNN FROST. + + Tensors are in TE's ``qkv_format``, as in the forward. ``softmax_lse`` is [b, h, s] BHSD, as + the forward returned it. The gradients come back in ``qkv_format``. + """ for name, tensor in (("q", q), ("k", k), ("v", v), ("out", out), ("dout", dout)): _check_layout(name, tensor) _check_dtype(name, tensor, q.dtype) _check_kv_match(k, v) # The same shape assumptions the forward makes, plus o/dO, which the graph declares with q's # shape. The forward runs first in autograd, but the CP ring calls this directly. - if k.shape[0] != q.shape[0] or k.shape[3] != q.shape[3]: - raise ValueError(f"k must match q in batch and head_dim; got q {q.shape} and k {k.shape}") - if q.shape[1] % k.shape[1] != 0: - raise ValueError( - f"num_heads must be divisible by num_gqa_groups; got {q.shape[1]} and {k.shape[1]}" - ) - o_shape, o_stride = _o_shape_stride(*q.shape, v.shape[3], q.stride()) + qd, _ = _bhsd(q, qkv_format) + kd, _ = _bhsd(k, qkv_format) + vd, _ = _bhsd(v, qkv_format) + if kd[0] != qd[0] or kd[3] != qd[3]: + raise ValueError(f"k must match q in batch and head_dim; got q {qd} and k {kd} in BHSD") + if qd[1] % kd[1] != 0: + raise ValueError(f"num_heads must be divisible by num_gqa_groups; got {qd[1]} and {kd[1]}") + o_shape, o_stride = _o_shape_stride(q.shape, vd[3], q.stride()) for name, tensor in (("out", out), ("dout", dout)): if list(tensor.shape) != o_shape: raise ValueError(f"{name} must be shaped {o_shape}; got {list(tensor.shape)}") if softmax_lse.dtype != torch.float32: raise ValueError(f"softmax_lse must be fp32; got {softmax_lse.dtype}") - if tuple(softmax_lse.shape[:3]) != tuple(q.shape[:3]): + # Compared against the BHSD description, not against q's own shape: the LSE is always + # [b, h, s] whatever format the tensors arrived in. + if tuple(softmax_lse.shape[:3]) != tuple(qd[:3]): raise ValueError( f"softmax_lse must be [b, h, s] matching q; got {tuple(softmax_lse.shape)} and" - f" {tuple(q.shape)}" + f" {tuple(qd[:3])}" ) mask = _mask_spec(attn_mask_type, window_size) - scale = attn_scale if attn_scale is not None else q.shape[-1] ** -0.5 - entry = _cached("bwd", _key(q, k, v, mask, scale, deterministic)) + scale = attn_scale if attn_scale is not None else qd[3] ** -0.5 + entry = _cached("bwd", _key(q, k, v, qkv_format, mask, scale, deterministic)) h = entry["handles"] if softmax_lse.dim() == 3: @@ -857,9 +858,10 @@ def fused_attn_fwd( ) out, softmax_lse = frost_attn_fwd( - to_frost_layout(q.contiguous(), qkv_format), - to_frost_layout(k.contiguous(), qkv_format), - to_frost_layout(v.contiguous(), qkv_format), + q.contiguous(), + k.contiguous(), + v.contiguous(), + qkv_format, attn_scale=attn_scale, attn_mask_type=mask_type, window_size=window, @@ -867,7 +869,7 @@ def fused_attn_fwd( # A real tensor, not None: it is saved for backward and handed to the activation offload # hooks, neither of which accepts None. FROST has no dropout, so nothing reads it. rng_state = torch.empty(2, dtype=torch.int64, device=q.device) - return from_frost_layout(out, qkv_format), [softmax_lse, rng_state] + return out, [softmax_lse, rng_state] def fused_attn_bwd( @@ -918,26 +920,30 @@ def fused_attn_bwd( cuda_graph_capture=cuda_graph, ) qkv_format = _qkv_format_from_layout(qkv_layout) + # o and dO used to carry their own format into the permute; they now share qkv_format, so a + # divergence would silently describe them with the wrong strides. The selector already + # declines it, but this is reached directly too. + for name, fmt in (("o_format", o_format), ("do_format", do_format)): + if fmt != qkv_format: + raise NotImplementedError( + f"FROST attention needs {name} to match qkv_format; got {fmt}/{qkv_format}" + ) mask_type, window = _te_mask_spec( attn_mask_type, window_size, _bottom_right_diagonal(attn_mask_type, bottom_right_diagonal) ) softmax_lse = aux_ctx_tensors[0] dq, dk, dv = frost_attn_bwd( - to_frost_layout(q.contiguous(), qkv_format), - to_frost_layout(k.contiguous(), qkv_format), - to_frost_layout(v.contiguous(), qkv_format), - to_frost_layout(o.contiguous(), o_format), + q.contiguous(), + k.contiguous(), + v.contiguous(), + o.contiguous(), softmax_lse, - to_frost_layout(d_o.contiguous(), do_format), + d_o.contiguous(), + qkv_format, attn_scale=attn_scale, attn_mask_type=mask_type, deterministic=deterministic, window_size=window, ) - return ( - from_frost_layout(dq, qkv_format), - from_frost_layout(dk, qkv_format), - from_frost_layout(dv, qkv_format), - None, - ) + return dq, dk, dv, None From 7bc01fbb9eacf15b62e13db62bbbd33993af9288 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 22:22:22 -0700 Subject: [PATCH 69/97] refactor(attention): drop the FROST kernel API, keep the sub-backend one frost_attn_fwd/bwd existed as a layout-agnostic [b, h, s, d] kernel API with fused_attn_fwd/bwd as a thin adapter over it. Once the tensors stopped being permuted that seam had little left to separate: both halves work in TE's format, and the adapter was mostly forwarding. The sub-backend entry points are what cpp_extensions.fused_attn dispatches to and the only ones any caller reaches, so they are what remains. The guards are not glue and all of them move: layout and dtype on every bound tensor, the k/v agreement, the GQA divisibility, the o/dO shape, and the LSE dtype and shape. They are grouped in _validate_qkv so the shims stay readable and so neither direction can quietly acquire a different set. Dropping the double mask translation comes free: the shim normalised through _te_mask_spec and the kernel then re-validated the result with _mask_spec on every call. Verified that the spec _te_mask_spec produces yields byte-identical cuDNN band kwargs across all twelve mask and window combinations the tests cover, so cuDNN sees the same graph. Tests drive the fused signature through two local adapters. One coverage change worth naming: the shim makes its inputs contiguous, so a test can no longer hand the backend a non-dense tensor. That path was never reachable from the dispatcher either, which is the reason the seam was not worth keeping. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_frost_attention.py | 103 ++++--- .../dot_product_attention/frost_attention.py | 269 +++++++----------- 2 files changed, 170 insertions(+), 202 deletions(-) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 548cc585f74..da89f4b3389 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -74,6 +74,66 @@ def _shape_id(s): return "b%d_hq%d_hkv%d_sq%d_skv%d_d%d_dv%d" % s +def _fwd(q, k, v, mask, scale, window=None): + """The forward through the fused signature, which is the only entry point the backend has. + + Everything is bshd here, because the shim derives one qkv_format and makes the tensors + contiguous; that is exactly what the dispatcher hands it in production. + """ + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + fused_attn_fwd, + ) + + out, aux = fused_attn_fwd( + True, + q.shape[1], + k.shape[1], + None, + None, + q, + k, + v, + None, + None, + attn_scale=scale, + qkv_layout="bshd_bshd_bshd", + o_format="bshd", + attn_mask_type=mask, + window_size=(-1, -1) if window is None else window, + ) + return out, aux[0] + + +def _bwd(q, k, v, out, lse, dout, mask, scale, window=None): + """The backward through the fused signature. aux_ctx_tensors is what the forward returned.""" + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + fused_attn_bwd, + ) + + dq, dk, dv, _ = fused_attn_bwd( + q.shape[1], + k.shape[1], + None, + None, + q, + k, + v, + out, + dout, + None, + [lse, torch.empty(2, dtype=torch.int64, device=q.device)], + None, + attn_scale=scale, + qkv_layout="bshd_bshd_bshd", + o_format="bshd", + do_format="bshd", + dqkv_layout="bshd_bshd_bshd", + attn_mask_type=mask, + window_size=(-1, -1) if window is None else window, + ) + return dq, dk, dv + + def _bhsd(t): """A [b, h, s, d] view of a bshd tensor. The reference works in that order; the kernel does not, since it takes TE's format and reorders the cuDNN descriptors instead.""" @@ -138,10 +198,6 @@ def _floor(q32, k32, v32, scale, mask, dtype, window=None): @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) def test_frost_forward_matches_reference(shape, mask, dtype): """Forward output and LSE against an independent float64 reference.""" - from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( - frost_attn_fwd, - ) - b, hq, hkv, sq, skv, d, d_v = shape torch.manual_seed(0) # Generate in fp32 so there is a true high-precision original to measure against, then cast @@ -151,7 +207,7 @@ def test_frost_forward_matches_reference(shape, mask, dtype): q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) scale = 1.0 / math.sqrt(d) - out, lse = frost_attn_fwd(q, k, v, "bshd", attn_scale=scale, attn_mask_type=mask) + out, lse = _fwd(q, k, v, mask, scale) floor_o, floor_l, ref_o, ref_lse = _floor( _bhsd(q32), _bhsd(k32), _bhsd(v32), scale, mask, dtype @@ -189,10 +245,6 @@ def test_frost_sliding_window_matches_reference(mask, window, sq, skv): is worth its own test because a left bound that is off by one, or silently dropped, still produces finite plausible-looking output -- the reference is the only thing that catches it. """ - from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( - frost_attn_fwd, - ) - # The rectangular case is the one that matters for alignment: top-left and bottom-right # coincide when sq == skv, so a swapped alignment is invisible in square shapes. b, hq, hkv, d = 2, 8, 4, 512 @@ -203,9 +255,7 @@ def test_frost_sliding_window_matches_reference(mask, window, sq, skv): q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) scale = 1.0 / math.sqrt(d) - out, _ = frost_attn_fwd( - q, k, v, "bshd", attn_scale=scale, attn_mask_type=mask, window_size=window - ) + out, _ = _fwd(q, k, v, mask, scale, window) floor_o, _, ref_o, _ = _floor(_bhsd(q32), _bhsd(k32), _bhsd(v32), scale, mask, dtype, window) err = (_bhsd(out).double() - ref_o).abs().max().item() @@ -218,7 +268,7 @@ def test_frost_sliding_window_matches_reference(mask, window, sq, skv): # A window must actually change the result; if the bound were dropped this would match the # unwindowed output and the check above would still pass. - full, _ = frost_attn_fwd(q, k, v, "bshd", attn_scale=scale, attn_mask_type=mask) + full, _ = _fwd(q, k, v, mask, scale) assert not torch.equal(out, full), "window %s produced the same output as no window" % (window,) @@ -234,11 +284,6 @@ def test_frost_backward_matches_reference(shape, mask, window, dtype): that would show first -- the gradient of a softmax involves a subtraction of similarly sized terms, so a range problem surfaces there before it surfaces in the forward. """ - from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( - frost_attn_bwd, - frost_attn_fwd, - ) - b, hq, hkv, sq, skv, d, d_v = shape torch.manual_seed(0) mk = lambda s_, h_, d_: torch.randn(b, s_, h_, d_, device="cuda") @@ -246,13 +291,9 @@ def test_frost_backward_matches_reference(shape, mask, window, dtype): q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) scale = 1.0 / math.sqrt(d) - out, lse = frost_attn_fwd( - q, k, v, "bshd", attn_scale=scale, attn_mask_type=mask, window_size=window - ) + out, lse = _fwd(q, k, v, mask, scale, window) dout = torch.randn_like(out) - dq, dk, dv = frost_attn_bwd( - q, k, v, out, lse, dout, "bshd", attn_scale=scale, attn_mask_type=mask, window_size=window - ) + dq, dk, dv = _bwd(q, k, v, out, lse, dout, mask, scale, window) # The reference works in [b, h, s, d], so it takes views and returns grads in that order. qr = _bhsd(q32).detach().clone().requires_grad_(True) @@ -465,19 +506,15 @@ def test_frost_sliding_window_selection_by_cp_comm_type(cp_comm_type, window, ex @requires_frost def test_frost_rejects_mismatched_kv(): """v must index the same KV positions as k. head_dim is free; the rest is not.""" - from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( - frost_attn_fwd, - ) - b, h, s, d = 2, 4, 512, 512 dtype = torch.bfloat16 mk = lambda hh: torch.randn(b, s, hh, d, device="cuda", dtype=dtype) q, k = mk(h), mk(h) with pytest.raises(ValueError, match="batch, heads and seqlen"): - frost_attn_fwd(q, k, mk(h * 2), "bshd") + _fwd(q, k, mk(h * 2), "no_mask", 1.0) with pytest.raises(ValueError, match="match q"): - frost_attn_fwd(q, k, k.to(torch.float32), "bshd") + _fwd(q, k, k.to(torch.float32), "no_mask", 1.0) @requires_frost @@ -491,10 +528,6 @@ def test_frost_serves_v_with_its_own_head_dim_and_layout(): v cannot differ from k in qkv_format: one format describes all three, which is what the fused path produces and what the selector enforces. """ - from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( - frost_attn_fwd, - ) - b, h, s, d, d_v = 2, 4, 512, 512, 320 dtype = torch.bfloat16 torch.manual_seed(0) @@ -507,7 +540,7 @@ def test_frost_serves_v_with_its_own_head_dim_and_layout(): assert v.stride(3) == 1, "the head dim must stay contiguous" scale = 1.0 / math.sqrt(d) - out, lse = frost_attn_fwd(q, k, v, "bshd", attn_scale=scale, attn_mask_type="causal") + out, lse = _fwd(q, k, v, "causal", scale) out = _bhsd(out) floor_o, floor_l, ref_o, ref_lse = _floor( diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index f71d5ef8c8d..6ceaa33c548 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -30,8 +30,6 @@ "is_frost_attention_supported", "fused_attn_fwd", "fused_attn_bwd", - "frost_attn_fwd", - "frost_attn_bwd", ] @@ -608,6 +606,30 @@ def _cached(kind: str, key): return entry +def _validate_qkv(q, k, v, qkv_format): + """Check the tensors the graph will bind, and return their BHSD descriptions. + + These are not stylistic guards. ``execute`` binds raw pointers, so a tensor whose shape, + dtype or layout disagrees with the node it is bound to is reinterpreted rather than rejected. + The context-parallel ring calls the backward outside autograd, so neither direction may + assume the other ran first. + """ + for name, tensor in (("q", q), ("k", k), ("v", v)): + _check_layout(name, tensor) + _check_dtype(name, tensor, q.dtype) + _check_kv_match(k, v) + qd, _ = _bhsd(q, qkv_format) + kd, _ = _bhsd(k, qkv_format) + vd, _ = _bhsd(v, qkv_format) + if kd[0] != qd[0] or kd[3] != qd[3]: + # The graph declares k and v with q's batch and head_dim, so a mismatch would bind a + # differently shaped buffer to that node and read the wrong elements silently. + raise ValueError(f"k must match q in batch and head_dim; got q {qd} and k {kd} in BHSD") + if qd[1] % kd[1] != 0: + raise ValueError(f"num_heads must be divisible by num_gqa_groups; got {qd[1]} and {kd[1]}") + return qd, kd, vd + + def _bhsd(t: torch.Tensor, qkv_format: str): """``t`` described in cuDNN's logical BHSD, without permuting it.""" return cudnn_pygraph.bhsd_dim_stride(t, qkv_format, backend_name=_BACKEND_NAME) @@ -646,148 +668,6 @@ def _key(q, k, v, qkv_format, mask, scale, deterministic=False): ) -def frost_attn_fwd( - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - qkv_format: str = "bshd", - attn_scale: Optional[float] = None, - attn_mask_type: str = "causal", - window_size: Optional[Tuple[int, int]] = None, -) -> Tuple[torch.Tensor, torch.Tensor]: - """Forward attention via cuDNN FROST. - - q, k, v are in TE's ``qkv_format`` and are never permuted: cuDNN takes dims and strides, so - the descriptors are reordered into its logical BHSD instead. That is what serves bshd and - sbhd alike without a transpose. GQA is supported directly (h_kv may differ from h_q), SQ need - not equal SKV, which is what lets a CP ring step use this, and v may carry its own head_dim, - in which case out follows q's layout with v's head_dim. ``out`` comes back in ``qkv_format``; - softmax_lse is [b, h, s] fp32 natural-log logsumexp, the layout and convention the CP ring - correction expects, and is BHSD regardless of the input format. - """ - for name, tensor in (("q", q), ("k", k), ("v", v)): - _check_layout(name, tensor) - _check_dtype(name, tensor, q.dtype) - _check_kv_match(k, v) - qd, _ = _bhsd(q, qkv_format) - kd, _ = _bhsd(k, qkv_format) - vd, _ = _bhsd(v, qkv_format) - if kd[0] != qd[0] or kd[3] != qd[3]: - # The graph declares k and v with q's batch and head_dim, so a mismatch would bind a - # differently shaped buffer to that node and read the wrong elements silently. - raise ValueError(f"k must match q in batch and head_dim; got q {qd} and k {kd} in BHSD") - if qd[1] % kd[1] != 0: - raise ValueError(f"num_heads must be divisible by num_gqa_groups; got {qd[1]} and {kd[1]}") - - mask = _mask_spec(attn_mask_type, window_size) - scale = attn_scale if attn_scale is not None else qd[3] ** -0.5 - entry = _cached("fwd", _key(q, k, v, qkv_format, mask, scale)) - tq, tk, tv, tout, tlse = entry["handles"] - - b, hq, sq = qd[0], qd[1], qd[2] - # Allocated in the caller's format, so no permute is needed on the way out either. - out_shape, out_stride = _o_shape_stride(q.shape, vd[3], q.stride()) - # Allocated per call so concurrent uses cannot alias; the cache holds only the plan. - # empty_strided, not empty_like: the latter does not preserve an arbitrary permuted stride. - out = torch.empty_strided(out_shape, out_stride, device=q.device, dtype=q.dtype) - lse = torch.empty(b, hq, sq, 1, device=q.device, dtype=torch.float32) - workspace = torch.empty(entry["workspace"], device=q.device, dtype=torch.uint8) - entry["graph"].execute( - {tq: q, tk: k, tv: v, tout: out, tlse: lse}, - workspace, - handle=cudnn_pygraph.handle_for(q.device, backend_name=_BACKEND_NAME), - ) - return out, lse.squeeze(-1) - - -def frost_attn_bwd( - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - out: torch.Tensor, - softmax_lse: torch.Tensor, - dout: torch.Tensor, - qkv_format: str = "bshd", - attn_scale: Optional[float] = None, - attn_mask_type: str = "causal", - deterministic: bool = False, - window_size: Optional[Tuple[int, int]] = None, -) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Backward attention via cuDNN FROST. - - Tensors are in TE's ``qkv_format``, as in the forward. ``softmax_lse`` is [b, h, s] BHSD, as - the forward returned it. The gradients come back in ``qkv_format``. - """ - for name, tensor in (("q", q), ("k", k), ("v", v), ("out", out), ("dout", dout)): - _check_layout(name, tensor) - _check_dtype(name, tensor, q.dtype) - _check_kv_match(k, v) - # The same shape assumptions the forward makes, plus o/dO, which the graph declares with q's - # shape. The forward runs first in autograd, but the CP ring calls this directly. - qd, _ = _bhsd(q, qkv_format) - kd, _ = _bhsd(k, qkv_format) - vd, _ = _bhsd(v, qkv_format) - if kd[0] != qd[0] or kd[3] != qd[3]: - raise ValueError(f"k must match q in batch and head_dim; got q {qd} and k {kd} in BHSD") - if qd[1] % kd[1] != 0: - raise ValueError(f"num_heads must be divisible by num_gqa_groups; got {qd[1]} and {kd[1]}") - o_shape, o_stride = _o_shape_stride(q.shape, vd[3], q.stride()) - for name, tensor in (("out", out), ("dout", dout)): - if list(tensor.shape) != o_shape: - raise ValueError(f"{name} must be shaped {o_shape}; got {list(tensor.shape)}") - if softmax_lse.dtype != torch.float32: - raise ValueError(f"softmax_lse must be fp32; got {softmax_lse.dtype}") - # Compared against the BHSD description, not against q's own shape: the LSE is always - # [b, h, s] whatever format the tensors arrived in. - if tuple(softmax_lse.shape[:3]) != tuple(qd[:3]): - raise ValueError( - f"softmax_lse must be [b, h, s] matching q; got {tuple(softmax_lse.shape)} and" - f" {tuple(qd[:3])}" - ) - - mask = _mask_spec(attn_mask_type, window_size) - scale = attn_scale if attn_scale is not None else qd[3] ** -0.5 - entry = _cached("bwd", _key(q, k, v, qkv_format, mask, scale, deterministic)) - h = entry["handles"] - - if softmax_lse.dim() == 3: - softmax_lse = softmax_lse.unsqueeze(-1) - softmax_lse = softmax_lse.contiguous() - - # The graph expects o and dO in the layout the forward wrote, and dO comes from autograd - # with strides we do not control, so restride rather than silently reading the wrong elements. - def _as(t, stride): - if list(t.stride()) == list(stride): - return t - buf = torch.empty_strided(t.shape, stride, device=t.device, dtype=t.dtype) - buf.copy_(t) - return buf - - out = _as(out, o_stride) - dout = _as(dout, o_stride) - - dq = torch.empty_strided(q.shape, q.stride(), device=q.device, dtype=q.dtype) - dk = torch.empty_strided(k.shape, k.stride(), device=k.device, dtype=k.dtype) - dv = torch.empty_strided(v.shape, v.stride(), device=v.device, dtype=v.dtype) - workspace = torch.empty(entry["workspace"], device=q.device, dtype=torch.uint8) - entry["graph"].execute( - { - h["q"]: q, - h["k"]: k, - h["v"]: v, - h["o"]: out, - h["do"]: dout, - h["stats"]: softmax_lse, - h["dq"]: dq, - h["dk"]: dk, - h["dv"]: dv, - }, - workspace, - handle=cudnn_pygraph.handle_for(q.device, backend_name=_BACKEND_NAME), - ) - return dq, dk, dv - - def _frost_only(**unsupported): """Raise if any feature the selector should have declined reached the kernels anyway.""" for name, value in unsupported.items(): @@ -853,23 +733,34 @@ def fused_attn_fwd( raise NotImplementedError( f"FROST attention needs o_format to match qkv_format; got {o_format}/{qkv_format}" ) - mask_type, window = _te_mask_spec( + # _te_mask_spec validates as it normalises, so what it returns is the spec the plan keys on. + mask = _te_mask_spec( attn_mask_type, window_size, _bottom_right_diagonal(attn_mask_type, bottom_right_diagonal) ) - out, softmax_lse = frost_attn_fwd( - q.contiguous(), - k.contiguous(), - v.contiguous(), - qkv_format, - attn_scale=attn_scale, - attn_mask_type=mask_type, - window_size=window, + q, k, v = q.contiguous(), k.contiguous(), v.contiguous() + qd, _, vd = _validate_qkv(q, k, v, qkv_format) + scale = attn_scale if attn_scale is not None else qd[3] ** -0.5 + entry = _cached("fwd", _key(q, k, v, qkv_format, mask, scale)) + tq, tk, tv, tout, tlse = entry["handles"] + + # Allocated per call so concurrent uses cannot alias; the cache holds only the plan. + # empty_strided, not empty_like: the latter does not preserve an arbitrary permuted stride. + # The output takes the caller's format with v's head_dim, so nothing is converted on the way + # out; the LSE is BHSD whatever the inputs were, which is what the ring correction expects. + out_shape, out_stride = _o_shape_stride(q.shape, vd[3], q.stride()) + out = torch.empty_strided(out_shape, out_stride, device=q.device, dtype=q.dtype) + lse = torch.empty(qd[0], qd[1], qd[2], 1, device=q.device, dtype=torch.float32) + workspace = torch.empty(entry["workspace"], device=q.device, dtype=torch.uint8) + entry["graph"].execute( + {tq: q, tk: k, tv: v, tout: out, tlse: lse}, + workspace, + handle=cudnn_pygraph.handle_for(q.device, backend_name=_BACKEND_NAME), ) # A real tensor, not None: it is saved for backward and handed to the activation offload # hooks, neither of which accepts None. FROST has no dropout, so nothing reads it. rng_state = torch.empty(2, dtype=torch.int64, device=q.device) - return out, [softmax_lse, rng_state] + return out, [lse.squeeze(-1), rng_state] def fused_attn_bwd( @@ -928,22 +819,66 @@ def fused_attn_bwd( raise NotImplementedError( f"FROST attention needs {name} to match qkv_format; got {fmt}/{qkv_format}" ) - mask_type, window = _te_mask_spec( + mask = _te_mask_spec( attn_mask_type, window_size, _bottom_right_diagonal(attn_mask_type, bottom_right_diagonal) ) softmax_lse = aux_ctx_tensors[0] - dq, dk, dv = frost_attn_bwd( - q.contiguous(), - k.contiguous(), - v.contiguous(), - o.contiguous(), - softmax_lse, - d_o.contiguous(), - qkv_format, - attn_scale=attn_scale, - attn_mask_type=mask_type, - deterministic=deterministic, - window_size=window, + q, k, v = q.contiguous(), k.contiguous(), v.contiguous() + o, d_o = o.contiguous(), d_o.contiguous() + qd, _, vd = _validate_qkv(q, k, v, qkv_format) + for name, tensor in (("o", o), ("d_o", d_o)): + _check_layout(name, tensor) + _check_dtype(name, tensor, q.dtype) + o_shape, o_stride = _o_shape_stride(q.shape, vd[3], q.stride()) + for name, tensor in (("o", o), ("d_o", d_o)): + if list(tensor.shape) != o_shape: + raise ValueError(f"{name} must be shaped {o_shape}; got {list(tensor.shape)}") + if softmax_lse.dtype != torch.float32: + raise ValueError(f"softmax_lse must be fp32; got {softmax_lse.dtype}") + # Compared against the BHSD description, not q's own shape: the LSE is always [b, h, s] + # whatever format the tensors arrived in. + if tuple(softmax_lse.shape[:3]) != tuple(qd[:3]): + raise ValueError( + f"softmax_lse must be [b, h, s] matching q; got {tuple(softmax_lse.shape)} and" + f" {tuple(qd[:3])}" + ) + + scale = attn_scale if attn_scale is not None else qd[3] ** -0.5 + entry = _cached("bwd", _key(q, k, v, qkv_format, mask, scale, deterministic)) + h = entry["handles"] + + if softmax_lse.dim() == 3: + softmax_lse = softmax_lse.unsqueeze(-1) + softmax_lse = softmax_lse.contiguous() + + # The graph expects o and dO in the layout the forward wrote, and dO comes from autograd + # with strides we do not control, so restride rather than silently reading the wrong elements. + def _as(t, stride): + if list(t.stride()) == list(stride): + return t + buf = torch.empty_strided(t.shape, stride, device=t.device, dtype=t.dtype) + buf.copy_(t) + return buf + + o, d_o = _as(o, o_stride), _as(d_o, o_stride) + dq = torch.empty_strided(q.shape, q.stride(), device=q.device, dtype=q.dtype) + dk = torch.empty_strided(k.shape, k.stride(), device=k.device, dtype=k.dtype) + dv = torch.empty_strided(v.shape, v.stride(), device=v.device, dtype=v.dtype) + workspace = torch.empty(entry["workspace"], device=q.device, dtype=torch.uint8) + entry["graph"].execute( + { + h["q"]: q, + h["k"]: k, + h["v"]: v, + h["o"]: o, + h["do"]: d_o, + h["stats"]: softmax_lse, + h["dq"]: dq, + h["dk"]: dk, + h["dv"]: dv, + }, + workspace, + handle=cudnn_pygraph.handle_for(q.device, backend_name=_BACKEND_NAME), ) return dq, dk, dv, None From 314a91d106d23856895a6fe4cdb5d28be514fa98 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 22:26:43 -0700 Subject: [PATCH 70/97] refactor(attention): share the graph cache mechanism and the per-tensor key Both backends memoize a built graph the same way; only what they key on differs. cached_graph takes the mechanism, tensor_key takes the per-tensor fragment, and the key composition stays where it is -- flex keys on a score_mod callback identity that frost has no use for, and forcing both through one builder is how a shifted key becomes two configurations sharing a graph. This is a fix for flex, not only unification. flex built its graphs under whatever device was ambient, with only the handle naming the right one, while frost scoped the build because the plans are JIT-compiled and a compile path may read the CUDA context rather than the handle. flex now gets that scoping. frost's key goes from sixteen flat fields to seven structured ones, so the builders destructure by name rather than by position, and _device_from_key goes away: the device is passed to the cache rather than decoded back out of the key it was flattened into. Verified rather than assumed, since a key change is the kind that fails by silently sharing a plan: - Fourteen configurations differing in one axis each -- format, every dimension, dtype, device index, mask, scale, determinism -- give the same answer for all 91 pairs under the old key and the new one. - flex's key fragments are byte-identical across both formats and all three device spellings. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/cudnn_pygraph.py | 47 +++++++++ .../dot_product_attention/flex_attention.py | 40 +++----- .../dot_product_attention/frost_attention.py | 97 +++++++------------ 3 files changed, 99 insertions(+), 85 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index 4ce7d178991..5ed9ea569fa 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -18,6 +18,7 @@ from __future__ import annotations +import contextlib import importlib import os from typing import Any, Dict, Optional, Sequence, Tuple @@ -165,6 +166,52 @@ def bhsd_graph_tensor( return graph.tensor(dim=dim, stride=stride, data_type=tensor.dtype) +def device_key(device: torch.device) -> Tuple[Any, ...]: + """Normalize a device for a cache key. + + ``index is None`` is resolved to the current device, so ``cuda`` and ``cuda:0`` cannot key + two entries for one physical device. The type is part of the key too, so CPU cannot alias it. + """ + if device.type == "cuda" and device.index is None: + return ("cuda", torch.cuda.current_device()) + return (device.type, device.index) + + +def tensor_key( + tensor: torch.Tensor, tensor_format: str, *, backend_name: str = "cuDNN attention" +) -> Tuple[Any, ...]: + """A tensor as the graph will see it: BHSD dims, BHSD strides, dtype. + + Strides belong in the key because the graph is built for this exact layout -- that is what + lets bshd and sbhd run without a transpose -- and the dtype because every node is declared + with one. Two formats that produce the same description are the same graph and should share + a plan, which is why the format itself is not keyed. + """ + dim, stride = bhsd_dim_stride(tensor, tensor_format, backend_name=backend_name) + return (tuple(dim), tuple(stride), tensor.dtype) + + +def cached_graph(cache: Dict[Any, Any], key: Optional[Any], build, *, device=None): + """Memoize a built graph. ``key=None`` means uncacheable: build and return without storing. + + The build runs under ``device`` when one is given, not merely with its handle: the plans are + JIT-compiled, and a compile path may read the ambient CUDA context rather than the handle. + """ + if key is None: + return build() + entry = cache.get(key) + if entry is None: + scope = ( + torch.cuda.device(device) + if device is not None and device.type == "cuda" + else contextlib.nullcontext() + ) + with scope: + entry = build() + cache[key] = entry + return entry + + def finalize_plans( graph, *, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index 38caf5b740b..f5a7c739c85 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -128,12 +128,7 @@ def _score_mod_callback_cache_key(callback: Optional[Callable]) -> Any: def _score_mod_device_key(device: torch.device) -> Tuple[Any, ...]: """Normalize a tensor device for graph cache keys.""" - if device.type == "cuda": - index = device.index - if index is None: - index = torch.cuda.current_device() - return (device.type, index) - return (device.type, device.index) + return cudnn_pygraph.device_key(device) def _score_mod_tensor_metadata(tensor: torch.Tensor) -> Tuple[Any, ...]: @@ -157,8 +152,9 @@ def _score_mod_tensor_dict_metadata( def _score_mod_bhsd_tensor_metadata(tensor: torch.Tensor, tensor_format: str) -> Tuple[Any, ...]: """Describe an SBHD/BSHD runtime tensor as a cuDNN BHSD graph tensor.""" - dim, stride = _bhsd_dim_stride(tensor, tensor_format) - return (dim, stride, tensor.dtype, _score_mod_device_key(tensor.device)) + return cudnn_pygraph.tensor_key(tensor, tensor_format, backend_name=_BACKEND_NAME) + ( + cudnn_pygraph.device_key(tensor.device), + ) def _make_cudnn_graph_tensor_dict(graph, tensors: Optional[Dict[str, torch.Tensor]]): @@ -415,14 +411,12 @@ def _get_cudnn_score_mod_fwd_graph( output_layer, stats, ) - key = _cudnn_score_mod_fwd_cache_key(*build_args) - if key is None: - return _build_cudnn_score_mod_fwd_graph(*build_args) - entry = _cudnn_score_mod_graph_cache.get(key) - if entry is None: - entry = _build_cudnn_score_mod_fwd_graph(*build_args) - _cudnn_score_mod_graph_cache[key] = entry - return entry + return cudnn_pygraph.cached_graph( + _cudnn_score_mod_graph_cache, + _cudnn_score_mod_fwd_cache_key(*build_args), + lambda: _build_cudnn_score_mod_fwd_graph(*build_args), + device=query_layer.device, + ) def _build_cudnn_score_mod_bwd_graph( @@ -534,14 +528,12 @@ def _get_cudnn_score_mod_bwd_graph( score_mod_bprop_tensors, deterministic, ) - key = _cudnn_score_mod_bwd_cache_key(*build_args) - if key is None: - return _build_cudnn_score_mod_bwd_graph(*build_args) - entry = _cudnn_score_mod_graph_cache.get(key) - if entry is None: - entry = _build_cudnn_score_mod_bwd_graph(*build_args) - _cudnn_score_mod_graph_cache[key] = entry - return entry + return cudnn_pygraph.cached_graph( + _cudnn_score_mod_graph_cache, + _cudnn_score_mod_bwd_cache_key(*build_args), + lambda: _build_cudnn_score_mod_bwd_graph(*build_args), + device=query_layer.device, + ) class FusedAttentionWithScoreModFunc(torch.autograd.Function): diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 6ceaa33c548..90da637bb6a 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -15,7 +15,6 @@ from __future__ import annotations -import contextlib import os from importlib.metadata import PackageNotFoundError, version as get_pkg_version from typing import Any, Dict, Optional, Sequence, Tuple @@ -93,12 +92,6 @@ def _diagonal_band_kwargs(cudnn, attn_mask_type: str, window: Tuple[int, int]) - return opts -def _device_from_key(device_key) -> torch.device: - """Rebuild the torch.device that _key recorded, for building under the right device.""" - kind, index = device_key - return torch.device(kind) if index is None else torch.device(kind, index) - - def _pkg_version(name: str, module=None) -> Tuple[Optional[PkgVersion], Optional[str]]: """(parsed version, raw string) for a package. Either element is None if undeterminable. @@ -503,21 +496,19 @@ def hint(): return name -def _build_fwd(key) -> dict: +def _build_fwd(key, device) -> dict: """Build (and JIT-compile) a forward graph. Expensive; always reached through the cache.""" cudnn = _import_cudnn_frontend() # deterministic is unused here: it selects a backward algorithm. Callers pass False for the # forward so the two never split the forward cache. - *_device, b, hq, hkv, sq, skv, d, d_v, dtype, mask, scale, qs, ks, vs, _deterministic = key - shq, shk, shv = [b, hq, sq, d], [b, hkv, skv, d], [b, hkv, skv, d_v] - sho, o_stride = _o_shape_stride([b, hq, sq, d], d_v, qs) - - graph = cudnn_pygraph.build_pygraph( - dtype, _device_from_key(_device), backend_name=_BACKEND_NAME - ) - tq = graph.tensor(name="q", dim=shq, stride=list(qs)) - tk = graph.tensor(name="k", dim=shk, stride=list(ks)) - tv = graph.tensor(name="v", dim=shv, stride=list(vs)) + _dev, (shq, qs, dtype), (shk, ks, _), (shv, vs, _), mask, scale, _deterministic = key + b, hq, sq = shq[0], shq[1], shq[2] + sho, o_stride = _o_shape_stride(shq, shv[3], qs) + + graph = cudnn_pygraph.build_pygraph(dtype, device, backend_name=_BACKEND_NAME) + tq = graph.tensor(name="q", dim=list(shq), stride=list(qs)) + tk = graph.tensor(name="k", dim=list(shk), stride=list(ks)) + tv = graph.tensor(name="v", dim=list(shv), stride=list(vs)) tout, tlse = graph.sdpa( name="frost_fwd", q=tq, @@ -540,17 +531,15 @@ def _build_fwd(key) -> dict: } -def _build_bwd(key) -> dict: +def _build_bwd(key, device) -> dict: """Build (and JIT-compile) a backward graph. Expensive; always reached through the cache.""" cudnn = _import_cudnn_frontend() - *_device, b, hq, hkv, sq, skv, d, d_v, dtype, mask, scale, qs, ks, vs, deterministic = key + _dev, (shq, qs, dtype), (shk, ks, _), (shv, vs, _), mask, scale, deterministic = key io_dt = cudnn_pygraph.io_data_type(cudnn, dtype, backend_name=_BACKEND_NAME) - shq, shk, shv = [b, hq, sq, d], [b, hkv, skv, d], [b, hkv, skv, d_v] - sho, o_stride = _o_shape_stride([b, hq, sq, d], d_v, qs) + b, hq, sq = shq[0], shq[1], shq[2] + sho, o_stride = _o_shape_stride(shq, shv[3], qs) - graph = cudnn_pygraph.build_pygraph( - dtype, _device_from_key(_device), backend_name=_BACKEND_NAME - ) + graph = cudnn_pygraph.build_pygraph(dtype, device, backend_name=_BACKEND_NAME) handles = {} # Each grad is declared with the layout of the tensor it differentiates. for name, shape, stride in ( @@ -560,7 +549,7 @@ def _build_bwd(key) -> dict: ("o", sho, o_stride), ("do", sho, o_stride), ): - handles[name] = graph.tensor(name=name, dim=shape, stride=list(stride)) + handles[name] = graph.tensor(name=name, dim=list(shape), stride=list(stride)) handles["stats"] = graph.tensor( name="stats", dim=[b, hq, sq, 1], @@ -591,19 +580,13 @@ def _build_bwd(key) -> dict: } -def _cached(kind: str, key): +def _cached(kind: str, key, device): """Plan cache. See module docstring: building dominates executing even once the JIT is cached, so this is required rather than an optimisation.""" - cache_key = (kind,) + key - entry = _PLAN_CACHE.get(cache_key) - if entry is None: - # Build under the device the key names, not merely with its handle: the plans are - # JIT-compiled, and a compile path may read the ambient CUDA context rather than the handle. - device = _device_from_key(key[:2]) - with torch.cuda.device(device) if device.type == "cuda" else contextlib.nullcontext(): - entry = _build_fwd(key) if kind == "fwd" else _build_bwd(key) - _PLAN_CACHE[cache_key] = entry - return entry + build = _build_fwd if kind == "fwd" else _build_bwd + return cudnn_pygraph.cached_graph( + _PLAN_CACHE, (kind,) + key, lambda: build(key, device), device=device + ) def _validate_qkv(q, k, v, qkv_format): @@ -636,32 +619,24 @@ def _bhsd(t: torch.Tensor, qkv_format: str): def _key(q, k, v, qkv_format, mask, scale, deterministic=False): - qd, qs = _bhsd(q, qkv_format) - kd, ks = _bhsd(k, qkv_format) - vd, vs = _bhsd(v, qkv_format) + """The plan cache key. + + Structured per tensor rather than flattened, so the builders destructure it by name instead + of by position, and so the per-tensor fragment is the same one flex_attention keys on. + """ + + def described(t): + return cudnn_pygraph.tensor_key(t, qkv_format, backend_name=_BACKEND_NAME) + return ( # Built under whichever device was current, so it must not be reused on another. Matches - # the C++ fused-attn cache, which keys on device_id. Type too, so CPU cannot alias cuda:0. - q.device.type, - q.device.index, - qd[0], - qd[1], - kd[1], - qd[2], - kd[2], - qd[3], - # v carries its own head_dim and strides, the way flex_attention keys each tensor - # separately. Without them an asymmetric v would reuse a plan built for k's shape. - vd[3], - q.dtype, + # the C++ fused-attn cache, which keys on device_id. + cudnn_pygraph.device_key(q.device), + described(q), + described(k), + described(v), mask, float(scale), - # Strides are part of the plan: the graph is built for this exact layout, which is what - # lets bshd and sbhd both run without a transpose. qkv_format does not need its own key - # entry, since two formats producing the same BHSD description are the same graph. - tuple(qs), - tuple(ks), - tuple(vs), # The deterministic backward is a different algorithm, not a flag on the same one, so a # plan built either way must not be handed to a call that asked for the other. bool(deterministic), @@ -741,7 +716,7 @@ def fused_attn_fwd( q, k, v = q.contiguous(), k.contiguous(), v.contiguous() qd, _, vd = _validate_qkv(q, k, v, qkv_format) scale = attn_scale if attn_scale is not None else qd[3] ** -0.5 - entry = _cached("fwd", _key(q, k, v, qkv_format, mask, scale)) + entry = _cached("fwd", _key(q, k, v, qkv_format, mask, scale), q.device) tq, tk, tv, tout, tlse = entry["handles"] # Allocated per call so concurrent uses cannot alias; the cache holds only the plan. @@ -845,7 +820,7 @@ def fused_attn_bwd( ) scale = attn_scale if attn_scale is not None else qd[3] ** -0.5 - entry = _cached("bwd", _key(q, k, v, qkv_format, mask, scale, deterministic)) + entry = _cached("bwd", _key(q, k, v, qkv_format, mask, scale, deterministic), q.device) h = entry["handles"] if softmax_lse.dim() == 3: From 2931c9581243deffa0ad1fd0296a27df5082f57f Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 22:31:52 -0700 Subject: [PATCH 71/97] fix(attention): scope the uncacheable graph build to its device too 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 #3527. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- .../attention/dot_product_attention/cudnn_pygraph.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index 5ed9ea569fa..71fdbf53666 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -196,10 +196,10 @@ def cached_graph(cache: Dict[Any, Any], key: Optional[Any], build, *, device=Non The build runs under ``device`` when one is given, not merely with its handle: the plans are JIT-compiled, and a compile path may read the ambient CUDA context rather than the handle. + That applies to the uncacheable path too, which is why there is one build site rather than + one per branch -- a second would be free to forget the scope. """ - if key is None: - return build() - entry = cache.get(key) + entry = None if key is None else cache.get(key) if entry is None: scope = ( torch.cuda.device(device) @@ -208,7 +208,8 @@ def cached_graph(cache: Dict[Any, Any], key: Optional[Any], build, *, device=Non ) with scope: entry = build() - cache[key] = entry + if key is not None: + cache[key] = entry return entry From e0f2690985d96915f9c89b5db3f47104b0435541 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 22:39:21 -0700 Subject: [PATCH 72/97] fix(attention): bar the cuDNN FROST engines by default, not on request 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 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_flex_attention.py | 121 ++++++++++++++++++ .../dot_product_attention/cudnn_pygraph.py | 39 +++++- .../dot_product_attention/frost_attention.py | 6 +- 3 files changed, 163 insertions(+), 3 deletions(-) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index beed4069917..ad4f37a25c3 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -12,6 +12,7 @@ from transformer_engine.pytorch import DotProductAttention, is_bf16_available from transformer_engine.pytorch.attention.dot_product_attention import _attention_backends import transformer_engine.pytorch.attention.dot_product_attention.flex_attention as flex_attention +from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph from transformer_engine.pytorch.utils import get_device_compute_capability _current_file = pathlib.Path(__file__).resolve() @@ -705,3 +706,123 @@ def test_dot_product_attention_score_mod(dtype, qkv_format, score_mod_case, scal torch.testing.assert_close(q.grad, q_ref.grad, **tols) torch.testing.assert_close(k.grad, k_ref.grad, **tols) torch.testing.assert_close(v.grad, v_ref.grad, **tols) + + +def test_flex_bars_the_frost_engines(): + """flex must tell cuDNN not to use a FROST engine, not merely decline to ask for them. + + The switch that offers those engines is process-wide, so another caller in the process, or a + user setting CUDNN_FRONTEND_ENABLE_FROST_ENGINES, puts them ahead of the backend engines for + these graphs too. They accept a score_mod graph, pass check_support, build, and then compute + without the callback. + + The bar is the shared helper's default rather than something flex asks for, so this checks + flex's own path reaches it: a caller that forgot to opt in is the failure being prevented. + + No GPU: this checks the instruction is passed, not what cuDNN does with it. + """ + barred = [] + + class FakeGraph: + """Records the engines flex bars, and stops at the first call it cannot serve.""" + + def validate(self): + pass + + def build_operation_graph(self): + pass + + def create_execution_plans(self, _heuristics): + pass + + def deselect_engines(self, names): + barred.extend(names) + + def check_support(self): + pass + + def build_plans(self, _policy): + pass + + def get_workspace_size(self): + return 4096 + + try: + flex_attention._import_cudnn_frontend() + except ImportError: + pytest.skip("cuDNN frontend Python package is required for score_mod attention.") + + assert flex_attention._finalize_cudnn_graph(FakeGraph()) == 4096 + assert barred, "flex did not ask cuDNN to exclude any engine" + assert set(cudnn_pygraph.FROST_PLAN_TOKENS) <= set(barred), barred + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required.") +def test_frost_switch_does_not_change_what_flex_computes(): + """Enabling the FROST engines must not change flex's output. + + This is the property the silent drop violates: with the engines on, an unpinned build selects + a FROST plan at every head dim measured on B200, and that plan returns plain attention with + the score_mod discarded. Comparing flex against itself across the switch needs no reference + and no knowledge of which plan ran; if the two differ, a different kernel answered. + """ + try: + flex_attention._import_cudnn_frontend() + except ImportError: + pytest.skip("cuDNN frontend Python package is required for score_mod attention.") + # Without this the test is vacuous: if the engines are absent, decline on arch, or sit below + # their version floors, both runs get a backend plan and agree no matter what flex does. The + # availability probe covers all three, where an arch check alone would miss the version ones. + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + is_frost_attention_available, + ) + + frost_ok, frost_reason = is_frost_attention_available() + if not frost_ok: + pytest.skip("the FROST engines must be reachable to test anything: %s" % frost_reason) + + env = "CUDNN_FRONTEND_ENABLE_FROST_ENGINES" + saved = os.environ.get(env) + torch.manual_seed(0) + b, h, s, d = 2, 4, 512, 64 + dtype = torch.bfloat16 if is_bf16_available() else torch.float16 + q, k, v = (torch.randn(b, s, h, d, device="cuda", dtype=dtype) for _ in range(3)) + + def bias_score_mod(score_mod_graph, score_tensor, _tensors): + """score += (row - col). Self-contained, and large enough that dropping it is obvious.""" + cudnn = flex_attention._import_cudnn_frontend() + row = score_mod_graph.gen_index(input=score_tensor, axis=2) + row.set_data_type(cudnn.data_type.INT32) + col = score_mod_graph.gen_index(input=score_tensor, axis=3) + col.set_data_type(cudnn.data_type.INT32) + bias = score_mod_graph.sub(a=row, b=col, compute_data_type=cudnn.data_type.FLOAT) + bias.set_data_type(cudnn.data_type.FLOAT) + return score_mod_graph.add(a=score_tensor, b=bias, compute_data_type=cudnn.data_type.FLOAT) + + def run(): + flex_attention._cudnn_score_mod_graph_cache.clear() + return flex_attention.FusedAttentionWithScoreModFunc.apply( + False, q, k, v, "bshd", "bshd", d**-0.5, bias_score_mod, None, None, None, False + ) + + try: + os.environ.pop(env, None) + without = run() + os.environ[env] = "1" + with_engines = run() + finally: + flex_attention._cudnn_score_mod_graph_cache.clear() + if saved is None: + os.environ.pop(env, None) + else: + os.environ[env] = saved + + torch.testing.assert_close( + with_engines, + without, + msg=lambda m: ( + "flex computed something different with the FROST engines enabled, which means a" + " FROST plan answered and dropped the score_mod:\n" + + m + ), + ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index 71fdbf53666..9908515ab2e 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -25,6 +25,17 @@ import torch +# The cuDNN FROST SDPA engines, named once so the two opposite instructions about them cannot +# drift: frost pins one of these by name, flex bars both. cuDNN has already renamed this family +# once (collapsing the per-head-dim ``..._d512`` rows), and if these strings stopped matching, +# frost would fail loudly while flex failed silently. +FROST_FWD_PLAN_TOKEN = "sdpa_fwd_prefill_sm100" +FROST_BWD_PLAN_TOKEN = "sdpa_bwd_sm100" +FROST_PLAN_TOKENS = (FROST_FWD_PLAN_TOKEN, FROST_BWD_PLAN_TOKEN) + +# Distinguishes "the caller said nothing" from "the caller asked for no exclusions at all". +_BAR_FROST_BY_DEFAULT = object() + _cudnn = None _frost_engines_enabled = False _HANDLES: Dict[Tuple[str, torch.device], Any] = {} @@ -221,6 +232,7 @@ def finalize_plans( build_policy: Any = None, require_plan_token: Optional[str] = None, not_found_hint: Any = "", + exclude_plan_tokens: Any = _BAR_FROST_BY_DEFAULT, ) -> Tuple[int, Optional[str]]: """Create plans, optionally pin one by name, build, and return (workspace size, plan name). @@ -241,6 +253,15 @@ def finalize_plans( ``_plan_pinned``, and only then is a decline fatal; unpinned, cuDNN records the decline and keeps walking. So the pin has to come first both because the check is scoped to the selected plan and because it is what makes the check binding at all. + + ``exclude_plan_tokens`` is the opposite instruction, and it defaults to barring the FROST + engines. That default is deliberate: the switch offering them is process-wide, so a caller + that merely declines to ask for them still gets them ranked first once anything else in the + process has enabled them, and those engines accept a score_mod graph and then compute without + it. Forgetting to exclude gives silently wrong numbers; excluding wrongly gives a slower plan + or a loud decline, so the burden belongs on the caller that wants them rather than the one + that does not. A caller pinning by name has already said which engine it wants, so exclusion + is skipped there rather than contradicting the pin. Pass an explicit value to override. """ cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() @@ -251,11 +272,27 @@ def finalize_plans( heuristics = [cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK] if require_plan_token is None: + if exclude_plan_tokens is _BAR_FROST_BY_DEFAULT: + exclude_plan_tokens = FROST_PLAN_TOKENS try: graph.create_execution_plans(list(heuristics)) + # Bar the named engines before the walk, so build_plans falls through to the first + # entry that is both unbarred and buildable. Inert where they are not on offer, which + # is every process that has not enabled them. getattr so a frontend predating + # deselect_engines degrades rather than raising. + deselect = getattr(graph, "deselect_engines", None) + if exclude_plan_tokens and deselect is not None: + deselect(list(exclude_plan_tokens)) graph.check_support() except cudnn.cudnnGraphNotSupportedError as exc: - raise RuntimeError(f"cuDNN {backend_name} SDPA graph is not supported: {exc}") from exc + # Name the bar in the message: if it removed the only viable plan, the graph is not + # what was unsupported. + barred = ( + f" (barred engines: {list(exclude_plan_tokens)})" if exclude_plan_tokens else "" + ) + raise RuntimeError( + f"cuDNN {backend_name} SDPA graph is not supported: {exc}{barred}" + ) from exc if build_policy is None: build_policy = cudnn.build_plan_policy.HEURISTICS_CHOICE graph.build_plans(build_policy) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 90da637bb6a..1f2ef1aabd8 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -35,8 +35,10 @@ # cudnn-frontend declares cutlass-dsl >= 4.6.2 but FROST enforces >= 4.7.0 at plan-build time. # Below that floor every FROST engine declines silently and backend plans come back instead, so # the selected plan is checked by NAME rather than trusting that the engine was used. -_FROST_FWD_PLAN_TOKEN = "sdpa_fwd_prefill_sm100" -_FROST_BWD_PLAN_TOKEN = "sdpa_bwd_sm100" +# Defined in cudnn_pygraph so the pin here and the bar in flex_attention cannot name different +# engines: one of them failing loudly and the other silently is exactly the drift to avoid. +_FROST_FWD_PLAN_TOKEN = cudnn_pygraph.FROST_FWD_PLAN_TOKEN +_FROST_BWD_PLAN_TOKEN = cudnn_pygraph.FROST_BWD_PLAN_TOKEN _MIN_CUTLASS_DSL = PkgVersion("4.7.0") # 1.29.0 is the first release carrying the head_dim=512 backward. 1.28.0 ships the forward only, From 631d8be4e908bbfe89b805322da537fc3ea6178e Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 22:49:22 -0700 Subject: [PATCH 73/97] docs(attention): bring the new docstrings into the repo's register 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/cudnn_pygraph.py | 110 +++++++----------- .../dot_product_attention/frost_attention.py | 72 ++++-------- 2 files changed, 61 insertions(+), 121 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index 9908515ab2e..b828268a75e 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -44,22 +44,14 @@ def import_cudnn_frontend(enable_frost_engines: bool = False): """Import cuDNN Frontend, enabling the FROST engines if this caller needs them. - ``enable_frost_engines`` is not merely additive: the switch also ranks FROST ahead of the - backend engines everywhere, so a caller that does not want FROST must not ask for it. Hence - the default is off, and FROST asks explicitly. - - The enabling is deliberately outside the import memo. Both backends call this, and whichever - one reaches it first would otherwise decide for the process: with the flag inside the memo, a - flex call would cache the module with FROST off and every later FROST call would get a cuDNN - that offers no FROST engine, which surfaces much later as "no cuDNN engine matching ... was - offered". Enabling late is sound because the switch is read per graph rather than at import: - in cuDNN Frontend 1.29.0 ``engines/manifest.py`` consults the environment inside - ``offered_ids()``, reached from ``engines_for(graph)`` on every ``create_execution_plans``. - - Note the switch is process-wide and never unset, so enabling it for FROST also reorders the - candidates a concurrent score_mod graph sees. Callers that require a particular engine should - verify by plan name rather than rely on the switch, which is what - ``finalize_plans(require_plan_token=...)`` does. + The switch ranks FROST ahead of the backend engines process-wide, so it defaults off and + FROST asks explicitly. A caller needing a particular engine should verify by plan name rather + than trust the switch, which is what ``finalize_plans(require_plan_token=...)`` does. + + The enabling sits outside the import memo deliberately. Inside it, whichever backend imported + cuDNN first would decide for the process, and a flex-first import would leave every later + FROST call with no engine on offer. Enabling late works because cuDNN re-reads the environment + per graph, inside ``offered_ids()`` on every ``create_execution_plans``. """ global _cudnn, _frost_engines_enabled # pylint: disable=global-statement if _cudnn is None: @@ -82,8 +74,7 @@ def import_cudnn_frontend(enable_frost_engines: bool = False): def cudnn_module(): """The imported frontend, or None if nothing has imported it yet. - For callers that want to inspect the module without triggering an import, such as a version - probe that must not enable anything as a side effect. + For a caller that must inspect it without triggering an import, such as a version probe. """ return _cudnn @@ -96,16 +87,13 @@ def frost_engines_enabled() -> bool: def handle_for(device: torch.device, *, backend_name: str = "cuDNN attention"): """A cuDNN handle for ``device``, rebound to PyTorch's current stream on every call. - Without the rebinding, cuDNN runs on its handle's own stream while the tensors and workspace - are allocated on PyTorch's current stream, and nothing orders the two. That is not - hypothetical: the p2p context-parallel ring issues attention inside - ``with torch.cuda.stream(cp_stream)``, so on alternating ring steps the kernel and its buffers - would otherwise be on different streams. The same cached plan is executed from different - streams across steps, so this has to happen per call rather than once per handle. + Without the rebinding cuDNN runs on its handle's own stream while the tensors and workspace + sit on PyTorch's, with nothing ordering the two. The p2p context-parallel ring issues + attention inside ``with torch.cuda.stream(cp_stream)``, and executes one cached plan from + different streams across ring steps, so this is per call rather than per handle. - Keyed on ``(backend_name, device)`` rather than on the device alone. A cuDNN handle is not - thread-safe and ``set_stream`` mutates it, so one handle shared by two backends widens an - existing within-backend race into a cross-backend one for no benefit. + Keyed on ``(backend_name, device)``: a cuDNN handle is not thread-safe and ``set_stream`` + mutates it, so one handle shared by two backends widens an existing race for no benefit. """ if device.type != "cuda": raise ValueError(f"{backend_name} only supports CUDA tensors, got device {device}.") @@ -125,8 +113,7 @@ def handle_for(device: torch.device, *, backend_name: str = "cuDNN attention"): def io_data_type(cudnn, dtype: torch.dtype, *, backend_name: str = "cuDNN attention"): """Map a torch dtype to the cuDNN enum these SDPA graphs are declared with. - Takes ``cudnn`` rather than importing it, so a dtype lookup cannot import the frontend or - flip the FROST switch as a side effect. + Takes ``cudnn`` rather than importing it, so a dtype lookup cannot flip the FROST switch. """ if dtype == torch.float16: return cudnn.data_type.HALF @@ -153,8 +140,8 @@ def bhsd_dim_stride( ) -> Tuple[Tuple[int, ...], Tuple[int, ...]]: """Describe an SBHD/BSHD tensor as cuDNN frontend's logical BHSD format. - The tensor is never permuted. cuDNN takes dims and strides, so reordering the descriptors - says the same thing as permuting the tensor and costs nothing. + The tensor is never permuted: cuDNN takes dims and strides, so reordering the descriptors + says the same thing and costs nothing. """ if tensor_format == "sbhd": return ( @@ -178,11 +165,7 @@ def bhsd_graph_tensor( def device_key(device: torch.device) -> Tuple[Any, ...]: - """Normalize a device for a cache key. - - ``index is None`` is resolved to the current device, so ``cuda`` and ``cuda:0`` cannot key - two entries for one physical device. The type is part of the key too, so CPU cannot alias it. - """ + """Normalize a device for a cache key, so ``cuda`` and ``cuda:0`` cannot key two entries.""" if device.type == "cuda" and device.index is None: return ("cuda", torch.cuda.current_device()) return (device.type, device.index) @@ -193,22 +176,20 @@ def tensor_key( ) -> Tuple[Any, ...]: """A tensor as the graph will see it: BHSD dims, BHSD strides, dtype. - Strides belong in the key because the graph is built for this exact layout -- that is what - lets bshd and sbhd run without a transpose -- and the dtype because every node is declared - with one. Two formats that produce the same description are the same graph and should share - a plan, which is why the format itself is not keyed. + Strides are keyed because the graph is built for this exact layout, which is what lets bshd + and sbhd run without a transpose. The format itself is not: two formats giving the same + description are the same graph. """ dim, stride = bhsd_dim_stride(tensor, tensor_format, backend_name=backend_name) return (tuple(dim), tuple(stride), tensor.dtype) def cached_graph(cache: Dict[Any, Any], key: Optional[Any], build, *, device=None): - """Memoize a built graph. ``key=None`` means uncacheable: build and return without storing. + """Memoize a built graph. ``key=None`` means uncacheable: build, return, do not store. - The build runs under ``device`` when one is given, not merely with its handle: the plans are - JIT-compiled, and a compile path may read the ambient CUDA context rather than the handle. - That applies to the uncacheable path too, which is why there is one build site rather than - one per branch -- a second would be free to forget the scope. + The build runs under ``device`` rather than merely with its handle, because the plans are + JIT-compiled and a compile path may read the ambient CUDA context. One build site, not one + per branch: a second is free to forget that. """ entry = None if key is None else cache.get(key) if entry is None: @@ -236,32 +217,19 @@ def finalize_plans( ) -> Tuple[int, Optional[str]]: """Create plans, optionally pin one by name, build, and return (workspace size, plan name). - ``require_plan_token`` makes the choice strict: only a plan whose name contains the token is - acceptable, and anything else raises. That is not a stylistic preference. Without a pin, - ``build_plans`` walks the ranked list from index 0 and finalizes the first plan that builds, - logging each decline at INFO, so a graph that the intended engine declines runs on whatever - cuDNN ranked next with nothing in the return value to say so. At head_dim 512 that matters in - the forward, where an ordinary engine may well build and compute a different function from the - FROST kernel. The backward is self-limiting, since no non-FROST d512 backward exists, so an - unpinned backward would fail loudly on its own. - - The token is matched as a substring rather than by equality on purpose: cuDNN has already - collapsed per-head-dim engine names (``..._d512`` and friends) into a single row once, and the - substring test survived that. - - Pinning also changes what ``check_support`` means. Selecting a plan sets cuDNN's internal - ``_plan_pinned``, and only then is a decline fatal; unpinned, cuDNN records the decline and - keeps walking. So the pin has to come first both because the check is scoped to the selected - plan and because it is what makes the check binding at all. - - ``exclude_plan_tokens`` is the opposite instruction, and it defaults to barring the FROST - engines. That default is deliberate: the switch offering them is process-wide, so a caller - that merely declines to ask for them still gets them ranked first once anything else in the - process has enabled them, and those engines accept a score_mod graph and then compute without - it. Forgetting to exclude gives silently wrong numbers; excluding wrongly gives a slower plan - or a loud decline, so the burden belongs on the caller that wants them rather than the one - that does not. A caller pinning by name has already said which engine it wants, so exclusion - is skipped there rather than contradicting the pin. Pass an explicit value to override. + ``require_plan_token`` makes the choice strict: a plan whose name lacks the token raises. + Unpinned, ``build_plans`` walks the ranked list and finalizes the first that builds, so a + graph the intended engine declines runs on whatever cuDNN ranked next with nothing in the + return value saying so. The token is matched as a substring because cuDNN has already + collapsed per-head-dim engine names into one row once. The pin has to precede + ``check_support``: selecting a plan sets cuDNN's ``_plan_pinned``, and only then is a decline + fatal rather than recorded. + + ``exclude_plan_tokens`` is the opposite instruction, defaulting to barring the FROST engines, + which accept a score_mod graph and then compute without it. The switch offering them is + process-wide, so declining to ask is not enough. Forgetting to exclude is silently wrong while + excluding wrongly costs a slower plan or a loud decline, so the default favours the caller + that does not want them; it is skipped when a plan is pinned, which has already named one. """ cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 1f2ef1aabd8..54c59b5a653 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -60,9 +60,8 @@ def _import_cudnn_frontend(enable_frost_engines: bool = True): """Import cuDNN Frontend with the FROST engines on, which is what this backend needs. - The default differs from the shared module's, where it is off. Every use site here wants the - engines; a caller that does not must not ask for them, because the switch is process-wide. - See ``cudnn_pygraph.import_cudnn_frontend`` for why the enabling sits outside the import memo. + The shared default is off; every use site here wants them. See + ``cudnn_pygraph.import_cudnn_frontend``. """ return cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=enable_frost_engines) @@ -95,13 +94,7 @@ def _diagonal_band_kwargs(cudnn, attn_mask_type: str, window: Tuple[int, int]) - def _pkg_version(name: str, module=None) -> Tuple[Optional[PkgVersion], Optional[str]]: - """(parsed version, raw string) for a package. Either element is None if undeterminable. - - Distribution metadata first, matching the sibling check in fused_mla_q_uproj.py, with the - module attribute as a fallback so a source or vendored install is not misreported as absent. - The raw string is returned separately so callers can tell "not installed" from "installed but - unparseable"; those warrant different answers, and conflating them declines valid installs. - """ + """A package's version, or None when it is absent or unparseable, with the raw string.""" raw = None for candidate in (lambda: get_pkg_version(name), lambda: getattr(module, "__version__", None)): try: @@ -261,11 +254,9 @@ def _te_mask_spec(attn_mask_type: str, window_size, bottom_right_diagonal: bool) def _bottom_right_diagonal(attn_mask_type: str, bottom_right_diagonal) -> bool: - """Resolve the anchor flag the same way cpp_extensions.fused_attn does. + """Resolve the anchor flag the way cpp_extensions.fused_attn does. - ``None`` means "read it off the mask name". The dispatcher resolves it before calling in, - so this only matters for a direct call, where ``bool(None)`` would quietly give a top-left - band to a caller that asked for bottom-right. + ``None`` means read it off the mask name. Only a direct call sees it unresolved. """ if bottom_right_diagonal is None: return attn_mask_type in {"causal_bottom_right", "padding_causal_bottom_right"} @@ -283,14 +274,12 @@ def _name_for(table, value, default=None): def is_frost_attention_supported(params) -> Tuple[int, str]: """Whether this fused-attention config should run on the FROST sub-backend. - Takes a FusedAttentionParams and returns (sub-backend value, reject message), the same shape - as tex.get_fused_attn_backend, so get_attention_backend can fall through to it when the C++ - backends decline. + Returns (sub-backend value, reject message) like tex.get_fused_attn_backend, so + get_attention_backend can fall through to it when the C++ backends decline. - Deliberately does not probe availability. That imports cuDNN Frontend with the FROST engines + Deliberately does not probe availability: that imports cuDNN Frontend with the FROST engines enabled, which changes the engine pool for every cuDNN consumer in the process, and this runs - for every attention config on the machine. get_attention_backend checks availability once at - the end, the way it checks flash-attn versions. + for every attention config. get_attention_backend checks availability once at the end. """ # pylint: disable-next=import-outside-toplevel from ...cpp_extensions.fused_attn import ( @@ -398,12 +387,7 @@ def is_frost_attention_supported(params) -> Tuple[int, str]: def _check_layout(name: str, t: torch.Tensor) -> None: - """Validate a [b, h, s, d] view. - - The graphs are built from each tensor's ACTUAL strides rather than one fixed layout, so bshd - and sbhd are both served without a transpose. The only hard requirement is that the head - dimension is contiguous, which the kernels assume. - """ + """Validate a 4D view whose head dimension is contiguous, which the kernels assume.""" if t.dim() != 4: raise ValueError(f"{name} must be 4D [b, h, s, d]; got {tuple(t.shape)}") if t.stride(3) != 1: @@ -416,9 +400,7 @@ def _check_layout(name: str, t: torch.Tensor) -> None: def _check_dtype(name: str, t: torch.Tensor, expected: torch.dtype) -> None: """Require a tensor to carry the dtype its graph node was declared with. - Every node but `stats` is declared from q's dtype, and execute() binds raw pointers, so a - tensor of another dtype would have its bits reinterpreted with no error at all. `dout` - matters most: it arrives from autograd and is not this module's to control. + ``dout`` matters most: it arrives from autograd and is not this module's to control. """ if t.dtype != expected: raise ValueError(f"{name} must be {expected} to match q; got {t.dtype}") @@ -427,9 +409,8 @@ def _check_dtype(name: str, t: torch.Tensor, expected: torch.dtype) -> None: def _check_kv_match(k: torch.Tensor, v: torch.Tensor) -> None: """Require v to agree with k on batch, heads and sequence length. - head_dim is free: v has its own graph node and its own cache-key entry, so an asymmetric - pair builds its own plan. The other three index the same KV positions as k by definition, - and a mismatch would bind a differently shaped buffer with no error at all. + head_dim is free, since v has its own graph node and its own cache-key entry. The other + three index the same KV positions as k by definition. """ if tuple(k.shape[:3]) != tuple(v.shape[:3]): raise ValueError( @@ -438,12 +419,7 @@ def _check_kv_match(k: torch.Tensor, v: torch.Tensor) -> None: def _head_dim_strides(shape: Sequence[int], ref_strides: Sequence[int]) -> list: - """Dense strides for ``shape`` in the memory order ``ref_strides`` describes. - - O, dO and the O-shaped grads follow q's layout but carry v's head_dim, so when the two head - dims differ they cannot reuse q's strides. The graph node and the allocation both go through - here so they cannot drift apart. - """ + """Dense strides for ``shape`` in the memory order ``ref_strides`` describes.""" order = sorted(range(len(shape)), key=lambda i: ref_strides[i], reverse=True) strides = [0] * len(shape) acc = 1 @@ -456,10 +432,9 @@ def _head_dim_strides(shape: Sequence[int], ref_strides: Sequence[int]) -> list: def _o_shape_stride(shape, d_v, ref_strides): """Shape and strides for an O-shaped tensor: ``shape``'s layout carrying v's head_dim. - Works in either space. The head dim is last in both TE's bshd/sbhd and cuDNN's BHSD, and the - rule only reorders by stride magnitude, so the graph node and the allocation can each apply it - in their own space and still agree. Equal head dims keep the reference strides untouched, - which preserves a caller's non-dense view. + Works in either space, since the head dim is last in both TE's bshd/sbhd and cuDNN's BHSD and + the rule only reorders by stride magnitude. Equal head dims keep the reference strides, which + preserves a caller's non-dense view. """ out = list(shape[:3]) + [d_v] return out, (list(ref_strides) if d_v == shape[3] else _head_dim_strides(out, ref_strides)) @@ -468,11 +443,9 @@ def _o_shape_stride(shape, d_v, ref_strides): def _select_frost_plan(graph, token: str, what: str): """Select a plan whose name proves a FROST engine was chosen. - Falling back to whatever plan happens to be first would defeat the purpose. A too-old - nvidia-cutlass-dsl makes the FROST engines decline silently, and in the forward an ordinary - engine may then build and compute something else; the pin turns that into a named error at - the first forward rather than a wrong number or a backward that fails later for no visible - reason. + A too-old nvidia-cutlass-dsl makes the FROST engines decline silently, and in the forward an + ordinary engine may then build and compute something else. The pin turns that into a named + error at the first forward rather than a wrong number. """ # Both versions, because either floor can cause this and blaming one misdirects. Looked up @@ -621,10 +594,9 @@ def _bhsd(t: torch.Tensor, qkv_format: str): def _key(q, k, v, qkv_format, mask, scale, deterministic=False): - """The plan cache key. + """The plan cache key, structured per tensor rather than flattened. - Structured per tensor rather than flattened, so the builders destructure it by name instead - of by position, and so the per-tensor fragment is the same one flex_attention keys on. + The builders destructure it by name, and the per-tensor fragment is the one flex keys on. """ def described(t): From 6654a684ccbe66f6ef2a78fb1f713b60be0b0c82 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 6 Oct 2026 23:34:41 -0700 Subject: [PATCH 74/97] fix(attention): decline asymmetric head_dim again, measured on B200 cc3f6016 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 Signed-off-by: Nitin Vegesna --- docs/envvars.rst | 4 +- .../pytorch/attention/test_frost_attention.py | 66 ++++--------------- .../dot_product_attention/frost_attention.py | 14 +++- 3 files changed, 27 insertions(+), 57 deletions(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index 7353b717175..1df2ed5bcca 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -179,7 +179,7 @@ In PyTorch, the broad preference order is ``FlashAttention > FusedAttention > UnfusedDotProductAttention`` on supported pre-Hopper GPUs such as Ampere/Ada, and ``FusedAttention > FlashAttention > UnfusedDotProductAttention`` on Hopper and newer GPUs, including Blackwell. On Blackwell SM100/SM103, FusedAttention has an extra sub-backend, FROST, -which is selected only for ``head_dim`` in (256, 512] and only when the cuDNN +which is selected only for symmetric ``head_dim`` in (256, 512] and only when the cuDNN sub-backends decline; it does not change the order above. In JAX, Transformer Engine uses cuDNN fused attention when ``NVTE_FUSED_ATTN=1`` and an eligible cuDNN kernel is available; otherwise it falls back to the JAX-native implementation. See :doc:`examples/attention/attention` for a @@ -219,7 +219,7 @@ longer backend-selection overview. :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 ``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. + :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. .. envvar:: NVTE_UNFUSED_ATTN diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index da89f4b3389..461a6c8bd4a 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -59,19 +59,16 @@ def _frost_availability(): # head_dim 512 is the whole point of the backend; 320 checks the interior of the (256, 512] range # rather than only its endpoint. _SHAPES = [ - # b, hq, hkv, sq, skv, d, d_v - (2, 8, 4, 1024, 1024, 512, 512), # Gemma-4 global layer, GQA - (2, 8, 8, 512, 512, 512, 512), # MHA - (1, 4, 4, 256, 512, 512, 512), # sq != skv, which is where mask alignment matters - (2, 4, 4, 512, 512, 320, 320), # interior head_dim - # d_v != d_qk. O and the O-shaped grads take q's layout with v's head_dim, so this is the - # case that catches a plan or an allocation still built from q's trailing dimension. - (2, 8, 4, 512, 512, 512, 320), + # b, hq, hkv, sq, skv, d + (2, 8, 4, 1024, 1024, 512), # Gemma-4 global layer, GQA + (2, 8, 8, 512, 512, 512), # MHA + (1, 4, 4, 256, 512, 512), # sq != skv, which is where mask alignment matters + (2, 4, 4, 512, 512, 320), # interior head_dim ] def _shape_id(s): - return "b%d_hq%d_hkv%d_sq%d_skv%d_d%d_dv%d" % s + return "b%d_hq%d_hkv%d_sq%d_skv%d_d%d" % s def _fwd(q, k, v, mask, scale, window=None): @@ -198,12 +195,12 @@ def _floor(q32, k32, v32, scale, mask, dtype, window=None): @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) def test_frost_forward_matches_reference(shape, mask, dtype): """Forward output and LSE against an independent float64 reference.""" - b, hq, hkv, sq, skv, d, d_v = shape + b, hq, hkv, sq, skv, d = shape torch.manual_seed(0) # Generate in fp32 so there is a true high-precision original to measure against, then cast # for the kernel. bshd is what the backend takes now: it is never permuted, only described. mk = lambda s_, h_, d_: torch.randn(b, s_, h_, d_, device="cuda") - q32, k32, v32 = mk(sq, hq, d), mk(skv, hkv, d), mk(skv, hkv, d_v) + q32, k32, v32 = mk(sq, hq, d), mk(skv, hkv, d), mk(skv, hkv, d) q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) scale = 1.0 / math.sqrt(d) @@ -284,10 +281,10 @@ def test_frost_backward_matches_reference(shape, mask, window, dtype): that would show first -- the gradient of a softmax involves a subtraction of similarly sized terms, so a range problem surfaces there before it surfaces in the forward. """ - b, hq, hkv, sq, skv, d, d_v = shape + b, hq, hkv, sq, skv, d = shape torch.manual_seed(0) mk = lambda s_, h_, d_: torch.randn(b, s_, h_, d_, device="cuda") - q32, k32, v32 = mk(sq, hq, d), mk(skv, hkv, d), mk(skv, hkv, d_v) + q32, k32, v32 = mk(sq, hq, d), mk(skv, hkv, d), mk(skv, hkv, d) q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) scale = 1.0 / math.sqrt(d) @@ -368,13 +365,13 @@ def test_frost_declines_unsupported_configs(): assert ( is_frost_attention_supported(_frost_params())[0] == FusedAttnBackend.FROST ), "the supported case must be accepted" - assert ( - is_frost_attention_supported(_frost_params(head_dim_v=320))[0] == FusedAttnBackend.FROST - ), "an asymmetric head_dim pair inside the range must be accepted" for override, why in ( (dict(head_dim_qk=256, head_dim_v=256), "head_dim at the exclusive lower bound"), (dict(head_dim_v=256), "head_dim_v below the range"), + # The forward serves an asymmetric pair and the backward does not, so the selector + # declines it rather than accepting a config whose backward cannot build. + (dict(head_dim_v=320), "asymmetric head_dim"), (dict(qkv_dtype=TE_DType[torch.float32]), "fp32"), (dict(dropout=0.1), "dropout"), (dict(bias_type=AttnBiasType["post_scale_bias"]), "attention bias"), @@ -517,43 +514,6 @@ def test_frost_rejects_mismatched_kv(): _fwd(q, k, k.to(torch.float32), "no_mask", 1.0) -@requires_frost -def test_frost_serves_v_with_its_own_head_dim_and_layout(): - """v is keyed and declared separately, the way flex_attention keys each tensor. - - Both halves of that are exercised: v carries its own head_dim, and its strides differ from - k's because it is a non-contiguous slice of a wider buffer rather than a fresh allocation. - A plan built from k alone would compute either case wrongly without raising. - - v cannot differ from k in qkv_format: one format describes all three, which is what the - fused path produces and what the selector enforces. - """ - b, h, s, d, d_v = 2, 4, 512, 512, 320 - dtype = torch.bfloat16 - torch.manual_seed(0) - q32 = torch.randn(b, s, h, d, device="cuda") - k32 = torch.randn(b, s, h, d, device="cuda") - # A slice of a wider buffer, so v's strides are its own rather than k's shape re-derived. - v32 = torch.randn(b, s, h, d_v + 64, device="cuda")[..., :d_v] - q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) - assert v.stride()[:3] != k.stride()[:3], "v must not share k's strides here" - assert v.stride(3) == 1, "the head dim must stay contiguous" - scale = 1.0 / math.sqrt(d) - - out, lse = _fwd(q, k, v, "causal", scale) - out = _bhsd(out) - - floor_o, floor_l, ref_o, ref_lse = _floor( - _bhsd(q32), _bhsd(k32), _bhsd(v32), scale, "causal", dtype - ) - assert out.shape == (b, h, s, d_v), "out takes v's head_dim; got %s" % (tuple(out.shape),) - assert out.stride(3) == 1, "out must stay head-contiguous; got stride %s" % (out.stride(),) - err_o = (out.double() - ref_o).abs().max().item() - err_l = (lse.double() - ref_lse).abs().max().item() - assert err_o <= 2 * floor_o + 1e-3, "out err %.3e exceeds 2x the floor %.3e" % (err_o, floor_o) - assert err_l <= 2 * floor_l + 1e-3, "lse err %.3e exceeds 2x the floor %.3e" % (err_l, floor_l) - - def test_frost_engines_are_enabled_even_if_cudnn_was_imported_without_them(): """Enabling the FROST engines must not depend on who imported cuDNN first. diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 54c59b5a653..93cda61b564 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -297,8 +297,18 @@ def is_frost_attention_supported(params) -> Tuple[int, str]: if int(os.environ.get("NVTE_FROST_ATTN", "1")) == 0: return no_backend, "FROST is disabled by NVTE_FROST_ATTN=0" - # Each head_dim is checked on its own: q/k and v get separate graph nodes, so an - # asymmetric pair is served as long as both dims land in the range. + if params.head_dim_qk != params.head_dim_v: + # Measured on B200 with cuDNN Frontend 1.29.0: the forward serves an asymmetric pair, the + # backward does not. Its d_qk > 128 path covers only 192/128 and 256/256, and nothing + # proposes a plan otherwise. Declined outright rather than for training alone, because + # is_training is module.training and eval() does not disable autograd, so it is no + # guarantee that no backward follows. + return ( + no_backend, + f"FROST requires symmetric head_dim; got {params.head_dim_qk}/{params.head_dim_v}", + ) + # Still checked per dimension: v has its own graph node and its own cache-key entry, so the + # range applies to each rather than to one standing in for both. for name, head_dim in ( ("head_dim_qk", params.head_dim_qk), ("head_dim_v", params.head_dim_v), From 16915e9d550ed8941d19b17741869d817b9c672e Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 7 Oct 2026 00:28:07 -0700 Subject: [PATCH 75/97] fix(attention): close two guards the fused path was bypassing 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/frost_attention.py | 44 ++++++++++++++----- 1 file changed, 32 insertions(+), 12 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 93cda61b564..6efa614a79f 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -175,6 +175,29 @@ def _no(reason): _NO_WINDOW = (-1, -1) +def _window_pair(window_size) -> Tuple[int, int]: + """Normalise window_size to a (left, right) pair, declining anything that is not one. + + Its own function because _te_mask_spec unpacks the window before _mask_spec ever sees it, so + leaving these two guards inside _mask_spec left them dead on the fused path. Raised as + NotImplementedError so the selector declines, that being the only thing + is_frost_attention_supported catches; a TypeError would escape backend selection instead. + """ + if window_size is None: + return _NO_WINDOW + try: + window = tuple(window_size) + except TypeError: + raise NotImplementedError( + f"window_size must be a (left, right) pair; got {window_size!r}" + ) from None + if len(window) != 2 or not all(isinstance(bound, int) for bound in window): + raise NotImplementedError( + f"window_size must be a pair of ints (left, right); got {window!r}" + ) + return window + + def _mask_spec(attn_mask_type: str, window_size=None): """Validate a TE mask type and window, returning the hashable spec the plan is keyed on.""" if attn_mask_type not in _SUPPORTED_MASKS: @@ -182,16 +205,7 @@ def _mask_spec(attn_mask_type: str, window_size=None): f"FROST attention supports attn_mask_type in {str(_SUPPORTED_MASKS)}; got" f" {attn_mask_type!r}" ) - try: - window = _NO_WINDOW if window_size is None else tuple(window_size) - except TypeError: - # Raised as NotImplementedError so the selector declines instead of propagating out of - # backend selection, which is the only thing is_frost_attention_supported catches. - raise NotImplementedError( - f"window_size must be a (left, right) pair; got {window_size!r}" - ) from None - if len(window) != 2: - raise NotImplementedError(f"window_size must be a (left, right) pair; got {window!r}") + window = _window_pair(window_size) if window[0] < -1: # cuDNN's left bound must be >= 1, so a left of -2 would build diagonal_band_left_bound=-1 # and fail at plan build rather than declining here. @@ -243,7 +257,7 @@ def _te_mask_spec(attn_mask_type: str, window_size, bottom_right_diagonal: bool) raise NotImplementedError( f"FROST attention does not support a padding mask; got {attn_mask_type!r}" ) - left, right = _NO_WINDOW if window_size is None else tuple(window_size) + left, right = _window_pair(window_size) if "causal" in attn_mask_type and right == -1: right = 0 if right == 0: @@ -773,7 +787,13 @@ def fused_attn_bwd( # o and dO used to carry their own format into the permute; they now share qkv_format, so a # divergence would silently describe them with the wrong strides. The selector already # declines it, but this is reached directly too. - for name, fmt in (("o_format", o_format), ("do_format", do_format)): + # dqkv_layout joins them: the grads go back through qkv_format, so a divergence would return + # them in a layout the caller did not ask for, with nothing raised. + for name, fmt in ( + ("o_format", o_format), + ("do_format", do_format), + ("dqkv_layout", _qkv_format_from_layout(dqkv_layout)), + ): if fmt != qkv_format: raise NotImplementedError( f"FROST attention needs {name} to match qkv_format; got {fmt}/{qkv_format}" From 86609bae13f329e338afd4ee020f1f04fb0a426d Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 7 Oct 2026 01:13:19 -0700 Subject: [PATCH 76/97] refactor(attention): move the diagonal-band translation to cudnn_pygraph 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/cudnn_pygraph.py | 34 ++++++++++++++-- .../dot_product_attention/frost_attention.py | 39 ++----------------- 2 files changed, 34 insertions(+), 39 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index b828268a75e..92c5ab557ba 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -4,9 +4,10 @@ """Mechanics of driving cuDNN Frontend's Python graph API from PyTorch. -Importing the frontend, holding one stream-current handle per device, describing TE tensors in -cuDNN's logical BHSD form, and creating, selecting and building plans. No attention semantics, and -no knowledge of any backend's cache-key layout. +Importing the frontend, holding one stream-current handle per device, translating TE's tensor +and mask vocabulary into cuDNN's, and creating, selecting and building plans. It holds no backend +policy and no knowledge of any backend's cache-key layout: what it knows is how to say a thing to +cuDNN, not which thing to say. ``flex_attention.py`` and ``frost_attention.py`` both drive cuDNN through this API. They share this module for ownership rather than for line count: the state below is process-global -- one @@ -164,6 +165,33 @@ def bhsd_graph_tensor( return graph.tensor(dim=dim, stride=stride, data_type=tensor.dtype) +def diagonal_band_kwargs(cudnn, attn_mask_type: str, window: Tuple[int, int]) -> Dict[str, Any]: + """cuDNN sdpa kwargs for a TE (mask type, window): a diagonal alignment plus a band. + + Note the off-by-one. cuDNN's left bound counts the diagonal itself and TE's window_size does + not, so a window of w becomes a left bound of w + 1. Passing it through unconverted silently + drops one token of context per layer, which no shape-level test would catch. + + These kwargs are mutually exclusive with score_mod. cuDNN enforces that in the backward node + only ("Attention score mod enabled and hence other subgraphs are disabled"); its forward node + composes the two without complaint. Callers must still refuse the pair on both sides, because + forward and backward have to carry the same mask or the gradients belong to a different + attention than the output does. + """ + left, right = window + opts: Dict[str, Any] = {} + if attn_mask_type in ("causal", "causal_bottom_right") or right == 0: + opts["diagonal_alignment"] = ( + cudnn.diagonal_alignment.BOTTOM_RIGHT + if attn_mask_type == "causal_bottom_right" + else cudnn.diagonal_alignment.TOP_LEFT + ) + opts["diagonal_band_right_bound"] = 0 + if left != -1: + opts["diagonal_band_left_bound"] = left + 1 + return opts + + def device_key(device: torch.device) -> Tuple[Any, ...]: """Normalize a device for a cache key, so ``cuda`` and ``cuda:0`` cannot key two entries.""" if device.type == "cuda" and device.index is None: diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 6efa614a79f..f6bf184c536 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -66,33 +66,6 @@ def _import_cudnn_frontend(enable_frost_engines: bool = True): return cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=enable_frost_engines) -def _diagonal_band_kwargs(cudnn, attn_mask_type: str, window: Tuple[int, int]) -> Dict[str, Any]: - """cuDNN sdpa kwargs for a TE (mask type, window): a diagonal alignment plus a band. - - Note the off-by-one. cuDNN's left bound counts the diagonal itself and TE's window_size does - not, so a window of w becomes a left bound of w + 1. Passing it through unconverted silently - drops one token of context per layer, which no shape-level test would catch. - - These kwargs are mutually exclusive with score_mod. cuDNN enforces that in the backward node - only ("Attention score mod enabled and hence other subgraphs are disabled"); its forward node - composes the two without complaint. Callers must still refuse the pair on both sides, because - forward and backward have to carry the same mask or the gradients belong to a different - attention than the output does. - """ - left, right = window - opts: Dict[str, Any] = {} - if attn_mask_type in ("causal", "causal_bottom_right") or right == 0: - opts["diagonal_alignment"] = ( - cudnn.diagonal_alignment.BOTTOM_RIGHT - if attn_mask_type == "causal_bottom_right" - else cudnn.diagonal_alignment.TOP_LEFT - ) - opts["diagonal_band_right_bound"] = 0 - if left != -1: - opts["diagonal_band_left_bound"] = left + 1 - return opts - - def _pkg_version(name: str, module=None) -> Tuple[Optional[PkgVersion], Optional[str]]: """A package's version, or None when it is absent or unparseable, with the raw string.""" raw = None @@ -167,7 +140,7 @@ def _no(reason): # cuDNN expresses causal, bottom-right and sliding-window masking as one mechanism, a diagonal -# alignment plus a two-sided band, which is what _mask_options builds. Both alignments are needed: +# alignment plus a two-sided band, which is what diagonal_band_kwargs builds. Both alignments are needed: # the p2p ring produces square tiles where they coincide, while all_gather trims KV so they differ. _SUPPORTED_MASKS = ("no_mask", "causal", "causal_bottom_right") @@ -217,12 +190,6 @@ def _mask_spec(attn_mask_type: str, window_size=None): return attn_mask_type, window -def _mask_options(cudnn, spec): - """cuDNN sdpa kwargs for a (mask type, window) spec: a diagonal alignment plus a band.""" - attn_mask_type, window = spec - return _diagonal_band_kwargs(cudnn, attn_mask_type, window) - - _SUPPORTED_QKV_FORMATS = ("bshd", "sbhd") @@ -515,7 +482,7 @@ def _build_fwd(key, device) -> dict: v=tv, generate_stats=True, # the CP ring needs the LSE, and it is cheap attn_scale=scale, - **_mask_options(cudnn, mask), + **cudnn_pygraph.diagonal_band_kwargs(cudnn, *mask), ) tout.set_output(True).set_dim(sho).set_stride(list(o_stride)) # out: q's layout, v's head_dim tlse.set_output(True).set_dim([b, hq, sq, 1]).set_stride([hq * sq, sq, 1, 1]).set_data_type( @@ -565,7 +532,7 @@ def _build_bwd(key, device) -> dict: stats=handles["stats"], attn_scale=scale, use_deterministic_algorithm=deterministic, - **_mask_options(cudnn, mask), + **cudnn_pygraph.diagonal_band_kwargs(cudnn, *mask), ) for tensor, stride in ((tdq, qs), (tdk, ks), (tdv, vs)): tensor.set_output(True).set_data_type(io_dt).set_stride(list(stride)) From 639cdad31149092ffa240f00bb1e7abd5d47f8cf Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 7 Oct 2026 01:18:49 -0700 Subject: [PATCH 77/97] Revert "refactor(attention): move the diagonal-band translation to cudnn_pygraph" This reverts commit 86609bae. 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/cudnn_pygraph.py | 34 ++-------------- .../dot_product_attention/frost_attention.py | 39 +++++++++++++++++-- 2 files changed, 39 insertions(+), 34 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index 92c5ab557ba..b828268a75e 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -4,10 +4,9 @@ """Mechanics of driving cuDNN Frontend's Python graph API from PyTorch. -Importing the frontend, holding one stream-current handle per device, translating TE's tensor -and mask vocabulary into cuDNN's, and creating, selecting and building plans. It holds no backend -policy and no knowledge of any backend's cache-key layout: what it knows is how to say a thing to -cuDNN, not which thing to say. +Importing the frontend, holding one stream-current handle per device, describing TE tensors in +cuDNN's logical BHSD form, and creating, selecting and building plans. No attention semantics, and +no knowledge of any backend's cache-key layout. ``flex_attention.py`` and ``frost_attention.py`` both drive cuDNN through this API. They share this module for ownership rather than for line count: the state below is process-global -- one @@ -165,33 +164,6 @@ def bhsd_graph_tensor( return graph.tensor(dim=dim, stride=stride, data_type=tensor.dtype) -def diagonal_band_kwargs(cudnn, attn_mask_type: str, window: Tuple[int, int]) -> Dict[str, Any]: - """cuDNN sdpa kwargs for a TE (mask type, window): a diagonal alignment plus a band. - - Note the off-by-one. cuDNN's left bound counts the diagonal itself and TE's window_size does - not, so a window of w becomes a left bound of w + 1. Passing it through unconverted silently - drops one token of context per layer, which no shape-level test would catch. - - These kwargs are mutually exclusive with score_mod. cuDNN enforces that in the backward node - only ("Attention score mod enabled and hence other subgraphs are disabled"); its forward node - composes the two without complaint. Callers must still refuse the pair on both sides, because - forward and backward have to carry the same mask or the gradients belong to a different - attention than the output does. - """ - left, right = window - opts: Dict[str, Any] = {} - if attn_mask_type in ("causal", "causal_bottom_right") or right == 0: - opts["diagonal_alignment"] = ( - cudnn.diagonal_alignment.BOTTOM_RIGHT - if attn_mask_type == "causal_bottom_right" - else cudnn.diagonal_alignment.TOP_LEFT - ) - opts["diagonal_band_right_bound"] = 0 - if left != -1: - opts["diagonal_band_left_bound"] = left + 1 - return opts - - def device_key(device: torch.device) -> Tuple[Any, ...]: """Normalize a device for a cache key, so ``cuda`` and ``cuda:0`` cannot key two entries.""" if device.type == "cuda" and device.index is None: diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index f6bf184c536..6efa614a79f 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -66,6 +66,33 @@ def _import_cudnn_frontend(enable_frost_engines: bool = True): return cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=enable_frost_engines) +def _diagonal_band_kwargs(cudnn, attn_mask_type: str, window: Tuple[int, int]) -> Dict[str, Any]: + """cuDNN sdpa kwargs for a TE (mask type, window): a diagonal alignment plus a band. + + Note the off-by-one. cuDNN's left bound counts the diagonal itself and TE's window_size does + not, so a window of w becomes a left bound of w + 1. Passing it through unconverted silently + drops one token of context per layer, which no shape-level test would catch. + + These kwargs are mutually exclusive with score_mod. cuDNN enforces that in the backward node + only ("Attention score mod enabled and hence other subgraphs are disabled"); its forward node + composes the two without complaint. Callers must still refuse the pair on both sides, because + forward and backward have to carry the same mask or the gradients belong to a different + attention than the output does. + """ + left, right = window + opts: Dict[str, Any] = {} + if attn_mask_type in ("causal", "causal_bottom_right") or right == 0: + opts["diagonal_alignment"] = ( + cudnn.diagonal_alignment.BOTTOM_RIGHT + if attn_mask_type == "causal_bottom_right" + else cudnn.diagonal_alignment.TOP_LEFT + ) + opts["diagonal_band_right_bound"] = 0 + if left != -1: + opts["diagonal_band_left_bound"] = left + 1 + return opts + + def _pkg_version(name: str, module=None) -> Tuple[Optional[PkgVersion], Optional[str]]: """A package's version, or None when it is absent or unparseable, with the raw string.""" raw = None @@ -140,7 +167,7 @@ def _no(reason): # cuDNN expresses causal, bottom-right and sliding-window masking as one mechanism, a diagonal -# alignment plus a two-sided band, which is what diagonal_band_kwargs builds. Both alignments are needed: +# alignment plus a two-sided band, which is what _mask_options builds. Both alignments are needed: # the p2p ring produces square tiles where they coincide, while all_gather trims KV so they differ. _SUPPORTED_MASKS = ("no_mask", "causal", "causal_bottom_right") @@ -190,6 +217,12 @@ def _mask_spec(attn_mask_type: str, window_size=None): return attn_mask_type, window +def _mask_options(cudnn, spec): + """cuDNN sdpa kwargs for a (mask type, window) spec: a diagonal alignment plus a band.""" + attn_mask_type, window = spec + return _diagonal_band_kwargs(cudnn, attn_mask_type, window) + + _SUPPORTED_QKV_FORMATS = ("bshd", "sbhd") @@ -482,7 +515,7 @@ def _build_fwd(key, device) -> dict: v=tv, generate_stats=True, # the CP ring needs the LSE, and it is cheap attn_scale=scale, - **cudnn_pygraph.diagonal_band_kwargs(cudnn, *mask), + **_mask_options(cudnn, mask), ) tout.set_output(True).set_dim(sho).set_stride(list(o_stride)) # out: q's layout, v's head_dim tlse.set_output(True).set_dim([b, hq, sq, 1]).set_stride([hq * sq, sq, 1, 1]).set_data_type( @@ -532,7 +565,7 @@ def _build_bwd(key, device) -> dict: stats=handles["stats"], attn_scale=scale, use_deterministic_algorithm=deterministic, - **cudnn_pygraph.diagonal_band_kwargs(cudnn, *mask), + **_mask_options(cudnn, mask), ) for tensor, stride in ((tdq, qs), (tdk, ks), (tdv, vs)): tensor.set_output(True).set_data_type(io_dt).set_stride(list(stride)) From 0d25aa35cff46833fe62d33a9811585ff86fff0b Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 7 Oct 2026 18:13:09 -0700 Subject: [PATCH 78/97] refactor(attention): share the graph execution step 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/cudnn_pygraph.py | 19 +++++++++++++++++ .../dot_product_attention/flex_attention.py | 21 ++----------------- .../dot_product_attention/frost_attention.py | 18 +++++++++------- 3 files changed, 31 insertions(+), 27 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index b828268a75e..63e41baf731 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -290,3 +290,22 @@ def finalize_plans( f" {exc}{(' ' + hint) if hint else ''}" ) from exc return max(graph.get_workspace_size(), 1), names[hits[0]] + + +def execute_graph( + graph, + variant_pack: Dict[Any, torch.Tensor], + workspace_size: int, + device: torch.device, + *, + backend_name: str = "cuDNN attention", +): + """Allocate a built graph's workspace and run it on the device's current stream. + + The handle is resolved here rather than by the caller so it is always rebound immediately + before ``execute``, which is what keeps cuDNN on the same stream as the tensors. + """ + if device.type == "cuda" and device.index is None: + device = torch.device("cuda", torch.cuda.current_device()) + workspace = torch.empty(workspace_size, device=device, dtype=torch.uint8) + graph.execute(variant_pack, workspace, handle=handle_for(device, backend_name=backend_name)) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index f5a7c739c85..4381e58c914 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -176,12 +176,6 @@ def _wrapped_score_mod(sdpa_graph, score_tensor): return _wrapped_score_mod -def _get_cudnn_current_stream_handle(cudnn, device: torch.device): - """Return a cuDNN handle for device, bound to PyTorch's current stream.""" - del cudnn # the shared module owns the import - return cudnn_pygraph.handle_for(device, backend_name=_BACKEND_NAME) - - def _build_cudnn_pygraph(dtype: torch.dtype, device: torch.device): """Create a cuDNN frontend Python graph for F16/BF16 SDPA.""" return cudnn_pygraph.build_pygraph(dtype, device, backend_name=_BACKEND_NAME) @@ -233,19 +227,8 @@ def _execute_cudnn_graph( device: torch.device, ): """Execute a built cuDNN frontend Python graph.""" - cudnn = _import_cudnn_frontend() - - if device.type == "cuda" and device.index is None: - device = torch.device("cuda", torch.cuda.current_device()) - workspace = torch.empty( - workspace_size, - device=device, - dtype=torch.uint8, - ) - graph.execute( - variant_pack, - workspace, - handle=_get_cudnn_current_stream_handle(cudnn, device), + cudnn_pygraph.execute_graph( + graph, variant_pack, workspace_size, device, backend_name=_BACKEND_NAME ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 6efa614a79f..b05419193a5 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -724,11 +724,12 @@ def fused_attn_fwd( out_shape, out_stride = _o_shape_stride(q.shape, vd[3], q.stride()) out = torch.empty_strided(out_shape, out_stride, device=q.device, dtype=q.dtype) lse = torch.empty(qd[0], qd[1], qd[2], 1, device=q.device, dtype=torch.float32) - workspace = torch.empty(entry["workspace"], device=q.device, dtype=torch.uint8) - entry["graph"].execute( + cudnn_pygraph.execute_graph( + entry["graph"], {tq: q, tk: k, tv: v, tout: out, tlse: lse}, - workspace, - handle=cudnn_pygraph.handle_for(q.device, backend_name=_BACKEND_NAME), + entry["workspace"], + q.device, + backend_name=_BACKEND_NAME, ) # A real tensor, not None: it is saved for backward and handed to the activation offload # hooks, neither of which accepts None. FROST has no dropout, so nothing reads it. @@ -844,8 +845,8 @@ def _as(t, stride): dq = torch.empty_strided(q.shape, q.stride(), device=q.device, dtype=q.dtype) dk = torch.empty_strided(k.shape, k.stride(), device=k.device, dtype=k.dtype) dv = torch.empty_strided(v.shape, v.stride(), device=v.device, dtype=v.dtype) - workspace = torch.empty(entry["workspace"], device=q.device, dtype=torch.uint8) - entry["graph"].execute( + cudnn_pygraph.execute_graph( + entry["graph"], { h["q"]: q, h["k"]: k, @@ -857,7 +858,8 @@ def _as(t, stride): h["dk"]: dk, h["dv"]: dv, }, - workspace, - handle=cudnn_pygraph.handle_for(q.device, backend_name=_BACKEND_NAME), + entry["workspace"], + q.device, + backend_name=_BACKEND_NAME, ) return dq, dk, dv, None From 8c6c36add932b61714a22750ee4a00b79d08b475 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 7 Oct 2026 18:16:40 -0700 Subject: [PATCH 79/97] docs(attention): mark the FROST engine bar as a workaround with an expiry 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 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/dot_product_attention/cudnn_pygraph.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index 63e41baf731..2870ee6703c 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -230,6 +230,10 @@ def finalize_plans( process-wide, so declining to ask is not enough. Forgetting to exclude is silently wrong while excluding wrongly costs a slower plan or a loud decline, so the default favours the caller that does not want them; it is skipped when a plan is pinned, which has already named one. + + Remove the default once cuDNN is fixed: the gate that should decline these graphs reads a key + ``sdpa()`` never writes, so it never fires. The bar is not self-verifying either, matching + engine names by a substring that a rename would silently break. """ cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() From f01a0ee6f6a131953830de697cad18473e633c70 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 7 Oct 2026 18:30:01 -0700 Subject: [PATCH 80/97] docs(attention): say what actually verifies the FROST engine bar 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 Signed-off-by: Nitin Vegesna --- .../attention/dot_product_attention/cudnn_pygraph.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index 2870ee6703c..5de005bd139 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -232,8 +232,10 @@ def finalize_plans( that does not want them; it is skipped when a plan is pinned, which has already named one. Remove the default once cuDNN is fixed: the gate that should decline these graphs reads a key - ``sdpa()`` never writes, so it never fires. The bar is not self-verifying either, matching - engine names by a substring that a rename would silently break. + ``sdpa()`` never writes, so it never fires. The match is on a substring of the engine name, so + a rename would break it; what catches that is the end-to-end test comparing flex's output + across the switch, not anything here. Measured: deselect marks rather than removes, so the + plan list still names a barred engine afterwards and cannot be asserted on. """ cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() From 5019bf33228502fadd3b48ca73d760992505811f Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 7 Oct 2026 18:32:03 -0700 Subject: [PATCH 81/97] test(attention): let a lane require the flex cuDNN tests to actually run 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 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_flex_attention.py | 23 +++++++++++++++++++ 1 file changed, 23 insertions(+) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index ad4f37a25c3..0d9794e003c 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -23,6 +23,29 @@ get_available_attention_backends, ) + +def _flex_availability(): + """Why the cuDNN score_mod path cannot run here, or None if it can.""" + if not torch.cuda.is_available(): + return "no CUDA device" + try: + flex_attention._import_cudnn_frontend() + except ImportError as exc: + return "nvidia-cudnn-frontend not importable: %s" % exc + return None + + +_SKIP = _flex_availability() +# Mirrors NVTE_GDN_TEST_REQUIRED in test_gdn_attention.py. Most tests here run on CPU with the +# builder monkeypatched; the two that reach cuDNN -- the numerical one and the FROST-switch one -- +# 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 usual: the switch test is what verifies the +# engine bar in cudnn_pygraph, which cannot be checked from the plan list. +# +# Deliberately not set in qa/L0_pytorch_unittest, which cannot be assumed to carry the package. +if os.getenv("NVTE_FLEX_TEST_REQUIRED", "0") == "1" and _SKIP is not None: + raise RuntimeError("NVTE_FLEX_TEST_REQUIRED=1, but flex attention is unavailable: %s" % _SKIP) + param_types = [torch.float16] if torch.cuda.is_available() and is_bf16_available(): param_types.append(torch.bfloat16) From 598db0de25acb9df1ce11040886cc3f86fa473b1 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 7 Oct 2026 18:34:15 -0700 Subject: [PATCH 82/97] test(attention): trim the flex required-guard comment 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 Signed-off-by: Nitin Vegesna --- tests/pytorch/attention/test_flex_attention.py | 8 +------- 1 file changed, 1 insertion(+), 7 deletions(-) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index 0d9794e003c..39431ce3462 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -36,13 +36,7 @@ def _flex_availability(): _SKIP = _flex_availability() -# Mirrors NVTE_GDN_TEST_REQUIRED in test_gdn_attention.py. Most tests here run on CPU with the -# builder monkeypatched; the two that reach cuDNN -- the numerical one and the FROST-switch one -- -# 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 usual: the switch test is what verifies the -# engine bar in cudnn_pygraph, which cannot be checked from the plan list. -# -# Deliberately not set in qa/L0_pytorch_unittest, which cannot be assumed to carry the package. +# Mirrors NVTE_GDN_TEST_REQUIRED: the tests that reach cuDNN skip silently without the frontend. if os.getenv("NVTE_FLEX_TEST_REQUIRED", "0") == "1" and _SKIP is not None: raise RuntimeError("NVTE_FLEX_TEST_REQUIRED=1, but flex attention is unavailable: %s" % _SKIP) From 8176c3012fa6623f6f76ac637e75867eefeae06d Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 7 Oct 2026 18:35:00 -0700 Subject: [PATCH 83/97] test(attention): trim the FROST guard comments to match the siblings 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 Signed-off-by: Nitin Vegesna --- tests/pytorch/attention/test_frost_attention.py | 9 +++------ 1 file changed, 3 insertions(+), 6 deletions(-) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 461a6c8bd4a..ea3fc854d0b 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -46,14 +46,11 @@ def _frost_availability(): _SKIP = _frost_availability() -# Mirrors NVTE_GDN_TEST_REQUIRED in test_gdn_attention.py. These tests skip on any machine that -# cannot reach the backend, which on most CI hardware is every machine; setting this on a lane -# that is supposed to cover FROST turns a silent skip into a loud failure. +# Mirrors NVTE_GDN_TEST_REQUIRED: these skip on any machine that cannot reach the backend. if os.getenv("NVTE_FROST_TEST_REQUIRED", "0") == "1" and _SKIP is not None: raise RuntimeError("NVTE_FROST_TEST_REQUIRED=1, but FrostAttention is unavailable: %s" % _SKIP) -# Applied per test rather than as a module-level pytestmark: the ONNX-export regression -# below guards a code path that runs on every GPU, so gating it on Blackwell would skip it -# exactly where the bug it covers can still occur. +# Per test, not a module-level pytestmark: the ONNX-export regression below runs on every GPU, so +# gating it on Blackwell would skip it exactly where its bug can still occur. requires_frost = pytest.mark.skipif(_SKIP is not None, reason=str(_SKIP)) # head_dim 512 is the whole point of the backend; 320 checks the interior of the (256, 512] range From a88556be0bd2e1a245f5c3672e2912146e8c86a0 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 7 Oct 2026 18:49:08 -0700 Subject: [PATCH 84/97] test(attention): cut the FROST suite's CI cost where it buys nothing 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 Signed-off-by: Nitin Vegesna --- .../attention/test_attention_with_cp.py | 62 ++++++++++------ .../pytorch/attention/test_frost_attention.py | 71 ++++--------------- 2 files changed, 54 insertions(+), 79 deletions(-) diff --git a/tests/pytorch/attention/test_attention_with_cp.py b/tests/pytorch/attention/test_attention_with_cp.py index ebe91da63fc..529001faee0 100644 --- a/tests/pytorch/attention/test_attention_with_cp.py +++ b/tests/pytorch/attention/test_attention_with_cp.py @@ -461,11 +461,14 @@ def test_cp_with_flash_attention_softcap(cp_pool, cp_comm_type): # cuDNN FROST: head_dim in (256, 512] on SM100/SM103, the range no other backend # serves together with context parallelism. Shapes are Gemma-4 global layers, which is what -# motivated the backend. seqlen must stay divisible by cp_size * 2 for causal load balancing. +# motivated the backend. seqlen must stay divisible by cp_size * 2 for causal load balancing, +# and is otherwise the cheapest axis there is: attention is O(s^2) and the ring does the same +# work per step whatever the length, so 2048 exercises every path 4096 would at a quarter the +# cost. These are d512, four times the head_dim of the fused and flash configs beside them. model_configs_frost_attn = { # test: ModelConfig(b, sq, hq, dqk) - "cp_hd512_0": ModelConfig(2, 4096, 8, 512, num_gqa_groups=4, attn_mask_type="causal"), - "cp_hd512_1": ModelConfig(2, 4096, 8, 512, num_gqa_groups=4, attn_mask_type="no_mask"), + "cp_hd512_0": ModelConfig(2, 2048, 8, 512, num_gqa_groups=4, attn_mask_type="causal"), + "cp_hd512_1": ModelConfig(2, 2048, 8, 512, num_gqa_groups=4, attn_mask_type="no_mask"), "cp_hd512_2": ModelConfig(2, 2048, 8, 512, num_gqa_groups=8, attn_mask_type="causal"), } @@ -777,17 +780,12 @@ def _frost_availability(): @pytest.mark.parametrize("model", model_configs_frost_attn.keys()) @pytest.mark.parametrize("qkv_format", ["bshd", "sbhd"]) -@pytest.mark.parametrize("cp_comm_type", ["p2p", "all_gather", "a2a", "a2a+p2p"]) +@pytest.mark.parametrize("cp_comm_type", ["p2p", "all_gather", "a2a"]) def test_cp_with_frost_attention(cp_pool, model, qkv_format, cp_comm_type): """Context parallelism at head_dim 512, which no other backend serves. thd is excluded because the backend declines it: it needs varlen support that is not - implemented. - - a2a+p2p needs four ranks rather than two -- an a2a subgroup crossed with a p2p subgroup -- and - exercises no new attention code: it dispatches to the same AttnFuncWithCPAndKVP2P as plain p2p, - with an a2a communication stage on either side of the ring. It is covered here so that claim is - measured rather than assumed. + implemented. a2a+p2p is covered separately below, since it needs four ranks. """ reason = _frost_availability() if reason is not None: @@ -797,18 +795,8 @@ def test_cp_with_frost_attention(cp_pool, model, qkv_format, cp_comm_type): config.context_parallel = True config.cp_comm_type = cp_comm_type - # a2a requires num_heads and num_gqa_groups divisible by the a2a subgroup size; every config - # here satisfies that, but assert rather than rely on it staying true. - if cp_comm_type == "a2a+p2p": - assert config.num_heads % 2 == 0 and config.num_gqa_groups % 2 == 0, ( - f"cp_comm_type=a2a+p2p needs num_heads ({config.num_heads}) and num_gqa_groups" - f" ({config.num_gqa_groups}) divisible by the a2a subgroup size" - ) - - pool = cp_pool(4 if cp_comm_type == "a2a+p2p" else 2) - _submit( - pool, + cp_pool(2), dtype="bf16", model=model, qkv_format=qkv_format, @@ -819,6 +807,38 @@ def test_cp_with_frost_attention(cp_pool, model, qkv_format, cp_comm_type): ) +def test_cp_with_frost_attention_a2a_p2p(cp_pool): + """One case, because a2a+p2p composes rather than adding a path. + + It dispatches to the same AttnFuncWithCPAndKVP2P as plain p2p, with an a2a communication stage + on either side of the ring, and it needs four ranks rather than two. One case says the + composition works; six would pay four-rank prices to re-cover p2p. + """ + reason = _frost_availability() + if reason is not None: + pytest.skip(reason) + + config = model_configs_frost_attn["cp_hd512_0"] + config.context_parallel = True + config.cp_comm_type = "a2a+p2p" + # The a2a stage shards heads across its subgroup, so both counts must divide by it. + assert config.num_heads % 2 == 0 and config.num_gqa_groups % 2 == 0, ( + f"a2a+p2p needs num_heads ({config.num_heads}) and num_gqa_groups" + f" ({config.num_gqa_groups}) divisible by the a2a subgroup size" + ) + + _submit( + cp_pool(4), + dtype="bf16", + model="cp_hd512_0", + qkv_format="bshd", + kernel_backend="FrostAttention", + cp_comm_type="a2a+p2p", + is_training=True, + log_level=pytest_logging_level, + ) + + @pytest.mark.parametrize("cp_comm_type", ["p2p", "all_gather", "a2a"]) def test_cp_with_frost_attention_fp16(cp_pool, cp_comm_type): """One fp16 arm per comm type, since the matrix above is bf16 throughout. diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index ea3fc854d0b..3f3a9023632 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -186,10 +186,18 @@ def _floor(q32, k32, v32, scale, mask, dtype, window=None): ) +# bf16 everywhere, fp16 on two shapes. What fp16 risks that bf16 does not is its narrower +# exponent range, and that surfaces in the backward, which runs both dtypes on every shape it +# covers. Crossing it with every forward shape pays for the same information twice. +_FWD_CASES = [(s, torch.bfloat16) for s in _SHAPES] + [(s, torch.float16) for s in _SHAPES[:2]] +_FWD_IDS = [ + "%s_%s" % (_shape_id(s), "bf16" if d is torch.bfloat16 else "fp16") for s, d in _FWD_CASES +] + + @requires_frost -@pytest.mark.parametrize("shape", _SHAPES, ids=_shape_id) +@pytest.mark.parametrize("shape,dtype", _FWD_CASES, ids=_FWD_IDS) @pytest.mark.parametrize("mask", ["no_mask", "causal", "causal_bottom_right"]) -@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) def test_frost_forward_matches_reference(shape, mask, dtype): """Forward output and LSE against an independent float64 reference.""" b, hq, hkv, sq, skv, d = shape @@ -228,7 +236,9 @@ def test_frost_forward_matches_reference(shape, mask, dtype): @requires_frost -@pytest.mark.parametrize("window", [(256, 0), (128, 0), (0, 0)], ids=lambda w: "win%d" % w[0]) +# (128, 0) is the ordinary case and (0, 0) the degenerate diagonal-only one, which is where an +# off-by-one in the band would show. A second ordinary width tests the same arithmetic again. +@pytest.mark.parametrize("window", [(128, 0), (0, 0)], ids=lambda w: "win%d" % w[0]) @pytest.mark.parametrize("mask", ["causal", "causal_bottom_right", "no_mask"]) @pytest.mark.parametrize("sq,skv", [(1024, 1024), (512, 1024)], ids=["square", "rect"]) def test_frost_sliding_window_matches_reference(mask, window, sq, skv): @@ -442,61 +452,6 @@ def test_frost_mask_spec_rejects_malformed_windows(): _mask_spec("causal", window), why -@requires_frost -@pytest.mark.parametrize( - "cp_comm_type,window,expect_frost", - [ - ("all_gather", (128, 0), True), - ("a2a", (128, 0), True), - ("p2p", (128, 0), False), - ("a2a+p2p", (128, 0), False), - ("p2p", (-1, 0), True), - ("p2p", (-1, -1), True), - ], -) -def test_frost_sliding_window_selection_by_cp_comm_type(cp_comm_type, window, expect_frost): - """Which context-parallel paths may serve a sliding window. - - all_gather and a2a each see a contiguous KV range, so the window applies unchanged. The p2p - ring shards KV across steps, so a bound measured against the full sequence does not survive - the per-step tiles -- the same rule FusedAttention carries. The cases without a real window - must still select FROST, since the decline has to key on the window and not on p2p itself. - """ - from transformer_engine.pytorch.attention.dot_product_attention.utils import ( - AttentionParams, - get_attention_backend, - ) - - params = AttentionParams( - qkv_dtype=torch.bfloat16, - qkv_layout="bshd_bshd_bshd", - batch_size=2, - num_heads=8, - num_gqa_groups=4, - max_seqlen_q=4096, - max_seqlen_kv=4096, - head_dim_qk=512, - head_dim_v=512, - attn_mask_type="causal", - window_size=window, - context_parallel=True, - cp_comm_type=cp_comm_type, - is_training=True, - ) - from transformer_engine.pytorch.cpp_extensions.fused_attn import FusedAttnBackend - - use_fused, fused_backend = get_attention_backend(params)[2:4] - use_frost = bool(use_fused) and fused_backend == FusedAttnBackend.FROST - assert ( - use_frost == expect_frost - ), "cp_comm_type=%s window=%s: expected the FROST sub-backend=%s, got %s" % ( - cp_comm_type, - window, - expect_frost, - use_frost, - ) - - @requires_frost def test_frost_rejects_mismatched_kv(): """v must index the same KV positions as k. head_dim is free; the rest is not.""" From c01f48b9e5e88eb11c280a80abcef58205871113 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 7 Oct 2026 18:51:04 -0700 Subject: [PATCH 85/97] test(attention): put the cp_comm_type selector test back Removed in a88556be 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 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_frost_attention.py | 58 +++++++++++++++++++ 1 file changed, 58 insertions(+) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 3f3a9023632..8476bf27604 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -452,6 +452,64 @@ def test_frost_mask_spec_rejects_malformed_windows(): _mask_spec("causal", window), why +# Kept despite the CI-cost review: it launches no kernel, so removing it would have been +# coverage given up for no time back. It is the only end-to-end get_attention_backend call +# here, and the only place FROST, a sliding window and context parallelism meet. +@requires_frost +@pytest.mark.parametrize( + "cp_comm_type,window,expect_frost", + [ + ("all_gather", (128, 0), True), + ("a2a", (128, 0), True), + ("p2p", (128, 0), False), + ("a2a+p2p", (128, 0), False), + ("p2p", (-1, 0), True), + ("p2p", (-1, -1), True), + ], +) +def test_frost_sliding_window_selection_by_cp_comm_type(cp_comm_type, window, expect_frost): + """Which context-parallel paths may serve a sliding window. + + all_gather and a2a each see a contiguous KV range, so the window applies unchanged. The p2p + ring shards KV across steps, so a bound measured against the full sequence does not survive + the per-step tiles -- the same rule FusedAttention carries. The cases without a real window + must still select FROST, since the decline has to key on the window and not on p2p itself. + """ + from transformer_engine.pytorch.attention.dot_product_attention.utils import ( + AttentionParams, + get_attention_backend, + ) + + params = AttentionParams( + qkv_dtype=torch.bfloat16, + qkv_layout="bshd_bshd_bshd", + batch_size=2, + num_heads=8, + num_gqa_groups=4, + max_seqlen_q=4096, + max_seqlen_kv=4096, + head_dim_qk=512, + head_dim_v=512, + attn_mask_type="causal", + window_size=window, + context_parallel=True, + cp_comm_type=cp_comm_type, + is_training=True, + ) + from transformer_engine.pytorch.cpp_extensions.fused_attn import FusedAttnBackend + + use_fused, fused_backend = get_attention_backend(params)[2:4] + use_frost = bool(use_fused) and fused_backend == FusedAttnBackend.FROST + assert ( + use_frost == expect_frost + ), "cp_comm_type=%s window=%s: expected the FROST sub-backend=%s, got %s" % ( + cp_comm_type, + window, + expect_frost, + use_frost, + ) + + @requires_frost def test_frost_rejects_mismatched_kv(): """v must index the same KV positions as k. head_dim is free; the rest is not.""" From adbc9517328dbd7c2ee47e2b6751080b6ee96c53 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Wed, 7 Oct 2026 18:53:05 -0700 Subject: [PATCH 86/97] Revert "test(attention): put the cp_comm_type selector test back" This reverts commit c01f48b9. 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 Signed-off-by: Nitin Vegesna --- .../pytorch/attention/test_frost_attention.py | 58 ------------------- 1 file changed, 58 deletions(-) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py index 8476bf27604..3f3a9023632 100644 --- a/tests/pytorch/attention/test_frost_attention.py +++ b/tests/pytorch/attention/test_frost_attention.py @@ -452,64 +452,6 @@ def test_frost_mask_spec_rejects_malformed_windows(): _mask_spec("causal", window), why -# Kept despite the CI-cost review: it launches no kernel, so removing it would have been -# coverage given up for no time back. It is the only end-to-end get_attention_backend call -# here, and the only place FROST, a sliding window and context parallelism meet. -@requires_frost -@pytest.mark.parametrize( - "cp_comm_type,window,expect_frost", - [ - ("all_gather", (128, 0), True), - ("a2a", (128, 0), True), - ("p2p", (128, 0), False), - ("a2a+p2p", (128, 0), False), - ("p2p", (-1, 0), True), - ("p2p", (-1, -1), True), - ], -) -def test_frost_sliding_window_selection_by_cp_comm_type(cp_comm_type, window, expect_frost): - """Which context-parallel paths may serve a sliding window. - - all_gather and a2a each see a contiguous KV range, so the window applies unchanged. The p2p - ring shards KV across steps, so a bound measured against the full sequence does not survive - the per-step tiles -- the same rule FusedAttention carries. The cases without a real window - must still select FROST, since the decline has to key on the window and not on p2p itself. - """ - from transformer_engine.pytorch.attention.dot_product_attention.utils import ( - AttentionParams, - get_attention_backend, - ) - - params = AttentionParams( - qkv_dtype=torch.bfloat16, - qkv_layout="bshd_bshd_bshd", - batch_size=2, - num_heads=8, - num_gqa_groups=4, - max_seqlen_q=4096, - max_seqlen_kv=4096, - head_dim_qk=512, - head_dim_v=512, - attn_mask_type="causal", - window_size=window, - context_parallel=True, - cp_comm_type=cp_comm_type, - is_training=True, - ) - from transformer_engine.pytorch.cpp_extensions.fused_attn import FusedAttnBackend - - use_fused, fused_backend = get_attention_backend(params)[2:4] - use_frost = bool(use_fused) and fused_backend == FusedAttnBackend.FROST - assert ( - use_frost == expect_frost - ), "cp_comm_type=%s window=%s: expected the FROST sub-backend=%s, got %s" % ( - cp_comm_type, - window, - expect_frost, - use_frost, - ) - - @requires_frost def test_frost_rejects_mismatched_kv(): """v must index the same KV positions as k. head_dim is free; the rest is not.""" From 6c66bca8f48539d4cecd999836eac3145aeff0c9 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 8 Oct 2026 11:57:00 -0700 Subject: [PATCH 87/97] test(attention): put the CP sequence length back, it saved nothing 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 Signed-off-by: Nitin Vegesna --- tests/pytorch/attention/test_attention_with_cp.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/tests/pytorch/attention/test_attention_with_cp.py b/tests/pytorch/attention/test_attention_with_cp.py index 529001faee0..da42c71d616 100644 --- a/tests/pytorch/attention/test_attention_with_cp.py +++ b/tests/pytorch/attention/test_attention_with_cp.py @@ -461,14 +461,14 @@ def test_cp_with_flash_attention_softcap(cp_pool, cp_comm_type): # cuDNN FROST: head_dim in (256, 512] on SM100/SM103, the range no other backend # serves together with context parallelism. Shapes are Gemma-4 global layers, which is what -# motivated the backend. seqlen must stay divisible by cp_size * 2 for causal load balancing, -# and is otherwise the cheapest axis there is: attention is O(s^2) and the ring does the same -# work per step whatever the length, so 2048 exercises every path 4096 would at a quarter the -# cost. These are d512, four times the head_dim of the fused and flash configs beside them. +# motivated the backend. seqlen must stay divisible by cp_size * 2 for causal load balancing. +# Measured: halving it saves nothing. A distributed case costs ~7.8 s whatever the length, since +# pool spawn, plan JIT and NCCL setup dominate and the attention itself disappears into them. +# Case count is the only lever here. model_configs_frost_attn = { # test: ModelConfig(b, sq, hq, dqk) - "cp_hd512_0": ModelConfig(2, 2048, 8, 512, num_gqa_groups=4, attn_mask_type="causal"), - "cp_hd512_1": ModelConfig(2, 2048, 8, 512, num_gqa_groups=4, attn_mask_type="no_mask"), + "cp_hd512_0": ModelConfig(2, 4096, 8, 512, num_gqa_groups=4, attn_mask_type="causal"), + "cp_hd512_1": ModelConfig(2, 4096, 8, 512, num_gqa_groups=4, attn_mask_type="no_mask"), "cp_hd512_2": ModelConfig(2, 2048, 8, 512, num_gqa_groups=8, attn_mask_type="causal"), } From dbb95004fddc3dd411a06ec5625899c1f77cc323 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 8 Oct 2026 11:58:19 -0700 Subject: [PATCH 88/97] test(attention): trim the CP seqlen comment The measurement belongs in the commit that made it, not beside the configs. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- tests/pytorch/attention/test_attention_with_cp.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/tests/pytorch/attention/test_attention_with_cp.py b/tests/pytorch/attention/test_attention_with_cp.py index da42c71d616..69e38188a68 100644 --- a/tests/pytorch/attention/test_attention_with_cp.py +++ b/tests/pytorch/attention/test_attention_with_cp.py @@ -461,10 +461,8 @@ def test_cp_with_flash_attention_softcap(cp_pool, cp_comm_type): # cuDNN FROST: head_dim in (256, 512] on SM100/SM103, the range no other backend # serves together with context parallelism. Shapes are Gemma-4 global layers, which is what -# motivated the backend. seqlen must stay divisible by cp_size * 2 for causal load balancing. -# Measured: halving it saves nothing. A distributed case costs ~7.8 s whatever the length, since -# pool spawn, plan JIT and NCCL setup dominate and the attention itself disappears into them. -# Case count is the only lever here. +# motivated the backend. seqlen must stay divisible by cp_size * 2 for causal load balancing, +# and is not a CI-time lever: halving it was measured to save nothing. model_configs_frost_attn = { # test: ModelConfig(b, sq, hq, dqk) "cp_hd512_0": ModelConfig(2, 4096, 8, 512, num_gqa_groups=4, attn_mask_type="causal"), From 029e0d702fab159a3203ad1355c86550c5f53dd8 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 8 Oct 2026 11:58:53 -0700 Subject: [PATCH 89/97] test(attention): drop the CP seqlen note Back to the comment as it was. The divisibility constraint is the only thing a reader needs here. Co-Authored-By: Claude Opus 5 Signed-off-by: Nitin Vegesna --- tests/pytorch/attention/test_attention_with_cp.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/tests/pytorch/attention/test_attention_with_cp.py b/tests/pytorch/attention/test_attention_with_cp.py index 69e38188a68..3712de9a906 100644 --- a/tests/pytorch/attention/test_attention_with_cp.py +++ b/tests/pytorch/attention/test_attention_with_cp.py @@ -461,8 +461,7 @@ def test_cp_with_flash_attention_softcap(cp_pool, cp_comm_type): # cuDNN FROST: head_dim in (256, 512] on SM100/SM103, the range no other backend # serves together with context parallelism. Shapes are Gemma-4 global layers, which is what -# motivated the backend. seqlen must stay divisible by cp_size * 2 for causal load balancing, -# and is not a CI-time lever: halving it was measured to save nothing. +# motivated the backend. seqlen must stay divisible by cp_size * 2 for causal load balancing. model_configs_frost_attn = { # test: ModelConfig(b, sq, hq, dqk) "cp_hd512_0": ModelConfig(2, 4096, 8, 512, num_gqa_groups=4, attn_mask_type="causal"), From f23a8508125b241fef0fe25b82c737c44f9224fb Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 8 Oct 2026 12:00:53 -0700 Subject: [PATCH 90/97] test(attention): point the single a2a+p2p case at the shortest model 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 Signed-off-by: Nitin Vegesna --- tests/pytorch/attention/test_attention_with_cp.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/tests/pytorch/attention/test_attention_with_cp.py b/tests/pytorch/attention/test_attention_with_cp.py index 3712de9a906..dbd237b4b24 100644 --- a/tests/pytorch/attention/test_attention_with_cp.py +++ b/tests/pytorch/attention/test_attention_with_cp.py @@ -815,7 +815,9 @@ def test_cp_with_frost_attention_a2a_p2p(cp_pool): if reason is not None: pytest.skip(reason) - config = model_configs_frost_attn["cp_hd512_0"] + # The shortest model, since this is the only four-rank case and so the only one that would + # take all of a2a+p2p down with it if it timed out. + config = model_configs_frost_attn["cp_hd512_2"] config.context_parallel = True config.cp_comm_type = "a2a+p2p" # The a2a stage shards heads across its subgroup, so both counts must divide by it. @@ -827,7 +829,7 @@ def test_cp_with_frost_attention_a2a_p2p(cp_pool): _submit( cp_pool(4), dtype="bf16", - model="cp_hd512_0", + model="cp_hd512_2", qkv_format="bshd", kernel_backend="FrostAttention", cp_comm_type="a2a+p2p", From b66c078bef21e0159127cbe88626891b8ec8267c Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 8 Oct 2026 14:41:16 -0700 Subject: [PATCH 91/97] refactor(attention): build the cuDNN SDPA graphs through cudnn_pygraph 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/backends.py | 9 +- .../dot_product_attention/cudnn_pygraph.py | 188 +++++++++++++++++- .../dot_product_attention/flex_attention.py | 178 +++++++---------- .../dot_product_attention/frost_attention.py | 151 ++++++-------- .../attention/dot_product_attention/utils.py | 6 +- 5 files changed, 312 insertions(+), 220 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index 6e5c331b902..c132d49d9e4 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -2036,14 +2036,7 @@ def _fused_attn_setup_ctx( bwd_args.softmax_type = fwd_args.softmax_type bwd_args.window_size = fwd_args.window_size bwd_args.bottom_right_diagonal = fwd_args.bottom_right_diagonal - # FROST has to survive this: it is a python sub-backend, so re-deriving F16_arbitrary_seqlen - # here would send its backward to the C++ path, which does not serve these head dims. - saved_fused_attention_backend = ctx_attrs["fused_attention_backend"] - bwd_args.fused_attention_backend = ( - saved_fused_attention_backend - if fp8 or saved_fused_attention_backend == FusedAttnBackend["FROST"] - else FusedAttnBackend["F16_arbitrary_seqlen"] - ) + bwd_args.fused_attention_backend = ctx_attrs["fused_attention_backend"] bwd_args.use_FAv2_bwd = fwd_args.use_FAv2_bwd bwd_args.deterministic = fwd_args.deterministic diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index 5de005bd139..fb56016ee33 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -2,18 +2,18 @@ # # See LICENSE for license information. -"""Mechanics of driving cuDNN Frontend's Python graph API from PyTorch. +"""Shared helpers for driving cuDNN frontend's Python graph API from PyTorch. -Importing the frontend, holding one stream-current handle per device, describing TE tensors in -cuDNN's logical BHSD form, and creating, selecting and building plans. No attention semantics, and -no knowledge of any backend's cache-key layout. +Covers importing the frontend, keeping one stream-current handle per device, describing TE +tensors in cuDNN's logical BHSD form, building the SDPA forward and backward graphs, and +creating, selecting and building their plans. What a backend wants from a graph -- a mask, a +score_mod, which engine to pin or bar -- is passed in rather than decided here. -``flex_attention.py`` and ``frost_attention.py`` both drive cuDNN through this API. They share -this module for ownership rather than for line count: the state below is process-global -- one -``cudnn`` module, one ``CUDNN_FRONTEND_ENABLE_FROST_ENGINES`` switch, one engine ranking, one -handle per device -- and giving it two owners is how this code has produced bugs before. +``flex_attention.py`` and ``frost_attention.py`` both build their graphs through this module. Its +state is process-global -- the ``cudnn`` module, the ``CUDNN_FRONTEND_ENABLE_FROST_ENGINES`` +switch and the per-device handles -- so it is set up in one place instead of twice. -``backend_name`` is threaded through purely so a failure still says which backend was driving. +``backend_name`` is passed through so an error says which backend was running. """ from __future__ import annotations @@ -298,6 +298,176 @@ def finalize_plans( return max(graph.get_workspace_size(), 1), names[hits[0]] +def _declare_tensor(graph, name: str, spec: Any): + """Declare one graph input. + + ``spec`` is a ``torch.Tensor``, described with ``tensor_like``, or a ``(dim, stride)`` / + ``(dim, stride, data_type)`` descriptor. Omitting ``data_type`` lets the tensor inherit the + graph's ``io_data_type``. + """ + if isinstance(spec, torch.Tensor): + return graph.tensor_like(spec) + kwargs: Dict[str, Any] = {"name": name, "dim": list(spec[0]), "stride": list(spec[1])} + if len(spec) > 2 and spec[2] is not None: + kwargs["data_type"] = spec[2] + return graph.tensor(**kwargs) + + +def _mark_output(tensor, spec: Any): + """Mark a graph output and apply whichever of dim/stride/data_type ``spec`` supplies. + + Each field is optional because the backends specify different subsets, and setting one a + caller left out would be describing the tensor for it rather than from it. + """ + tensor.set_output(True) + if spec[0] is not None: + tensor.set_dim(list(spec[0])) + if spec[1] is not None: + tensor.set_stride(list(spec[1])) + if len(spec) > 2 and spec[2] is not None: + tensor.set_data_type(spec[2]) + return tensor + + +def _declare_aux(graph, aux_tensors): + """Declare the auxiliary runtime tensors an sdpa callback reads, grouped by role.""" + return { + group: {name: graph.tensor_like(t) for name, t in tensors.items()} + for group, tensors in (aux_tensors or {}).items() + } + + +def _resolve_sdpa_kwargs(sdpa_kwargs, aux): + """Extra ``sdpa``/``sdpa_backward`` arguments, as a dict or a callable taking ``aux``. + + The callable form exists because a score_mod closes over graph tensors that cannot be built + until the graph is, so the caller gets them handed back here rather than building its own. + """ + return sdpa_kwargs(aux) if callable(sdpa_kwargs) else dict(sdpa_kwargs or {}) + + +def build_fwd( + *, + dtype: torch.dtype, + device: torch.device, + backend_name: str, + name: str, + q: Any, + k: Any, + v: Any, + out: Any, + attn_scale: float, + stats: Any = None, + aux_tensors: Optional[Dict[str, Dict[str, torch.Tensor]]] = None, + sdpa_kwargs: Any = None, + heuristics: Optional[Sequence[Any]] = None, + require_plan_token: Optional[str] = None, + not_found_hint: Any = "", + exclude_plan_tokens: Any = _BAR_FROST_BY_DEFAULT, +) -> Dict[str, Any]: + """Build and plan an SDPA forward graph. + + ``stats`` is the LSE output descriptor, or ``None`` to skip generating it. Everything + backend-specific -- a mask, a score_mod, which engine to pin or bar -- arrives through + ``sdpa_kwargs`` and the plan arguments, which are passed to :func:`finalize_plans`. + """ + graph = build_pygraph(dtype, device, backend_name=backend_name) + tq = _declare_tensor(graph, "q", q) + tk = _declare_tensor(graph, "k", k) + tv = _declare_tensor(graph, "v", v) + aux = _declare_aux(graph, aux_tensors) + tout, tstats = graph.sdpa( + name=name, + q=tq, + k=tk, + v=tv, + generate_stats=stats is not None, + attn_scale=attn_scale, + **_resolve_sdpa_kwargs(sdpa_kwargs, aux), + ) + _mark_output(tout, out) + if stats is None: + tstats = None + else: + _mark_output(tstats, stats) + workspace, plan = finalize_plans( + graph, + backend_name=backend_name, + heuristics=heuristics, + require_plan_token=require_plan_token, + not_found_hint=not_found_hint, + exclude_plan_tokens=exclude_plan_tokens, + ) + return { + "graph": graph, + "q": tq, + "k": tk, + "v": tv, + "out": tout, + "stats": tstats, + "aux": aux, + "workspace": workspace, + "plan": plan, + } + + +def build_bwd( + *, + dtype: torch.dtype, + device: torch.device, + backend_name: str, + name: str, + q: Any, + k: Any, + v: Any, + o: Any, + do: Any, + stats: Any, + dq: Any, + dk: Any, + dv: Any, + attn_scale: float, + deterministic: bool = False, + aux_tensors: Optional[Dict[str, Dict[str, torch.Tensor]]] = None, + sdpa_kwargs: Any = None, + heuristics: Optional[Sequence[Any]] = None, + require_plan_token: Optional[str] = None, + not_found_hint: Any = "", + exclude_plan_tokens: Any = _BAR_FROST_BY_DEFAULT, +) -> Dict[str, Any]: + """Build and plan an SDPA backward graph. The counterpart of :func:`build_fwd`.""" + graph = build_pygraph(dtype, device, backend_name=backend_name) + handles = { + n: _declare_tensor(graph, n, spec) + for n, spec in (("q", q), ("k", k), ("v", v), ("o", o), ("do", do), ("stats", stats)) + } + aux = _declare_aux(graph, aux_tensors) + tdq, tdk, tdv = graph.sdpa_backward( + name=name, + q=handles["q"], + k=handles["k"], + v=handles["v"], + o=handles["o"], + dO=handles["do"], + stats=handles["stats"], + attn_scale=attn_scale, + use_deterministic_algorithm=deterministic, + **_resolve_sdpa_kwargs(sdpa_kwargs, aux), + ) + for handle, spec in ((tdq, dq), (tdk, dk), (tdv, dv)): + _mark_output(handle, spec) + workspace, plan = finalize_plans( + graph, + backend_name=backend_name, + heuristics=heuristics, + require_plan_token=require_plan_token, + not_found_hint=not_found_hint, + exclude_plan_tokens=exclude_plan_tokens, + ) + handles.update({"dq": tdq, "dk": tdk, "dv": tdv}) + return {"graph": graph, "aux": aux, "workspace": workspace, "plan": plan, **handles} + + def execute_graph( graph, variant_pack: Dict[Any, torch.Tensor], diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index 4381e58c914..74911250db5 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -33,11 +33,6 @@ def _bhsd_dim_stride( return cudnn_pygraph.bhsd_dim_stride(tensor, tensor_format, backend_name=_BACKEND_NAME) -def _bhsd_graph_tensor(graph, tensor: torch.Tensor, tensor_format: str): - """Create a cuDNN graph tensor with BHSD dims and TE-layout strides.""" - return cudnn_pygraph.bhsd_graph_tensor(graph, tensor, tensor_format, backend_name=_BACKEND_NAME) - - # score_mod graph cache helpers. def _freeze_score_mod_cache_key(value: Any) -> Any: """Convert a user-provided score_mod graph key into a hashable structure.""" @@ -157,13 +152,6 @@ def _score_mod_bhsd_tensor_metadata(tensor: torch.Tensor, tensor_format: str) -> ) -def _make_cudnn_graph_tensor_dict(graph, tensors: Optional[Dict[str, torch.Tensor]]): - """Create cuDNN graph tensors matching runtime tensors.""" - if tensors is None: - return {} - return {name: graph.tensor_like(tensor) for name, tensor in tensors.items()} - - # cuDNN frontend score_mod graph helpers. def _wrap_score_mod(score_mod: Optional[Callable], graph_tensors: Dict[str, Any]): """Adapt TE's score_mod signature to cuDNN frontend's two-argument callback.""" @@ -176,9 +164,10 @@ def _wrapped_score_mod(sdpa_graph, score_tensor): return _wrapped_score_mod -def _build_cudnn_pygraph(dtype: torch.dtype, device: torch.device): - """Create a cuDNN frontend Python graph for F16/BF16 SDPA.""" - return cudnn_pygraph.build_pygraph(dtype, device, backend_name=_BACKEND_NAME) +def _finalize_cudnn_graph(graph) -> int: + """Build a cuDNN frontend Python graph and return its workspace size.""" + workspace_size, _ = cudnn_pygraph.finalize_plans(graph, backend_name=_BACKEND_NAME) + return workspace_size @dataclass @@ -214,12 +203,6 @@ class _CudnnScoreModBwdGraphEntry: workspace_size: int -def _finalize_cudnn_graph(graph) -> int: - """Build a cuDNN frontend Python graph and return its workspace size.""" - workspace_size, _ = cudnn_pygraph.finalize_plans(graph, backend_name=_BACKEND_NAME) - return workspace_size - - def _execute_cudnn_graph( graph, variant_pack: Dict[Any, torch.Tensor], @@ -309,6 +292,12 @@ def _cudnn_score_mod_bwd_cache_key( ) +def _bhsd_tensor_spec(tensor: torch.Tensor, tensor_format: str): + """Describe a tensor for ``cudnn_pygraph``: BHSD dims with TE-layout strides.""" + dim, stride = _bhsd_dim_stride(tensor, tensor_format) + return (dim, stride, tensor.dtype) + + def _build_cudnn_score_mod_fwd_graph( is_training: bool, query_layer: torch.Tensor, @@ -324,46 +313,35 @@ def _build_cudnn_score_mod_fwd_graph( ) -> _CudnnScoreModFwdGraphEntry: """Build a cached cuDNN frontend graph for score_mod fprop.""" cudnn = _import_cudnn_frontend() + if is_training: + assert stats is not None - graph = _build_cudnn_pygraph(query_layer.dtype, query_layer.device) - q = _bhsd_graph_tensor(graph, query_layer, q_format) - k = _bhsd_graph_tensor(graph, key_layer, kv_format) - v = _bhsd_graph_tensor(graph, value_layer, kv_format) - - score_mod_graph_tensors = _make_cudnn_graph_tensor_dict(graph, score_mod_tensors) - wrapped_score_mod = _wrap_score_mod(score_mod, score_mod_graph_tensors) - - output_dim, output_stride = _bhsd_dim_stride(output_layer, q_format) - output, stats_tensor = graph.sdpa( + entry = cudnn_pygraph.build_fwd( + dtype=query_layer.dtype, + device=query_layer.device, + backend_name=_BACKEND_NAME, name="te_score_mod_sdpa", - q=q, - k=k, - v=v, - generate_stats=is_training, + q=_bhsd_tensor_spec(query_layer, q_format), + k=_bhsd_tensor_spec(key_layer, kv_format), + v=_bhsd_tensor_spec(value_layer, kv_format), + out=_bhsd_dim_stride(output_layer, q_format), + stats=((stats.size(), stats.stride(), cudnn.data_type.FLOAT) if is_training else None), attn_scale=attn_scale, - use_causal_mask=False, - score_mod=wrapped_score_mod, + aux_tensors={"score_mod": score_mod_tensors or {}}, + sdpa_kwargs=lambda aux: { + "use_causal_mask": False, + "score_mod": _wrap_score_mod(score_mod, aux["score_mod"]), + }, ) - output.set_output(True).set_dim(output_dim).set_stride(output_stride) - - if is_training: - assert stats is not None - stats_tensor.set_output(True).set_dim(stats.size()).set_stride( - stats.stride() - ).set_data_type(cudnn.data_type.FLOAT) - else: - stats_tensor = None - - workspace_size = _finalize_cudnn_graph(graph) return _CudnnScoreModFwdGraphEntry( - graph=graph, - q=q, - k=k, - v=v, - output=output, - stats=stats_tensor, - score_mod_graph_tensors=score_mod_graph_tensors, - workspace_size=workspace_size, + graph=entry["graph"], + q=entry["q"], + k=entry["k"], + v=entry["v"], + output=entry["out"], + stats=entry["stats"], + score_mod_graph_tensors=entry["aux"]["score_mod"], + workspace_size=entry["workspace"], ) @@ -419,62 +397,52 @@ def _build_cudnn_score_mod_bwd_graph( deterministic: bool, ) -> _CudnnScoreModBwdGraphEntry: """Build a cached cuDNN frontend graph for score_mod bprop.""" - graph = _build_cudnn_pygraph(query_layer.dtype, query_layer.device) - q = _bhsd_graph_tensor(graph, query_layer, q_format) - k = _bhsd_graph_tensor(graph, key_layer, kv_format) - v = _bhsd_graph_tensor(graph, value_layer, kv_format) - output = _bhsd_graph_tensor(graph, output_layer, q_format) - d_output = _bhsd_graph_tensor(graph, d_out, q_format) - stats_tensor = graph.tensor_like(stats) - - score_mod_graph_tensors = _make_cudnn_graph_tensor_dict(graph, score_mod_tensors) - score_mod_bprop_graph_tensors = ( - _make_cudnn_graph_tensor_dict(graph, score_mod_bprop_tensors) - if score_mod_bprop is not None - else {} - ) - wrapped_score_mod = _wrap_score_mod(score_mod, score_mod_graph_tensors) - wrapped_score_mod_bprop = _wrap_score_mod(score_mod_bprop, score_mod_bprop_graph_tensors) - dq_layer = torch.empty_like(query_layer) dk_layer = torch.empty_like(key_layer) dv_layer = torch.empty_like(value_layer) - dq_dim, dq_stride = _bhsd_dim_stride(dq_layer, q_format) - dk_dim, dk_stride = _bhsd_dim_stride(dk_layer, kv_format) - dv_dim, dv_stride = _bhsd_dim_stride(dv_layer, kv_format) - dq, dk, dv = graph.sdpa_backward( + + entry = cudnn_pygraph.build_bwd( + dtype=query_layer.dtype, + device=query_layer.device, + backend_name=_BACKEND_NAME, name="te_score_mod_sdpa_backward", - q=q, - k=k, - v=v, - o=output, - dO=d_output, - stats=stats_tensor, + q=_bhsd_tensor_spec(query_layer, q_format), + k=_bhsd_tensor_spec(key_layer, kv_format), + v=_bhsd_tensor_spec(value_layer, kv_format), + o=_bhsd_tensor_spec(output_layer, q_format), + do=_bhsd_tensor_spec(d_out, q_format), + stats=stats, + dq=_bhsd_dim_stride(dq_layer, q_format), + dk=_bhsd_dim_stride(dk_layer, kv_format), + dv=_bhsd_dim_stride(dv_layer, kv_format), attn_scale=attn_scale, - use_causal_mask=False, - score_mod=wrapped_score_mod, - score_mod_bprop=wrapped_score_mod_bprop, - use_deterministic_algorithm=deterministic, + deterministic=deterministic, + aux_tensors={ + "score_mod": score_mod_tensors or {}, + "score_mod_bprop": ( + score_mod_bprop_tensors or {} if score_mod_bprop is not None else {} + ), + }, + sdpa_kwargs=lambda aux: { + "use_causal_mask": False, + "score_mod": _wrap_score_mod(score_mod, aux["score_mod"]), + "score_mod_bprop": _wrap_score_mod(score_mod_bprop, aux["score_mod_bprop"]), + }, ) - dq.set_output(True).set_dim(dq_dim).set_stride(dq_stride) - dk.set_output(True).set_dim(dk_dim).set_stride(dk_stride) - dv.set_output(True).set_dim(dv_dim).set_stride(dv_stride) - - workspace_size = _finalize_cudnn_graph(graph) return _CudnnScoreModBwdGraphEntry( - graph=graph, - q=q, - k=k, - v=v, - output=output, - d_output=d_output, - stats=stats_tensor, - dq=dq, - dk=dk, - dv=dv, - score_mod_graph_tensors=score_mod_graph_tensors, - score_mod_bprop_graph_tensors=score_mod_bprop_graph_tensors, - workspace_size=workspace_size, + graph=entry["graph"], + q=entry["q"], + k=entry["k"], + v=entry["v"], + output=entry["o"], + d_output=entry["do"], + stats=entry["stats"], + dq=entry["dq"], + dk=entry["dk"], + dv=entry["dv"], + score_mod_graph_tensors=entry["aux"]["score_mod"], + score_mod_bprop_graph_tensors=entry["aux"]["score_mod_bprop"], + workspace_size=entry["workspace"], ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index b05419193a5..d2b6284251d 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -2,15 +2,12 @@ # # See LICENSE for license information. -"""cuDNN FROST attention backend for head_dim in (256, 512] on SM100/SM103. +"""cuDNN FROST attention. This feature is **experimental and subject to change**. -**Experimental and subject to change.** The engines this wraps are themselves experimental in -cuDNN Frontend, and if the fused path gains these shapes this backend may be folded into it. - -Why a separate Python backend rather than teaching the existing C++ fused path: FROST engines are -registered at Python import time behind CUDNN_FRONTEND_ENABLE_FROST_ENGINES and require the -nvidia-cutlass-dsl Python package, while TE's C++ builds against cuDNN Frontend headers only. -Reaching them requires a Python graph, which is what this module is. +FROST runs through cudnn-frontend's Python API and is registered behind the +CUDNN_FRONTEND_ENABLE_FROST_ENGINES=1 flag. Different from FusedAttnBackend.F16_arbitrary_seqlen +and FusedAttnBackend.FP8, FusedAttnBackend.FROST is Python only and requires the Python +installation of cudnn-frontend, not just its C++ header files. """ from __future__ import annotations @@ -130,7 +127,7 @@ def _no(reason): return _no("no CUDA device") if os.environ.get("CUDNN_FRONTEND_ENABLE_FROST_ENGINES", "1") == "0": # Explicitly switched off. Declining here is the difference between falling back cleanly - # and raising from _select_frost_plan once a plan is built. + # and raising from _frost_plan_args once a plan is built. return _no("CUDNN_FRONTEND_ENABLE_FROST_ENGINES=0 disables the FROST engines") if torch.cuda.get_device_capability() not in _SUPPORTED_ARCHS: major, minor = torch.cuda.get_device_capability() @@ -143,7 +140,7 @@ def _no(reason): return _no(f"nvidia-cudnn-frontend not importable: {exc}") # Decline only on positive evidence: a version below a floor, or a package absent outright. - # An unparseable version defers to _select_frost_plan, which checks the plan by name. + # An unparseable version defers to _frost_plan_args, which checks the plan by name. frontend, frontend_raw = _pkg_version("nvidia-cudnn-frontend", cudnn_pygraph.cudnn_module()) if frontend is not None and frontend < _MIN_CUDNN_FRONTEND: return _no( @@ -311,18 +308,8 @@ def is_frost_attention_supported(params) -> Tuple[int, str]: if int(os.environ.get("NVTE_FROST_ATTN", "1")) == 0: return no_backend, "FROST is disabled by NVTE_FROST_ATTN=0" - if params.head_dim_qk != params.head_dim_v: - # Measured on B200 with cuDNN Frontend 1.29.0: the forward serves an asymmetric pair, the - # backward does not. Its d_qk > 128 path covers only 192/128 and 256/256, and nothing - # proposes a plan otherwise. Declined outright rather than for training alone, because - # is_training is module.training and eval() does not disable autograd, so it is no - # guarantee that no backward follows. - return ( - no_backend, - f"FROST requires symmetric head_dim; got {params.head_dim_qk}/{params.head_dim_v}", - ) - # Still checked per dimension: v has its own graph node and its own cache-key entry, so the - # range applies to each rather than to one standing in for both. + # Checked per dimension: v has its own graph node and its own cache-key entry, so the range + # applies to each rather than to one standing in for both. for name, head_dim in ( ("head_dim_qk", params.head_dim_qk), ("head_dim_v", params.head_dim_v), @@ -334,6 +321,17 @@ def is_frost_attention_supported(params) -> Tuple[int, str]: no_backend, f"FROST needs {name} to be a multiple of {_HEAD_DIM_MULTIPLE}; got {head_dim}", ) + # Reached only once both dims are in range, which is the point: cuDNN's backward does plan + # asymmetric pairs below it (dqk192/dv128), so widening the range must revisit this rule. + if params.head_dim_qk != params.head_dim_v: + # Measured on B200 with cuDNN Frontend 1.29.0: above d_qk=128 the backward serves only + # 192/128 and 256/256, so no asymmetric pair in (256, 512] has a plan. Declined for both + # directions, since is_training follows module.training and eval() leaves autograd on. + return ( + no_backend, + "FROST requires symmetric head_dim in (256, 512]; got" + f" {params.head_dim_qk}/{params.head_dim_v}", + ) qkv_dtype = TORCH_DType.get(params.qkv_dtype) if qkv_dtype not in (torch.bfloat16, torch.float16): @@ -464,8 +462,8 @@ def _o_shape_stride(shape, d_v, ref_strides): return out, (list(ref_strides) if d_v == shape[3] else _head_dim_strides(out, ref_strides)) -def _select_frost_plan(graph, token: str, what: str): - """Select a plan whose name proves a FROST engine was chosen. +def _frost_plan_args(token: str, what: str) -> dict: + """Plan arguments that pin an engine whose name proves FROST was chosen. A too-old nvidia-cutlass-dsl makes the FROST engines decline silently, and in the forward an ordinary engine may then build and compute something else. The pin turns that into a named @@ -485,14 +483,11 @@ def hint(): ) cudnn = _import_cudnn_frontend() - _, name = cudnn_pygraph.finalize_plans( - graph, - backend_name=_BACKEND_NAME, - heuristics=[cudnn.heur_mode.A], - require_plan_token=token, - not_found_hint=hint, - ) - return name + return { + "heuristics": [cudnn.heur_mode.A], + "require_plan_token": token, + "not_found_hint": hint, + } def _build_fwd(key, device) -> dict: @@ -503,31 +498,21 @@ def _build_fwd(key, device) -> dict: _dev, (shq, qs, dtype), (shk, ks, _), (shv, vs, _), mask, scale, _deterministic = key b, hq, sq = shq[0], shq[1], shq[2] sho, o_stride = _o_shape_stride(shq, shv[3], qs) - - graph = cudnn_pygraph.build_pygraph(dtype, device, backend_name=_BACKEND_NAME) - tq = graph.tensor(name="q", dim=list(shq), stride=list(qs)) - tk = graph.tensor(name="k", dim=list(shk), stride=list(ks)) - tv = graph.tensor(name="v", dim=list(shv), stride=list(vs)) - tout, tlse = graph.sdpa( + return cudnn_pygraph.build_fwd( + dtype=dtype, + device=device, + backend_name=_BACKEND_NAME, name="frost_fwd", - q=tq, - k=tk, - v=tv, - generate_stats=True, # the CP ring needs the LSE, and it is cheap + q=(shq, qs), + k=(shk, ks), + v=(shv, vs), + out=(sho, o_stride), # out: q's layout, v's head_dim + # the CP ring needs the LSE, and it is cheap + stats=([b, hq, sq, 1], [hq * sq, sq, 1, 1], cudnn.data_type.FLOAT), attn_scale=scale, - **_mask_options(cudnn, mask), + sdpa_kwargs=_mask_options(cudnn, mask), + **_frost_plan_args(_FROST_FWD_PLAN_TOKEN, "forward"), ) - tout.set_output(True).set_dim(sho).set_stride(list(o_stride)) # out: q's layout, v's head_dim - tlse.set_output(True).set_dim([b, hq, sq, 1]).set_stride([hq * sq, sq, 1, 1]).set_data_type( - cudnn.data_type.FLOAT - ) - plan = _select_frost_plan(graph, _FROST_FWD_PLAN_TOKEN, "forward") - return { - "graph": graph, - "handles": (tq, tk, tv, tout, tlse), - "workspace": max(graph.get_workspace_size(), 1), - "plan": plan, - } def _build_bwd(key, device) -> dict: @@ -537,46 +522,26 @@ def _build_bwd(key, device) -> dict: io_dt = cudnn_pygraph.io_data_type(cudnn, dtype, backend_name=_BACKEND_NAME) b, hq, sq = shq[0], shq[1], shq[2] sho, o_stride = _o_shape_stride(shq, shv[3], qs) - - graph = cudnn_pygraph.build_pygraph(dtype, device, backend_name=_BACKEND_NAME) - handles = {} - # Each grad is declared with the layout of the tensor it differentiates. - for name, shape, stride in ( - ("q", shq, qs), - ("k", shk, ks), - ("v", shv, vs), - ("o", sho, o_stride), - ("do", sho, o_stride), - ): - handles[name] = graph.tensor(name=name, dim=list(shape), stride=list(stride)) - handles["stats"] = graph.tensor( - name="stats", - dim=[b, hq, sq, 1], - stride=[hq * sq, sq, 1, 1], - data_type=cudnn.data_type.FLOAT, - ) - tdq, tdk, tdv = graph.sdpa_backward( + return cudnn_pygraph.build_bwd( + dtype=dtype, + device=device, + backend_name=_BACKEND_NAME, name="frost_bwd", - q=handles["q"], - k=handles["k"], - v=handles["v"], - o=handles["o"], - dO=handles["do"], - stats=handles["stats"], + q=(shq, qs), + k=(shk, ks), + v=(shv, vs), + o=(sho, o_stride), + do=(sho, o_stride), + stats=([b, hq, sq, 1], [hq * sq, sq, 1, 1], cudnn.data_type.FLOAT), + # Each grad is declared with the layout of the tensor it differentiates. + dq=(None, qs, io_dt), + dk=(None, ks, io_dt), + dv=(None, vs, io_dt), attn_scale=scale, - use_deterministic_algorithm=deterministic, - **_mask_options(cudnn, mask), + deterministic=deterministic, + sdpa_kwargs=_mask_options(cudnn, mask), + **_frost_plan_args(_FROST_BWD_PLAN_TOKEN, "backward"), ) - for tensor, stride in ((tdq, qs), (tdk, ks), (tdv, vs)): - tensor.set_output(True).set_data_type(io_dt).set_stride(list(stride)) - plan = _select_frost_plan(graph, _FROST_BWD_PLAN_TOKEN, "backward") - handles["dq"], handles["dk"], handles["dv"] = tdq, tdk, tdv - return { - "graph": graph, - "handles": handles, - "workspace": max(graph.get_workspace_size(), 1), - "plan": plan, - } def _cached(kind: str, key, device): @@ -715,7 +680,7 @@ def fused_attn_fwd( qd, _, vd = _validate_qkv(q, k, v, qkv_format) scale = attn_scale if attn_scale is not None else qd[3] ** -0.5 entry = _cached("fwd", _key(q, k, v, qkv_format, mask, scale), q.device) - tq, tk, tv, tout, tlse = entry["handles"] + tq, tk, tv, tout, tlse = (entry[n] for n in ("q", "k", "v", "out", "stats")) # Allocated per call so concurrent uses cannot alias; the cache holds only the plan. # empty_strided, not empty_like: the latter does not preserve an arbitrary permuted stride. @@ -826,7 +791,7 @@ def fused_attn_bwd( scale = attn_scale if attn_scale is not None else qd[3] ** -0.5 entry = _cached("bwd", _key(q, k, v, qkv_format, mask, scale, deterministic), q.device) - h = entry["handles"] + h = entry if softmax_lse.dim() == 3: softmax_lse = softmax_lse.unsqueeze(-1) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index b30429fc542..ca03b427f28 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -452,8 +452,6 @@ def _get_fused_attn_backend(**fused_attn_kwargs): params = FusedAttentionParams(**fused_attn_kwargs) fused_attention_backend, reject_message = tex.get_fused_attn_backend(params) if fused_attention_backend == FusedAttnBackend.No_Backend: - # A python sub-backend, invisible to the C++ selector. Availability is checked at the - # end of get_attention_backend, the way flash-attn's version is. from .frost_attention import ( # pylint: disable=import-outside-toplevel is_frost_attention_supported, ) @@ -461,7 +459,7 @@ def _get_fused_attn_backend(**fused_attn_kwargs): frost_backend, frost_reject = is_frost_attention_supported(params) if frost_backend != FusedAttnBackend.No_Backend: return int(frost_backend), frost_reject - reject_message = f"{reject_message} {frost_reject}" + reject_message = f"{reject_message.rstrip('. ')}. {frost_reject}" return int(fused_attention_backend), reject_message @@ -1878,8 +1876,6 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt use_flash_attention_4 = False use_flash_attention = use_flash_attention_2 or use_flash_attention_3 or use_flash_attention_4 if use_fused_attention and fused_attention_backend == FusedAttnBackend.FROST.value: - # Deferred to here: probing imports cuDNN Frontend with the engines enabled, which - # changes the engine pool for every cuDNN consumer in the process. from .frost_attention import ( # pylint: disable=import-outside-toplevel is_frost_attention_available, ) From 5c099bfa8e16f2d4fe40562d97cab2358ce0d80d Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 8 Oct 2026 21:45:27 +0000 Subject: [PATCH 92/97] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- .../attention/dot_product_attention/frost_attention.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index d2b6284251d..e5917150036 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -329,8 +329,10 @@ def is_frost_attention_supported(params) -> Tuple[int, str]: # directions, since is_training follows module.training and eval() leaves autograd on. return ( no_backend, - "FROST requires symmetric head_dim in (256, 512]; got" - f" {params.head_dim_qk}/{params.head_dim_v}", + ( + "FROST requires symmetric head_dim in (256, 512]; got" + f" {params.head_dim_qk}/{params.head_dim_v}" + ), ) qkv_dtype = TORCH_DType.get(params.qkv_dtype) From b1632699e86c90932a5e5fd244e4241e9fdbfcec Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 8 Oct 2026 15:19:10 -0700 Subject: [PATCH 93/97] refactor(attention): drop the FROST layout checks cuDNN already makes _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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/frost_attention.py | 24 ++++++------------- 1 file changed, 7 insertions(+), 17 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index e5917150036..8dc3293aa9e 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -410,17 +410,6 @@ def is_frost_attention_supported(params) -> Tuple[int, str]: return int(FusedAttnBackend.FROST), "" -def _check_layout(name: str, t: torch.Tensor) -> None: - """Validate a 4D view whose head dimension is contiguous, which the kernels assume.""" - if t.dim() != 4: - raise ValueError(f"{name} must be 4D [b, h, s, d]; got {tuple(t.shape)}") - if t.stride(3) != 1: - raise ValueError( - f"{name} must have a contiguous head dimension; got shape {tuple(t.shape)} stride" - f" {tuple(t.stride())}" - ) - - def _check_dtype(name: str, t: torch.Tensor, expected: torch.dtype) -> None: """Require a tensor to carry the dtype its graph node was declared with. @@ -558,13 +547,15 @@ def _cached(kind: str, key, device): def _validate_qkv(q, k, v, qkv_format): """Check the tensors the graph will bind, and return their BHSD descriptions. - These are not stylistic guards. ``execute`` binds raw pointers, so a tensor whose shape, - dtype or layout disagrees with the node it is bound to is reinterpreted rather than rejected. - The context-parallel ring calls the backward outside autograd, so neither direction may - assume the other ran first. + These are not stylistic guards. ``execute`` binds raw pointers, so a tensor whose shape or + dtype disagrees with the node it is bound to is reinterpreted rather than rejected. The + context-parallel ring calls the backward outside autograd, so neither direction may assume + the other ran first. + + Rank and head-dim contiguity are not checked here: cuDNN's frost_sdpa families register + _sdpa_validate.validate_graph, which enforces both in graph.validate(). """ for name, tensor in (("q", q), ("k", k), ("v", v)): - _check_layout(name, tensor) _check_dtype(name, tensor, q.dtype) _check_kv_match(k, v) qd, _ = _bhsd(q, qkv_format) @@ -775,7 +766,6 @@ def fused_attn_bwd( o, d_o = o.contiguous(), d_o.contiguous() qd, _, vd = _validate_qkv(q, k, v, qkv_format) for name, tensor in (("o", o), ("d_o", d_o)): - _check_layout(name, tensor) _check_dtype(name, tensor, q.dtype) o_shape, o_stride = _o_shape_stride(q.shape, vd[3], q.stride()) for name, tensor in (("o", o), ("d_o", d_o)): From 7364535205742b377d51e8465da0ef7e3360fbf2 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 8 Oct 2026 15:48:27 -0700 Subject: [PATCH 94/97] fix(attention): restore the FROST rank check, which cuDNN cannot make 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 Signed-off-by: Nitin Vegesna --- .../attention/dot_product_attention/frost_attention.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 8dc3293aa9e..cdf98c27094 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -552,10 +552,15 @@ def _validate_qkv(q, k, v, qkv_format): context-parallel ring calls the backward outside autograd, so neither direction may assume the other ran first. - Rank and head-dim contiguity are not checked here: cuDNN's frost_sdpa families register - _sdpa_validate.validate_graph, which enforces both in graph.validate(). + Rank is checked here because cuDNN cannot: _sdpa_validate inspects the descriptor this + module builds, and bhsd_dim_stride reads dims 0-3, so a higher-rank tensor is described as + 4D with its trailing dims dropped rather than rejected. The head-dim stride does reach the + descriptor, and _sdpa_validate rejects a non-unit one in graph.validate(), so it is not + repeated here. """ for name, tensor in (("q", q), ("k", k), ("v", v)): + if tensor.dim() != 4: + raise ValueError(f"{name} must be 4D; got shape {tuple(tensor.shape)}") _check_dtype(name, tensor, q.dtype) _check_kv_match(k, v) qd, _ = _bhsd(q, qkv_format) From 8a680d57fc6d848d312c4e77b99f5d83feaa5a26 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 8 Oct 2026 15:54:20 -0700 Subject: [PATCH 95/97] docs(attention): cite the engine capability for the FROST symmetry rule 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/frost_attention.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index cdf98c27094..1ad180acc0f 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -43,6 +43,8 @@ _MIN_CUDNN_FRONTEND = PkgVersion("1.29.0") _SUPPORTED_ARCHS = ((10, 0), (10, 3)) +# These mirror sdpa_bwd_sm100's declared envelope: d_envelope_floor=256, d={512}, +# d_pad_multiple=8. _MAX_HEAD_DIM = 512 _MIN_HEAD_DIM = 257 # below this the existing cuDNN/flash backends already serve the shape # The engine pads head_dim to a multiple of 8, so 260 is in range but not servable. Declined @@ -321,12 +323,12 @@ def is_frost_attention_supported(params) -> Tuple[int, str]: no_backend, f"FROST needs {name} to be a multiple of {_HEAD_DIM_MULTIPLE}; got {head_dim}", ) - # Reached only once both dims are in range, which is the point: cuDNN's backward does plan - # asymmetric pairs below it (dqk192/dv128), so widening the range must revisit this rule. + # Scoped to this range deliberately. sdpa_bwd_sm100 is the only f16 FROST backward on + # SM100/SM103 and leaves dqk_ge_dv unset, which cuDNN reads as requiring d_qk == d_v; its + # sm80 and sm120 siblings do set it and serve rectangular pairs such as 192/128. if params.head_dim_qk != params.head_dim_v: - # Measured on B200 with cuDNN Frontend 1.29.0: above d_qk=128 the backward serves only - # 192/128 and 256/256, so no asymmetric pair in (256, 512] has a plan. Declined for both - # directions, since is_training follows module.training and eval() leaves autograd on. + # Declined for both directions, since is_training follows module.training and eval() + # leaves autograd on, so an eval-mode call is no promise that no backward follows. return ( no_backend, ( From 80672e33bd90982da8151eb27be71b3529606852 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 8 Oct 2026 16:09:43 -0700 Subject: [PATCH 96/97] fix(attention): check the whole softmax_lse shape, not just its leading 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/frost_attention.py | 23 +++++++++++-------- 1 file changed, 13 insertions(+), 10 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py index 1ad180acc0f..081d4b9f49e 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -323,9 +323,9 @@ def is_frost_attention_supported(params) -> Tuple[int, str]: no_backend, f"FROST needs {name} to be a multiple of {_HEAD_DIM_MULTIPLE}; got {head_dim}", ) - # Scoped to this range deliberately. sdpa_bwd_sm100 is the only f16 FROST backward on - # SM100/SM103 and leaves dqk_ge_dv unset, which cuDNN reads as requiring d_qk == d_v; its - # sm80 and sm120 siblings do set it and serve rectangular pairs such as 192/128. + # The rule is the backward engine's, not this range's. sdpa_bwd_sm100 is the only f16 FROST + # backward on SM100/SM103 and leaves dqk_ge_dv unset, which cuDNN reads as requiring + # d_qk == d_v at every head dim. The forward serves 192/128 here, so only the backward binds. if params.head_dim_qk != params.head_dim_v: # Declined for both directions, since is_training follows module.training and eval() # leaves autograd on, so an eval-mode call is no promise that no backward follows. @@ -555,10 +555,10 @@ def _validate_qkv(q, k, v, qkv_format): the other ran first. Rank is checked here because cuDNN cannot: _sdpa_validate inspects the descriptor this - module builds, and bhsd_dim_stride reads dims 0-3, so a higher-rank tensor is described as - 4D with its trailing dims dropped rather than rejected. The head-dim stride does reach the - descriptor, and _sdpa_validate rejects a non-unit one in graph.validate(), so it is not - repeated here. + module builds, and bhsd_dim_stride reads dims 0-3, so a rank-5 tensor is described as 4D. + The plan cache key carries strides but not rank, so such a tensor can even reuse a valid 4D + plan and be bound by pointer. The head-dim stride does reach the descriptor, so + _sdpa_validate rejects a non-unit one in graph.validate() and it is not repeated here. """ for name, tensor in (("q", q), ("k", k), ("v", v)): if tensor.dim() != 4: @@ -782,10 +782,13 @@ def fused_attn_bwd( raise ValueError(f"softmax_lse must be fp32; got {softmax_lse.dtype}") # Compared against the BHSD description, not q's own shape: the LSE is always [b, h, s] # whatever format the tensors arrived in. - if tuple(softmax_lse.shape[:3]) != tuple(qd[:3]): + # The whole shape, not just the leading dims: the stats node is declared [b, h, s, 1], so a + # trailing dim of any other size would be bound by pointer and read with strides that do not + # describe it. Same hole the rank check on q/k/v closes. + lse_shape = tuple(softmax_lse.shape) + if lse_shape not in (tuple(qd[:3]), tuple(qd[:3]) + (1,)): raise ValueError( - f"softmax_lse must be [b, h, s] matching q; got {tuple(softmax_lse.shape)} and" - f" {tuple(qd[:3])}" + f"softmax_lse must be [b, h, s] or [b, h, s, 1] matching q; got {lse_shape}" ) scale = attn_scale if attn_scale is not None else qd[3] ** -0.5 From 0effc6ccab9d3de4a097e14af1c4cb87348c7ae0 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 8 Oct 2026 16:22:15 -0700 Subject: [PATCH 97/97] docs(attention): shorten the cudnn_pygraph module docstring 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 Signed-off-by: Nitin Vegesna --- .../dot_product_attention/cudnn_pygraph.py | 16 +++++----------- 1 file changed, 5 insertions(+), 11 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py index fb56016ee33..b49e13fd7eb 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -2,18 +2,12 @@ # # See LICENSE for license information. -"""Shared helpers for driving cuDNN frontend's Python graph API from PyTorch. +"""Helpers for driving cuDNN frontend's Python graph API from PyTorch. -Covers importing the frontend, keeping one stream-current handle per device, describing TE -tensors in cuDNN's logical BHSD form, building the SDPA forward and backward graphs, and -creating, selecting and building their plans. What a backend wants from a graph -- a mask, a -score_mod, which engine to pin or bar -- is passed in rather than decided here. - -``flex_attention.py`` and ``frost_attention.py`` both build their graphs through this module. Its -state is process-global -- the ``cudnn`` module, the ``CUDNN_FRONTEND_ENABLE_FROST_ENGINES`` -switch and the per-device handles -- so it is set up in one place instead of twice. - -``backend_name`` is passed through so an error says which backend was running. +Shared by flex_attention.py and frost_attention.py: importing the frontend, holding a handle per +device, describing tensors in cuDNN's BHSD form, and building and planning the SDPA forward and +backward graphs. Anything specific to one backend, such as a mask, a score_mod or which engine +to pin, is passed in by the caller. """ from __future__ import annotations