diff --git a/docs/envvars.rst b/docs/envvars.rst index 46b70bbe46..9a7933f6d1 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -178,10 +178,13 @@ 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 -backend-selection overview. +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 +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 @@ -213,6 +216,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. **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 :Type: ``int`` (0 or 1) diff --git a/qa/L0_pytorch_unittest/test.sh b/qa/L0_pytorch_unittest/test.sh index b3b6ccacac..fda1b91d68 100644 --- a/qa/L0_pytorch_unittest/test.sh +++ b/qa/L0_pytorch_unittest/test.sh @@ -64,6 +64,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_GDN2_TEST_REQUIRED=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_gdn2_attention.xml $TE_PATH/tests/pytorch/attention/test_gdn2_attention.py || test_fail "test_gdn2_attention.py" NVTE_GDP_TEST_REQUIRED=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_gdp_attention.xml $TE_PATH/tests/pytorch/attention/test_gdp_attention.py || test_fail "test_gdp_attention.py" diff --git a/tests/pytorch/attention/run_attention_with_cp.py b/tests/pytorch/attention/run_attention_with_cp.py index 0d2a142dbc..176e813ea5 100644 --- a/tests/pytorch/attention/run_attention_with_cp.py +++ b/tests/pytorch/attention/run_attention_with_cp.py @@ -20,6 +20,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 ( @@ -276,6 +277,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", @@ -596,6 +605,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_attention_with_cp.py b/tests/pytorch/attention/test_attention_with_cp.py index e4b5ad86ed..321fbdb171 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 @@ -747,6 +757,99 @@ 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", "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 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: + pytest.skip(reason) + + config = model_configs_frost_attn[model] + config.context_parallel = True + config.cp_comm_type = cp_comm_type + + # a2a requires num_heads and num_gqa_groups divisible by the a2a subgroup size; every config + # here satisfies that, but assert rather than rely on it staying true. + if cp_comm_type == "a2a+p2p": + assert config.num_heads % 2 == 0 and config.num_gqa_groups % 2 == 0, ( + f"cp_comm_type=a2a+p2p needs num_heads ({config.num_heads}) and num_gqa_groups" + f" ({config.num_gqa_groups}) divisible by the a2a subgroup size" + ) + + pool = cp_pool(4 if cp_comm_type == "a2a+p2p" else 2) + + _submit( + pool, + 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.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_device_compute_capability() < (9, 0), reason="FusedAttention THD requires sm90+." ) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index beed406991..42236812b2 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -705,3 +705,195 @@ 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 new file mode 100644 index 0000000000..9492267e78 --- /dev/null +++ b/tests/pytorch/attention/test_frost_attention.py @@ -0,0 +1,522 @@ +# 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 +float64 reference instead. + +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 +import os + +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() +# 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) +# 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. +_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, 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, + 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) + s = (qq @ kk.transpose(-1, -2)) * scale + 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) + + +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, window) + lossy, lossy_lse = _reference( + q32.to(dtype).double(), k32.to(dtype).double(), v32.to(dtype).double(), scale, mask, window + ) + return ( + (exact - lossy).abs().max().item(), + (exact_lse - lossy_lse).abs().max().item(), + exact, + exact_lse, + ) + + +@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]) +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 = 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) + 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.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. + 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 + + +@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"]) +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 + 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, + ) + + # 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) + # 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) + + 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,) + + +@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"]) +@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 + 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) + 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) + 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 + ) + + 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, window) + 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.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" % ( + name, + err, + want.abs().max().item(), + ) + + +@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, + ) + + 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"), + # 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"), + # 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 + assert reason, "a decline must explain itself" + + +@requires_frost +@pytest.mark.parametrize( + "cp_comm_type,window,expect_frost", + [ + ("all_gather", (128, 0), True), + ("a2a", (128, 0), True), + ("p2p", (128, 0), False), + ("a2a+p2p", (128, 0), False), + ("p2p", (-1, 0), True), + ("p2p", (-1, -1), True), + ], +) +def test_frost_sliding_window_selection_by_cp_comm_type(cp_comm_type, window, expect_frost): + """Which context-parallel paths may serve a sliding window. + + all_gather and a2a each see a contiguous KV range, so the window applies unchanged. The p2p + ring shards KV across steps, so a bound measured against the full sequence does not survive + the per-step tiles -- the same rule FusedAttention carries. The cases without a real window + must still select FROST, since the decline has to key on the window and not on p2p itself. + """ + from transformer_engine.pytorch.attention.dot_product_attention.utils import ( + AttentionParams, + get_attention_backend, + ) + + params = AttentionParams( + qkv_dtype=torch.bfloat16, + qkv_layout="bshd_bshd_bshd", + batch_size=2, + num_heads=8, + num_gqa_groups=4, + max_seqlen_q=4096, + max_seqlen_kv=4096, + head_dim_qk=512, + head_dim_v=512, + attn_mask_type="causal", + window_size=window, + context_parallel=True, + cp_comm_type=cp_comm_type, + is_training=True, + ) + 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), + ) + + +@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 ( + 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. 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)) + + +@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. + + 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 + + +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 + + 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.""" + + 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/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/tests/pytorch/test_torch_compile.py b/tests/pytorch/test_torch_compile.py index eae6f0a8a2..fd4412e26d 100644 --- a/tests/pytorch/test_torch_compile.py +++ b/tests/pytorch/test_torch_compile.py @@ -1294,6 +1294,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 62917a5c8d..6a8d50bf19 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 f90fc46c03..49845211ec 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -2369,6 +2369,210 @@ 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 `_: 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 3a153a3f99..a547532d5e 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1572,6 +1572,259 @@ 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 @@ -1618,6 +1871,7 @@ def forward( use_flash_attn_4, fp8_output, layer_number, + use_frost_attention, ): # pylint: disable=missing-function-docstring @@ -1962,7 +2216,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_, @@ -2033,7 +2289,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 % 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: ( out_per_step[i % 2], softmax_lse_per_step[i % 2], @@ -2062,7 +2328,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 % 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: ( out_per_step[i % 2], softmax_lse_per_step[i % 2], @@ -2091,7 +2367,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 % 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: ( out_per_step[i % 2], softmax_lse_per_step[i % 2], @@ -2121,7 +2407,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 % 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: ( out_per_step[i % 2], softmax_lse_per_step[i % 2], @@ -2386,6 +2680,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 @@ -2747,7 +3042,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, @@ -2819,7 +3122,14 @@ 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, + deterministic=ctx.deterministic, + ) + elif ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( *fused_attn_inputs, *prepare_outputs, section ) @@ -2832,7 +3142,14 @@ 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, + deterministic=ctx.deterministic, + ) + elif ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( *fused_attn_inputs, *prepare_outputs, section ) @@ -2845,7 +3162,14 @@ 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, + deterministic=ctx.deterministic, + ) + elif ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( *fused_attn_inputs, *prepare_outputs, section ) @@ -2858,7 +3182,14 @@ 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, + deterministic=ctx.deterministic, + ) + elif ctx.use_fused_attention: dq_, dk_, dv_, dbias_ = cp_p2p_bwd_fused_attn( *fused_attn_inputs, *prepare_outputs, section ) @@ -3199,6 +3530,7 @@ def backward(ctx, dout, *_args): None, None, None, + None, # use_frost_attention ) @@ -3288,6 +3620,7 @@ def forward( quantizers, fp8_output, load_balancing_strategy, + use_frost_attention, ): # pylint: disable=missing-function-docstring nvtx_range_push("transformer_engine.AttnFuncWithCPAndKVAllGather.forward") @@ -3321,11 +3654,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: @@ -3694,7 +4028,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] @@ -3964,6 +4308,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 @@ -4234,7 +4579,24 @@ 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, + deterministic=ctx.deterministic, + ) + 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] @@ -4561,6 +4923,7 @@ def backward(ctx, dout, *_args): None, None, None, + None, # use_frost_attention ) @@ -4605,6 +4968,7 @@ def forward( softmax_type, softmax_offset, fp8_output, + use_frost_attention, ): # pylint: disable=missing-function-docstring nvtx_range_push("transformer_engine.AttnFuncWithCPAndQKVOA2A.forward") @@ -4634,10 +4998,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, ( @@ -4790,7 +5156,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, 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 fp8: if fp8_recipe.mxfp8(): q_fp8, k_fp8, v_fp8, qkv_layout, qkv_scale_inv_format = combine_and_quantize( @@ -5026,6 +5404,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 @@ -5177,7 +5559,25 @@ 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: + # 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: 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 @@ -5410,6 +5810,7 @@ def backward(ctx, dout, *_args): None, d_softmax_offset, None, + None, # use_frost_attention ) @@ -5549,6 +5950,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, @@ -5686,8 +6088,11 @@ 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 + 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" @@ -5753,6 +6158,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": @@ -5768,6 +6174,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": @@ -5784,6 +6191,7 @@ 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/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py new file mode 100644 index 0000000000..7a07a08484 --- /dev/null +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -0,0 +1,284 @@ +# 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/dot_product_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py index 658dab5d88..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 @@ -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, } @@ -156,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( @@ -166,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, } ) @@ -996,6 +1000,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, @@ -2844,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 @@ -2858,6 +2875,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 +2885,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 +2904,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 +2914,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 +3090,27 @@ 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/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index b9593b42d9..6df8655f0a 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,42 @@ """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 + # 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) 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): @@ -173,6 +163,27 @@ 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: @@ -194,40 +205,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 @@ -263,19 +247,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.""" - 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, exclude_plan_tokens=_FROST_PLAN_TOKENS) + return workspace_size def _execute_cudnn_graph( @@ -285,20 +269,7 @@ 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) def _cudnn_score_mod_fwd_cache_key( @@ -313,6 +284,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. @@ -335,6 +307,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, ) @@ -353,6 +328,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) @@ -375,6 +351,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, ) @@ -390,8 +367,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) @@ -403,6 +387,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, @@ -410,8 +395,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) @@ -448,6 +432,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 = ( @@ -463,12 +448,15 @@ def _get_cudnn_score_mod_fwd_graph( output_layer, stats, ) - 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 @@ -488,8 +476,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) @@ -522,8 +513,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, ) @@ -564,6 +554,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 = ( @@ -582,12 +573,13 @@ def _get_cudnn_score_mod_bwd_graph( score_mod_bprop_tensors, deterministic, ) - 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 new file mode 100644 index 0000000000..b300be40a7 --- /dev/null +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -0,0 +1,649 @@ +# 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. + +**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 + +import contextlib +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 + +from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph + +__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 = 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 = 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. +_HEAD_DIM_MULTIPLE = 8 + +_cudnn = None +_availability: Optional[Tuple[bool, str]] = None +_PLAN_CACHE: dict = {} +_HANDLES = cudnn_pygraph._handles # pylint: disable=protected-access + + +def _import_cudnn(enable_frost_engines: bool = True): + """Import cuDNN Frontend, registering the FROST engines unless told not to. + + 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=enable_frost_engines) + 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 _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. + + Distribution metadata first, matching the sibling check in fused_mla_q_uproj.py, with the + module attribute as a fallback so a source or vendored install is not misreported as absent. + The raw string is returned separately so callers can tell "not installed" from "installed but + unparseable"; those warrant different answers, and conflating them declines valid installs. + """ + 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: + return PkgVersion(raw), raw + except InvalidVersion: + return None, raw + + +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 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: + 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. + _import_cudnn(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. + frontend, frontend_raw = _pkg_version("nvidia-cudnn-frontend", _cudnn) + if frontend is not None and frontend < _MIN_CUDNN_FRONTEND: + return _no( + 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(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( + 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, "") + 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. +_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_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( + f"FROST attention supports attn_mask_type in {str(_SUPPORTED_MASKS)}; got" + f" {attn_mask_type!r}" + ) + try: + window = _NO_WINDOW if window_size is None else tuple(window_size) + except TypeError: + # Raised as NotImplementedError so the selector declines instead of propagating out of + # backend selection, which is the only thing is_frost_attention_supported catches. + raise NotImplementedError( + f"window_size must be a (left, right) pair; got {window_size!r}" + ) from None + if len(window) != 2: + raise NotImplementedError(f"window_size must be a (left, right) pair; got {window!r}") + 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(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(f"FROST attention does not support a right window {window!r}") + return attn_mask_type, window + + +def _mask_options(cudnn, spec): + """cuDNN sdpa kwargs for a (mask type, window) spec: a diagonal alignment plus a band.""" + attn_mask_type, window = spec + return 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. + """ + if 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, f"FROST path 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}" + ), + ) + 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" + try: + _mask_spec(attn_mask_type, window_size) + except NotImplementedError as exc: + return False, str(exc) + ok, reason = is_frost_attention_available() + if not ok: + return False, reason + 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( + 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 _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(f"{name} must be 4D [b, h, s, d]; got {tuple(t.shape)}") + if t.stride(3) != 1: + raise ValueError( + f"{name} must have a contiguous head dimension; got shape {tuple(t.shape)} stride" + f" {tuple(t.stride())}" + ) + + +def _check_dtype(name: str, t: torch.Tensor, expected: torch.dtype) -> None: + """Require a tensor to carry the dtype its graph node was declared with. + + 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(f"{name} must be {expected} to match q; got {t.dtype}") + + +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(f"k and v must have the same shape; got {k.shape} and {v.shape}") + if k.stride() != v.stride(): + raise ValueError( + f"k and v must have the same layout; got strides {tuple(k.stride())} and" + f" {tuple(v.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. 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"Wanted the FROST {what} engine." + " 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})." + ) + + 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: + """Build (and JIT-compile) a forward graph. Expensive; always reached through the cache.""" + cudnn = _import_cudnn() + # 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( + 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)) + 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_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( + cudnn.data_type.FLOAT + ) + plan = _select_frost_plan(graph, _FROST_FWD_PLAN_TOKEN, "forward") + return { + "graph": graph, + "handles": (tq, tk, tv, tout, tlse), + "workspace": max(graph.get_workspace_size(), 1), + "plan": plan, + } + + +def _build_bwd(key) -> 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, deterministic = key + io_dt = _cudnn_dtype(dtype) + shq, shkv = [b, hq, sq, d], [b, hkv, skv, d] + + 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. + 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, + use_deterministic_algorithm=deterministic, + **_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)) + 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: building dominates executing even once the JIT is + cached, so this is required rather than an optimisation.""" + cache_key = (kind,) + key + entry = _PLAN_CACHE.get(cache_key) + if entry is None: + # Build under the device the key names, not merely with 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 + + +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. + q.device.type, + q.device.index, + 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()), + # 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), + ) + + +def frost_attn_fwd( + q: torch.Tensor, + k: torch.Tensor, + 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. + + 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(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]}" + ) + + 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"] + + 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, handle=_handle_for(q.device) + ) + 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", + 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)): + _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]}" + ) + 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 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]): + raise ValueError( + 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) + 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"] + + 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, + 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 fbbd899aa2..1ee6cf9487 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 @@ -613,6 +615,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: @@ -1859,6 +1862,170 @@ 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 @@ -1877,7 +2044,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]), @@ -1887,6 +2054,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), ) # Prefer FA2 for THD training with dropout on SM100/103, where FusedAttention has a known @@ -1916,13 +2087,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) @@ -1933,6 +2111,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, )