From 7765b76555034aec27fb5b8a719240aaabbaf350 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Tue, 15 Sep 2026 20:58:54 -0700 Subject: [PATCH 01/69] 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 0000000000..e7b2e2c054 --- /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/69] 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 e5d7169da1..723a2ebd07 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 6b66458985..0571153240 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 9a339233a4..355da80830 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 658dab5d88..13e9aec1ec 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 23b89287df..ded2de7455 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/69] 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 510d14ac63..024a52fab9 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/69] 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 8d1b870d3a..e820f212e9 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 8b85c30057..9542157930 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/69] [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 024a52fab9..f3478e611c 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 ded2de7455..b8670926c9 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/69] 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 9542157930..d1dd818721 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/69] 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 e7b2e2c054..37db3c9f5f 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/69] 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 37db3c9f5f..c15e114dea 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/69] 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 355da80830..8ce1a336b4 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 c15e114dea..9eb1f6c5d9 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/69] 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 e820f212e9..842be0400d 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 d665df4cef..d4618db126 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 f3478e611c..6cfc1e11ea 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 13e9aec1ec..dc16cd967e 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 9eb1f6c5d9..38c63fc598 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/69] 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 8ce1a336b4..58ef53e568 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 dc16cd967e..b4ecff5e83 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 38c63fc598..da0a0a4301 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 b8670926c9..c8b7b7abaa 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/69] 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 da0a0a4301..038f83c746 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 c8b7b7abaa..6674654996 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/69] 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 a78a99d7f9..6584ff8333 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 0000000000..7068ec835a --- /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 038f83c746..b35d1b6797 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 6674654996..bf6a01d1d5 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/69] 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 0fee105fd0..1e04f69790 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 7068ec835a..7d400393d0 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 58ef53e568..3ca708a715 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 6cfc1e11ea..43231e38b8 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 b35d1b6797..11754ef7bc 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 bf6a01d1d5..11fcccf1b9 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/69] 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 1e04f69790..9fe25fbad8 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 11fcccf1b9..1c7aadbf93 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/69] 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 43231e38b8..beb0ee20fc 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/69] 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 7d400393d0..377eaf7f60 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 1c7aadbf93..3477452119 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/69] 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 377eaf7f60..34e7bfff21 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/69] 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 34e7bfff21..62ee83d21e 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 3ca708a715..e67d012b2b 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 beb0ee20fc..482f669a28 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 11754ef7bc..ded9bb8c8c 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 3477452119..5096f05260 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/69] 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 482f669a28..162c56ea13 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 5096f05260..638797d6f3 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/69] 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 9fe25fbad8..5707111f2a 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 62ee83d21e..933c24e09f 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/69] 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 933c24e09f..bf904c8a73 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 162c56ea13..4a6d50cea1 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 ded9bb8c8c..6e9097e20c 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 638797d6f3..c9a4962ad4 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/69] 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 6e9097e20c..9a8d9b16b0 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/69] 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 bf904c8a73..cdd7255175 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/69] 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 c9a4962ad4..f9ecc975a2 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/69] 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 cdd7255175..e01a3481a2 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 4a6d50cea1..e02456f20b 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 b4ecff5e83..41f94352df 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 9a8d9b16b0..4725cc6735 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 f9ecc975a2..50031022ba 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/69] 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 50031022ba..ea3b267c64 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/69] 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 e01a3481a2..ddd46bd527 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/69] 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 5707111f2a..e51c856828 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 d1dd818721..08acfeccba 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 ea3b267c64..999150741a 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/69] 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 e67d012b2b..c15e3ea04a 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/69] 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 08acfeccba..ce0b79d51e 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 ddd46bd527..0ecee3b0f6 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/69] 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 e51c856828..e3bc619210 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/69] 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 c15e3ea04a..ed45c9aa75 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 4725cc6735..51b1436ca5 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 999150741a..6edf91b01f 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/69] 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 e3bc619210..d6723c0c21 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/69] 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 0000000000..08e29f1038 --- /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 b9593b42d9..e5decacd84 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 51b1436ca5..fc6e432742 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/69] 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 08e29f1038..ae64de77dd 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 e5decacd84..f441853788 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 fc6e432742..10f07a88cd 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/69] 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 10f07a88cd..7874de3aa4 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/69] 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 f441853788..8b3bcca5e7 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 7874de3aa4..57426a08e3 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/69] 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 beed406991..98233b1d21 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 0ecee3b0f6..7a3f6ca0bc 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 ae64de77dd..7ad0a32233 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 8b3bcca5e7..b500031039 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 57426a08e3..1fef234841 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/69] 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 98233b1d21..bcde093e50 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 b500031039..881f6d2259 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/69] 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 7a3f6ca0bc..305c78f273 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 7ad0a32233..a2e20d0153 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/69] [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 a2e20d0153..a10ae5f508 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 881f6d2259..5a15c234a5 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 1fef234841..d8c235c200 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/69] 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 bcde093e50..ba94e4db0b 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 305c78f273..9492267e78 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/69] 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 ba94e4db0b..52114e0a0e 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 a10ae5f508..7a07a08484 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 5a15c234a5..8230555079 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 d8c235c200..b300be40a7 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/69] [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 52114e0a0e..466d8d902c 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 8230555079..6df8655f0a 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/69] 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 466d8d902c..61fbab80d4 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/69] [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 61fbab80d4..42236812b2 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/69] 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 b2a2b56bf6..a547532d5e 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/69] 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 9a7933f6d1..1df2ed5bcc 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 176e813ea5..342084474c 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 9492267e78..45c301f703 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 d4618db126..d665df4cef 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 fd4412e26d..eae6f0a8a2 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 6a8d50bf19..62917a5c8d 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 49845211ec..3a81b5142b 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 a547532d5e..a98c5538bb 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 41f94352df..658dab5d88 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 b300be40a7..f1956aa2ae 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 1ee6cf9487..83e6b8e211 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 9a33df7634..8e731f813e 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/69] [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 342084474c..af5eea2172 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 3a81b5142b..6e5c331b90 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/69] 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 45c301f703..76686de768 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 f1956aa2ae..96cb07544c 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/69] [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 96cb07544c..a4a51662cc 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/69] 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 42236812b2..beed406991 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 76686de768..395f3f6beb 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 7a07a08484..0000000000 --- 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 6df8655f0a..b9593b42d9 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 a4a51662cc..4cb8a9ea4e 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/69] [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 4cb8a9ea4e..3afb9c939e 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/69] 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 a98c5538bb..abc244756f 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/69] 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 abc244756f..8cb10bcc48 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/69] 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 62917a5c8d..5a52003dc3 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/69] 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 7534b39a74..2a7cb1c449 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/69] 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 beed406991..29a15e2fea 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 b9593b42d9..e551d1c49d 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/69] [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 29a15e2fea..174caa5949 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/69] 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 e551d1c49d..b3e704d9f4 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 3afb9c939e..023aa7b5d7 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 83e6b8e211..b30429fc54 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 8e731f813e..9dd93f3ce1 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/69] 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 8cb10bcc48..f36bb46401 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/69] 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 321fbdb171..ebe91da63f 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 395f3f6beb..3f578ebe02 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 023aa7b5d7..9c3960d3a6 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/69] 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 174caa5949..beed406991 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 b3e704d9f4..b9593b42d9 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/69] 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 1df2ed5bcc..7353b71717 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/69] 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 9c3960d3a6..c51543c6c8 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/69] 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 3f578ebe02..2d2763b406 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 0000000000..4ce7d17899 --- /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 b9593b42d9..38caf5b740 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 c51543c6c8..15f60a7647 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/69] 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 2d2763b406..548cc585f7 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 15f60a7647..f71d5ef8c8 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/69] 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 548cc585f7..da89f4b338 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 f71d5ef8c8..6ceaa33c54 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