diff --git a/docs/envvars.rst b/docs/envvars.rst index 46b70bbe46a..1df2ed5bcca 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -178,10 +178,12 @@ Then it applies a performance-based preference order among the remaining eligibl In PyTorch, the broad preference order is ``FlashAttention > FusedAttention > UnfusedDotProductAttention`` on supported pre-Hopper GPUs such as Ampere/Ada, and ``FusedAttention > FlashAttention > UnfusedDotProductAttention`` on Hopper and newer GPUs, -including Blackwell. In JAX, Transformer Engine uses cuDNN fused attention when -``NVTE_FUSED_ATTN=1`` and an eligible cuDNN kernel is available; otherwise it falls back to the -JAX-native implementation. See :doc:`examples/attention/attention` for a longer -backend-selection overview. +including Blackwell. On Blackwell SM100/SM103, FusedAttention has an extra sub-backend, FROST, +which is selected only for symmetric ``head_dim`` in (256, 512] and only when the cuDNN +sub-backends decline; it does not change the order above. In JAX, Transformer Engine uses cuDNN +fused attention when ``NVTE_FUSED_ATTN=1`` and an eligible cuDNN kernel is available; otherwise +it falls back to the JAX-native implementation. See :doc:`examples/attention/attention` for a +longer backend-selection overview. .. envvar:: NVTE_FLASH_ATTN @@ -213,6 +215,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 the FROST sub-backend of FusedAttention for DotProductAttention. FROST wraps the cuDNN FROST CuTe-DSL SDPA kernels through the cuDNN Frontend python API, rather than the C++ fused-attention path the other sub-backends use. **It is experimental and subject to change**, as the underlying cuDNN FROST engines are. When set to ``0``, FROST will not be used. It is selected only where the cuDNN sub-backends decline and is the only released backend serving symmetric ``head_dim`` in (256, 512] together with context parallelism. It is limited to SM100/SM103 with BF16/FP16 inputs, a ``head_dim`` that is a multiple of 8, and ``nvidia-cudnn-frontend>=1.29.0`` and ``nvidia-cutlass-dsl>=4.7.0`` installed. It declines FP8, ``thd`` layouts, dropout, attention bias, KV caching, ``max_logit``, CUDA graph capture, and deterministic execution, the last because cuDNN offers no deterministic backward for these kernels. + .. envvar:: NVTE_UNFUSED_ATTN :Type: ``int`` (0 or 1) diff --git a/qa/L0_pytorch_unittest/test.sh b/qa/L0_pytorch_unittest/test.sh index b3b6ccacac7..fda1b91d68e 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 0d2a142dbc1..af5eea21722 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,15 @@ 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": + # FROST is a sub-backend of FusedAttention, so NVTE_FUSED_ATTN has to stay on. Flash is + # left off; nothing else serves head_dim > 256, so the selector reaches FROST on its own. + os.environ["NVTE_FUSED_ATTN"] = "1" + os.environ["NVTE_FROST_ATTN"] = "1" + if model in model_configs_frost_attn: + config = copy.deepcopy(model_configs_frost_attn[model]) + else: + assert False, f"{model=} is not a known FrostAttention CP config!" assert config.attn_mask_type in [ "causal", "no_mask", @@ -596,6 +606,22 @@ def run_dpa_with_cp( pad_between_seqs=pad_between_seqs, fp8_output=fp8_mha, ) + if kernel_backend == "FrostAttention": + # Assert the sub-backend actually used, not just the one requested. FROST is + # currently the only selectable backend for these configs -- flash is env-gated off + # and CP disables unfused -- so a silent substitution is impossible today and this + # would pass by construction. It is here so it stops passing if that stops being + # true, rather than quietly testing some other kernel. + from transformer_engine.pytorch.attention.dot_product_attention.dot_product_attention import ( # pylint: disable=import-outside-toplevel + _attention_backends, + ) + + # pylint: disable-next=import-outside-toplevel + from transformer_engine.pytorch.cpp_extensions.fused_attn import FusedAttnBackend + + assert ( + _attention_backends["fused_attention_backend"] == FusedAttnBackend.FROST + ), "expected the FROST sub-backend to be selected, got %s" % (_attention_backends,) if config.return_max_logit: out_, max_logit_ = out_ if is_training: diff --git a/tests/pytorch/attention/test_attention_with_cp.py b/tests/pytorch/attention/test_attention_with_cp.py index e4b5ad86edb..dbd237b4b24 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: 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,118 @@ def test_cp_with_fused_attention( ) +def _frost_availability(): + """Why FrostAttention cannot run here, or None if it can. + + The backend needs cuDNN Frontend >= 1.29.0 and, less obviously, + nvidia-cutlass-dsl >= 4.7.0: cudnn-frontend only declares >= 4.6.2, and below the FROST floor + every FROST engine silently declines and ordinary backend plans are returned with no error. + Reporting the reason as a skip keeps that distinguishable from a real failure. + """ + if get_device_compute_capability() not in ((10, 0), (10, 3)): + return "FrostAttention requires SM100/SM103 (the cuDNN d512 backward is Blackwell-only)." + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + is_frost_attention_available, + ) + + ok, reason = is_frost_attention_available() + return None if ok else reason + + +@pytest.mark.parametrize("model", model_configs_frost_attn.keys()) +@pytest.mark.parametrize("qkv_format", ["bshd", "sbhd"]) +@pytest.mark.parametrize("cp_comm_type", ["p2p", "all_gather", "a2a"]) +def test_cp_with_frost_attention(cp_pool, model, qkv_format, cp_comm_type): + """Context parallelism at head_dim 512, which no other backend serves. + + thd is excluded because the backend declines it: it needs varlen support that is not + implemented. a2a+p2p is covered separately below, since it needs four ranks. + """ + 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 + + _submit( + cp_pool(2), + 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, + ) + + +def test_cp_with_frost_attention_a2a_p2p(cp_pool): + """One case, because a2a+p2p composes rather than adding a path. + + It dispatches to the same AttnFuncWithCPAndKVP2P as plain p2p, with an a2a communication stage + on either side of the ring, and it needs four ranks rather than two. One case says the + composition works; six would pay four-rank prices to re-cover p2p. + """ + reason = _frost_availability() + if reason is not None: + pytest.skip(reason) + + # The shortest model, since this is the only four-rank case and so the only one that would + # take all of a2a+p2p down with it if it timed out. + config = model_configs_frost_attn["cp_hd512_2"] + config.context_parallel = True + config.cp_comm_type = "a2a+p2p" + # The a2a stage shards heads across its subgroup, so both counts must divide by it. + assert config.num_heads % 2 == 0 and config.num_gqa_groups % 2 == 0, ( + f"a2a+p2p needs num_heads ({config.num_heads}) and num_gqa_groups" + f" ({config.num_gqa_groups}) divisible by the a2a subgroup size" + ) + + _submit( + cp_pool(4), + dtype="bf16", + model="cp_hd512_2", + qkv_format="bshd", + kernel_backend="FrostAttention", + cp_comm_type="a2a+p2p", + is_training=True, + log_level=pytest_logging_level, + ) + + +@pytest.mark.parametrize("cp_comm_type", ["p2p", "all_gather", "a2a"]) +def test_cp_with_frost_attention_fp16(cp_pool, cp_comm_type): + """One fp16 arm per comm type, since the matrix above is bf16 throughout. + + 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 beed4069917..39431ce3462 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -12,6 +12,7 @@ from transformer_engine.pytorch import DotProductAttention, is_bf16_available from transformer_engine.pytorch.attention.dot_product_attention import _attention_backends import transformer_engine.pytorch.attention.dot_product_attention.flex_attention as flex_attention +from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph from transformer_engine.pytorch.utils import get_device_compute_capability _current_file = pathlib.Path(__file__).resolve() @@ -22,6 +23,23 @@ get_available_attention_backends, ) + +def _flex_availability(): + """Why the cuDNN score_mod path cannot run here, or None if it can.""" + if not torch.cuda.is_available(): + return "no CUDA device" + try: + flex_attention._import_cudnn_frontend() + except ImportError as exc: + return "nvidia-cudnn-frontend not importable: %s" % exc + return None + + +_SKIP = _flex_availability() +# Mirrors NVTE_GDN_TEST_REQUIRED: the tests that reach cuDNN skip silently without the frontend. +if os.getenv("NVTE_FLEX_TEST_REQUIRED", "0") == "1" and _SKIP is not None: + raise RuntimeError("NVTE_FLEX_TEST_REQUIRED=1, but flex attention is unavailable: %s" % _SKIP) + param_types = [torch.float16] if torch.cuda.is_available() and is_bf16_available(): param_types.append(torch.bfloat16) @@ -705,3 +723,123 @@ def test_dot_product_attention_score_mod(dtype, qkv_format, score_mod_case, scal torch.testing.assert_close(q.grad, q_ref.grad, **tols) torch.testing.assert_close(k.grad, k_ref.grad, **tols) torch.testing.assert_close(v.grad, v_ref.grad, **tols) + + +def test_flex_bars_the_frost_engines(): + """flex must tell cuDNN not to use a FROST engine, not merely decline to ask for them. + + The switch that offers those engines is process-wide, so another caller in the process, or a + user setting CUDNN_FRONTEND_ENABLE_FROST_ENGINES, puts them ahead of the backend engines for + these graphs too. They accept a score_mod graph, pass check_support, build, and then compute + without the callback. + + The bar is the shared helper's default rather than something flex asks for, so this checks + flex's own path reaches it: a caller that forgot to opt in is the failure being prevented. + + No GPU: this checks the instruction is passed, not what cuDNN does with it. + """ + barred = [] + + class FakeGraph: + """Records the engines flex bars, and stops at the first call it cannot serve.""" + + def validate(self): + pass + + def build_operation_graph(self): + pass + + def create_execution_plans(self, _heuristics): + pass + + def deselect_engines(self, names): + barred.extend(names) + + def check_support(self): + pass + + def build_plans(self, _policy): + pass + + def get_workspace_size(self): + return 4096 + + try: + flex_attention._import_cudnn_frontend() + except ImportError: + pytest.skip("cuDNN frontend Python package is required for score_mod attention.") + + assert flex_attention._finalize_cudnn_graph(FakeGraph()) == 4096 + assert barred, "flex did not ask cuDNN to exclude any engine" + assert set(cudnn_pygraph.FROST_PLAN_TOKENS) <= set(barred), barred + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required.") +def test_frost_switch_does_not_change_what_flex_computes(): + """Enabling the FROST engines must not change flex's output. + + This is the property the silent drop violates: with the engines on, an unpinned build selects + a FROST plan at every head dim measured on B200, and that plan returns plain attention with + the score_mod discarded. Comparing flex against itself across the switch needs no reference + and no knowledge of which plan ran; if the two differ, a different kernel answered. + """ + try: + flex_attention._import_cudnn_frontend() + except ImportError: + pytest.skip("cuDNN frontend Python package is required for score_mod attention.") + # Without this the test is vacuous: if the engines are absent, decline on arch, or sit below + # their version floors, both runs get a backend plan and agree no matter what flex does. The + # availability probe covers all three, where an arch check alone would miss the version ones. + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + is_frost_attention_available, + ) + + frost_ok, frost_reason = is_frost_attention_available() + if not frost_ok: + pytest.skip("the FROST engines must be reachable to test anything: %s" % frost_reason) + + env = "CUDNN_FRONTEND_ENABLE_FROST_ENGINES" + saved = os.environ.get(env) + torch.manual_seed(0) + b, h, s, d = 2, 4, 512, 64 + dtype = torch.bfloat16 if is_bf16_available() else torch.float16 + q, k, v = (torch.randn(b, s, h, d, device="cuda", dtype=dtype) for _ in range(3)) + + def bias_score_mod(score_mod_graph, score_tensor, _tensors): + """score += (row - col). Self-contained, and large enough that dropping it is obvious.""" + cudnn = flex_attention._import_cudnn_frontend() + row = score_mod_graph.gen_index(input=score_tensor, axis=2) + row.set_data_type(cudnn.data_type.INT32) + col = score_mod_graph.gen_index(input=score_tensor, axis=3) + col.set_data_type(cudnn.data_type.INT32) + bias = score_mod_graph.sub(a=row, b=col, compute_data_type=cudnn.data_type.FLOAT) + bias.set_data_type(cudnn.data_type.FLOAT) + return score_mod_graph.add(a=score_tensor, b=bias, compute_data_type=cudnn.data_type.FLOAT) + + def run(): + flex_attention._cudnn_score_mod_graph_cache.clear() + return flex_attention.FusedAttentionWithScoreModFunc.apply( + False, q, k, v, "bshd", "bshd", d**-0.5, bias_score_mod, None, None, None, False + ) + + try: + os.environ.pop(env, None) + without = run() + os.environ[env] = "1" + with_engines = run() + finally: + flex_attention._cudnn_score_mod_graph_cache.clear() + if saved is None: + os.environ.pop(env, None) + else: + os.environ[env] = saved + + torch.testing.assert_close( + with_engines, + without, + msg=lambda m: ( + "flex computed something different with the FROST engines enabled, which means a" + " FROST plan answered and dropped the score_mod:\n" + + m + ), + ) diff --git a/tests/pytorch/attention/test_frost_attention.py b/tests/pytorch/attention/test_frost_attention.py new file mode 100644 index 00000000000..3f3a9023632 --- /dev/null +++ b/tests/pytorch/attention/test_frost_attention.py @@ -0,0 +1,597 @@ +# 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: these skip on any machine that cannot reach the backend. +if os.getenv("NVTE_FROST_TEST_REQUIRED", "0") == "1" and _SKIP is not None: + raise RuntimeError("NVTE_FROST_TEST_REQUIRED=1, but FrostAttention is unavailable: %s" % _SKIP) +# Per test, not a module-level pytestmark: the ONNX-export regression below runs on every GPU, so +# gating it on Blackwell would skip it exactly where its bug can still occur. +requires_frost = pytest.mark.skipif(_SKIP is not None, reason=str(_SKIP)) + +# head_dim 512 is the whole point of the backend; 320 checks the interior of the (256, 512] range +# 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 _shape_id(s): + return "b%d_hq%d_hkv%d_sq%d_skv%d_d%d" % s + + +def _fwd(q, k, v, mask, scale, window=None): + """The forward through the fused signature, which is the only entry point the backend has. + + Everything is bshd here, because the shim derives one qkv_format and makes the tensors + contiguous; that is exactly what the dispatcher hands it in production. + """ + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + fused_attn_fwd, + ) + + out, aux = fused_attn_fwd( + True, + q.shape[1], + k.shape[1], + None, + None, + q, + k, + v, + None, + None, + attn_scale=scale, + qkv_layout="bshd_bshd_bshd", + o_format="bshd", + attn_mask_type=mask, + window_size=(-1, -1) if window is None else window, + ) + return out, aux[0] + + +def _bwd(q, k, v, out, lse, dout, mask, scale, window=None): + """The backward through the fused signature. aux_ctx_tensors is what the forward returned.""" + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + fused_attn_bwd, + ) + + dq, dk, dv, _ = fused_attn_bwd( + q.shape[1], + k.shape[1], + None, + None, + q, + k, + v, + out, + dout, + None, + [lse, torch.empty(2, dtype=torch.int64, device=q.device)], + None, + attn_scale=scale, + qkv_layout="bshd_bshd_bshd", + o_format="bshd", + do_format="bshd", + dqkv_layout="bshd_bshd_bshd", + attn_mask_type=mask, + window_size=(-1, -1) if window is None else window, + ) + return dq, dk, dv + + +def _bhsd(t): + """A [b, h, s, d] view of a bshd tensor. The reference works in that order; the kernel does + not, since it takes TE's format and reorders the cuDNN descriptors instead.""" + return t.permute(0, 2, 1, 3) + + +def _reference(q, k, v, scale, mask, window=None): + """Attention in float64, computed independently of TE and of cuDNN. + + 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, + ) + + +# bf16 everywhere, fp16 on two shapes. What fp16 risks that bf16 does not is its narrower +# exponent range, and that surfaces in the backward, which runs both dtypes on every shape it +# covers. Crossing it with every forward shape pays for the same information twice. +_FWD_CASES = [(s, torch.bfloat16) for s in _SHAPES] + [(s, torch.float16) for s in _SHAPES[:2]] +_FWD_IDS = [ + "%s_%s" % (_shape_id(s), "bf16" if d is torch.bfloat16 else "fp16") for s, d in _FWD_CASES +] + + +@requires_frost +@pytest.mark.parametrize("shape,dtype", _FWD_CASES, ids=_FWD_IDS) +@pytest.mark.parametrize("mask", ["no_mask", "causal", "causal_bottom_right"]) +def test_frost_forward_matches_reference(shape, mask, dtype): + """Forward output and LSE against an independent float64 reference.""" + b, hq, hkv, sq, skv, d = shape + torch.manual_seed(0) + # Generate in fp32 so there is a true high-precision original to measure against, then cast + # for the kernel. bshd is what the backend takes now: it is never permuted, only described. + mk = lambda s_, h_, d_: torch.randn(b, s_, h_, d_, device="cuda") + q32, k32, v32 = mk(sq, hq, d), mk(skv, hkv, d), mk(skv, hkv, d) + q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) + scale = 1.0 / math.sqrt(d) + + out, lse = _fwd(q, k, v, mask, scale) + + floor_o, floor_l, ref_o, ref_lse = _floor( + _bhsd(q32), _bhsd(k32), _bhsd(v32), scale, mask, dtype + ) + err_o = (_bhsd(out).double() - ref_o).abs().max().item() + err_l = (lse.double() - ref_lse).abs().max().item() + + assert torch.isfinite(out).all(), "forward produced non-finite values" + # 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 +# (128, 0) is the ordinary case and (0, 0) the degenerate diagonal-only one, which is where an +# off-by-one in the band would show. A second ordinary width tests the same arithmetic again. +@pytest.mark.parametrize("window", [(128, 0), (0, 0)], ids=lambda w: "win%d" % w[0]) +@pytest.mark.parametrize("mask", ["causal", "causal_bottom_right", "no_mask"]) +@pytest.mark.parametrize("sq,skv", [(1024, 1024), (512, 1024)], ids=["square", "rect"]) +def test_frost_sliding_window_matches_reference(mask, window, sq, skv): + """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. + """ + # The rectangular case is the one that matters for alignment: top-left and bottom-right + # coincide when sq == skv, so a swapped alignment is invisible in square shapes. + b, hq, hkv, d = 2, 8, 4, 512 + dtype = torch.bfloat16 + torch.manual_seed(0) + mk = lambda s_, h_: torch.randn(b, s_, h_, d, device="cuda") + 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, _ = _fwd(q, k, v, mask, scale, window) + + floor_o, _, ref_o, _ = _floor(_bhsd(q32), _bhsd(k32), _bhsd(v32), scale, mask, dtype, window) + err = (_bhsd(out).double() - ref_o).abs().max().item() + 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, _ = _fwd(q, k, v, mask, scale) + assert not torch.equal(out, full), "window %s produced the same output as no window" % (window,) + + +@requires_frost +@pytest.mark.parametrize("shape", _SHAPES[:2] + _SHAPES[-1:], ids=_shape_id) +@pytest.mark.parametrize("mask", ["no_mask", "causal"]) +@pytest.mark.parametrize("window", [None, (128, 0)], ids=["nowin", "win128"]) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +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. + """ + b, hq, hkv, sq, skv, d = shape + torch.manual_seed(0) + mk = lambda s_, h_, d_: torch.randn(b, s_, h_, d_, device="cuda") + q32, k32, v32 = mk(sq, hq, d), mk(skv, hkv, d), mk(skv, hkv, d) + q, k, v = q32.to(dtype), k32.to(dtype), v32.to(dtype) + scale = 1.0 / math.sqrt(d) + + out, lse = _fwd(q, k, v, mask, scale, window) + dout = torch.randn_like(out) + dq, dk, dv = _bwd(q, k, v, out, lse, dout, mask, scale, window) + + # The reference works in [b, h, s, d], so it takes views and returns grads in that order. + qr = _bhsd(q32).detach().clone().requires_grad_(True) + kr = _bhsd(k32).detach().clone().requires_grad_(True) + vr = _bhsd(v32).detach().clone().requires_grad_(True) + ref_o, _ = _reference(qr, kr, vr, scale, mask, window) + ref_o.backward(_bhsd(dout).double()) + + for name, got, want in ( + ("dq", _bhsd(dq), qr.grad), + ("dk", _bhsd(dk), kr.grad), + ("dv", _bhsd(dv), vr.grad), + ): + assert torch.isfinite(got).all(), "%s has non-finite values" % name + assert got.shape == want.shape, "%s shape %s != %s" % (name, got.shape, want.shape) + err = (got.double() - want).abs().max().item() + # Gradients accumulate over the sequence, so scale the bar with skv rather than reusing + # the forward's floor. This is a sanity bound on systematic error, not a tight check. + assert err <= 0.05 * want.abs().max().item() + 1e-2, "%s max|err|=%.3e vs ref max %.3e" % ( + name, + err, + want.abs().max().item(), + ) + + +def _frost_params(**overrides): + """A FusedAttentionParams for a config FROST serves, with fields overridable by name.""" + from transformer_engine.pytorch.attention.dot_product_attention.utils import ( + FusedAttentionParams, + ) + from transformer_engine.pytorch.cpp_extensions.fused_attn import ( + AttnBiasType, + AttnMaskType, + QKVFormat, + QKVLayout, + SoftmaxType, + ) + from transformer_engine.pytorch.constants import TE_DType + + fields = dict( + head_dim_qk=512, + head_dim_v=512, + qkv_dtype=TE_DType[torch.bfloat16], + attn_mask_type=AttnMaskType["causal"], + bias_type=AttnBiasType["no_bias"], + softmax_type=SoftmaxType["vanilla"], + qkv_layout=QKVLayout["bshd_bshd_bshd"], + o_format=QKVFormat["bshd"], + window_size_left=-1, + window_size_right=-1, + bottom_right_diagonal=False, + ) + fields.update(overrides) + return FusedAttentionParams(**fields) + + +@requires_frost +def test_frost_declines_unsupported_configs(): + """The selector must decline what the kernels do not serve, rather than computing wrongly.""" + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + is_frost_attention_supported, + ) + from transformer_engine.pytorch.cpp_extensions.fused_attn import ( + AttnBiasType, + AttnMaskType, + FusedAttnBackend, + QKVFormat, + QKVLayout, + ) + from transformer_engine.pytorch.constants import TE_DType + + assert ( + is_frost_attention_supported(_frost_params())[0] == FusedAttnBackend.FROST + ), "the supported case must be accepted" + + for override, why in ( + (dict(head_dim_qk=256, head_dim_v=256), "head_dim at the exclusive lower bound"), + (dict(head_dim_v=256), "head_dim_v below the range"), + # The forward serves an asymmetric pair and the backward does not, so the selector + # declines it rather than accepting a config whose backward cannot build. + (dict(head_dim_v=320), "asymmetric head_dim"), + (dict(qkv_dtype=TE_DType[torch.float32]), "fp32"), + (dict(dropout=0.1), "dropout"), + (dict(bias_type=AttnBiasType["post_scale_bias"]), "attention bias"), + (dict(attn_mask_type=AttnMaskType["padding_causal"]), "padding mask"), + (dict(qkv_layout=QKVLayout["thd_thd_thd"], o_format=QKVFormat["thd"]), "thd layout"), + (dict(o_format=QKVFormat["sbhd"]), "an output format that differs from the input"), + (dict(num_pages_k=4, num_pages_v=4), "paged KV"), + (dict(return_max_logit=True), "max_logit"), + (dict(cuda_graph=True), "CUDA graph capture"), + (dict(deterministic=True, is_training=True), "a deterministic backward"), + # window_size reaches _mask_spec through the selector, so its validation is part of the + # selector contract rather than an internal detail. + (dict(window_size_right=5), "a right window past the diagonal"), + (dict(window_size_left=-2, window_size_right=0), "a left window below -1"), + # A right-bounded window on a non-causal mask takes its anchor only from + # bottom_right_diagonal; the all-gather ring measures its window bottom-right. Those + # differ exactly when the lengths do. + ( + dict( + attn_mask_type=AttnMaskType["no_mask"], + window_size_left=128, + window_size_right=0, + max_seqlen_q=512, + max_seqlen_kv=640, + ), + "an ambiguous diagonal anchor", + ), + # The engine pads head_dim to a multiple of 8, so an in-range but unpadded dim has to be + # declined here rather than failing later at plan selection. + (dict(head_dim_qk=260, head_dim_v=260), "head_dim not a multiple of 8"), + # Each dim is checked on its own, so v has to be covered as well as q. + (dict(head_dim_v=260), "head_dim_v not a multiple of 8"), + ): + backend, reason = is_frost_attention_supported(_frost_params(**override)) + assert backend == FusedAttnBackend.No_Backend, "%s must be declined" % why + assert reason, "a decline must explain itself" + + +@requires_frost +def test_frost_serves_an_unambiguous_window_on_a_non_causal_mask(): + """The anchor is only ambiguous when the q and kv lengths differ; equal lengths must serve.""" + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + is_frost_attention_supported, + ) + from transformer_engine.pytorch.cpp_extensions.fused_attn import AttnMaskType, FusedAttnBackend + + params = _frost_params( + attn_mask_type=AttnMaskType["no_mask"], + window_size_left=128, + window_size_right=0, + max_seqlen_q=4096, + max_seqlen_kv=4096, + ) + assert is_frost_attention_supported(params)[0] == FusedAttnBackend.FROST + + +@requires_frost +def test_frost_mask_spec_rejects_malformed_windows(): + """_mask_spec is the only validation between a caller-supplied window and a built band.""" + from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import ( + _mask_spec, + ) + + for window, why in ( + ((128,), "a malformed window pair"), + (7, "a non-iterable window"), + ((-1, 5), "a right window past the diagonal"), + ((-2, 0), "a left window below -1"), + ): + with pytest.raises(NotImplementedError): + _mask_spec("causal", window), why + + +@requires_frost +def test_frost_rejects_mismatched_kv(): + """v must index the same KV positions as k. head_dim is free; the rest is not.""" + b, h, s, d = 2, 4, 512, 512 + dtype = torch.bfloat16 + mk = lambda hh: torch.randn(b, s, hh, d, device="cuda", dtype=dtype) + q, k = mk(h), mk(h) + + with pytest.raises(ValueError, match="batch, heads and seqlen"): + _fwd(q, k, mk(h * 2), "no_mask", 1.0) + with pytest.raises(ValueError, match="match q"): + _fwd(q, k, k.to(torch.float32), "no_mask", 1.0) + + +def test_frost_engines_are_enabled_even_if_cudnn_was_imported_without_them(): + """Enabling the FROST engines must not depend on who imported cuDNN first. + + is_frost_attention_available imports cuDNN WITHOUT the engines, because enabling them reorders + plan selection for every cuDNN consumer in the process and the checks after it may still + decline. So the enabling cannot sit inside the "already imported?" memo: a process that + probed availability first would otherwise leave FROST with a cuDNN that offers it no engine. + That surfaces far from its cause, as "no cuDNN engine matching 'sdpa_fwd_prefill_sm100' was + offered" on the first head_dim 512 forward, with a hint pointing at package versions that are + in fact fine. + + No GPU and no real cuDNN: a stub stands in for the package, because what is under test is the + order-dependence of our own wrapper. It also has to run in-process with the globals reset, + 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, + frost_attention, + ) + + # The import memo and the switch live in the shared module; frost's wrapper only supplies the + # default. Drive it through the wrapper, which is the real entry, and assert on the owner. + env = "CUDNN_FRONTEND_ENABLE_FROST_ENGINES" + saved = ( + 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) + + # The availability probe first, which must not enable anything. + frost_attention._import_cudnn_frontend(enable_frost_engines=False) + assert env not in os.environ, "the non-FROST caller must not set the switch" + assert not cudnn_pygraph.frost_engines_enabled() + + # A use site second, on an already-imported cuDNN. This is the case that used to be + # skipped. + frost_attention._import_cudnn_frontend(enable_frost_engines=True) + assert os.environ.get(env) == "1", "FROST was requested after the import and not enabled" + assert cudnn_pygraph.frost_engines_enabled() + 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, + frost_attention, + ) + + try: + cudnn = frost_attention._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(), + backend_name="FrostAttention", + 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/utils.py b/tests/pytorch/utils.py index 2a848842b52..2a7cb1c449d 100644 --- a/tests/pytorch/utils.py +++ b/tests/pytorch/utils.py @@ -511,7 +511,7 @@ def test(): _attention_backends["backend_selection_requires_update"] = False return available_backends, flash_attention_backend, fused_attention_backend - backends = {1: "F16_arbitrary_seqlen", 2: "FP8"} + backends = {1: "F16_arbitrary_seqlen", 2: "FP8", 3: "FROST"} if AttentionLogging._is_logging_setup is False: AttentionLogging.setup_logging() diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index f90fc46c032..c132d49d9e4 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -2036,9 +2036,7 @@ def _fused_attn_setup_ctx( bwd_args.softmax_type = fwd_args.softmax_type bwd_args.window_size = fwd_args.window_size bwd_args.bottom_right_diagonal = fwd_args.bottom_right_diagonal - bwd_args.fused_attention_backend = ( - ctx_attrs["fused_attention_backend"] if fp8 else FusedAttnBackend["F16_arbitrary_seqlen"] - ) + bwd_args.fused_attention_backend = ctx_attrs["fused_attention_backend"] bwd_args.use_FAv2_bwd = fwd_args.use_FAv2_bwd bwd_args.deterministic = fwd_args.deterministic @@ -2604,8 +2602,9 @@ def forward( ) if context_parallel: - assert ( - fp8 or fused_attention_backend == FusedAttnBackend["F16_arbitrary_seqlen"] + assert fp8 or fused_attention_backend in ( + FusedAttnBackend["F16_arbitrary_seqlen"], + FusedAttnBackend["FROST"], ), f"{fused_attention_backend} does not work with context parallelism!" assert core_attention_bias_type not in [ "alibi" @@ -2637,6 +2636,7 @@ def forward( attn_bias=core_attention_bias, deterministic=self.deterministic, use_fused_attention=True, + fused_attention_backend=fused_attention_backend, window_size=window_size, fp8=fp8, fp8_meta=fp8_meta, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index 3a153a3f994..f36bb46401c 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1109,8 +1109,9 @@ def cp_p2p_fwd_fused_attn( softmax_lse_per_step, rng_states, *rest = aux_ctx_tensors attn_bias = rest[0] if len(rest) > 0 else None + # Indexed rather than starred: every caller unpacks exactly five values. if return_max_logit: - return out_per_step, softmax_lse_per_step, rng_states, attn_bias, *max_logit + return out_per_step, softmax_lse_per_step, rng_states, attn_bias, max_logit[0] return out_per_step, softmax_lse_per_step, rng_states, attn_bias, None @@ -1605,6 +1606,7 @@ def forward( attn_bias, deterministic, use_fused_attention, + fused_attention_backend, return_max_logit, softcap, fp8, @@ -1780,7 +1782,9 @@ def forward( # q, k, v: torch.Tensor, dtype=fwd_nominal_dtype q_f16 = q if use_fused_attention: - fused_attn_backend = FusedAttnBackend["F16_arbitrary_seqlen"] + fused_attn_backend = ( + fused_attention_backend or FusedAttnBackend["F16_arbitrary_seqlen"] + ) if return_max_logit: max_logit_per_step = [ torch.empty(q.shape[-2], dtype=q.dtype, device=q.device) for _ in range(2) @@ -2386,6 +2390,7 @@ def forward( ctx.deterministic = deterministic ctx.softcap = softcap ctx.use_fused_attention = use_fused_attention + ctx.fused_attention_backend = fused_attention_backend ctx.pad_between_seqs = pad_between_seqs ctx.softmax_lse_in_packed_format = softmax_lse_in_packed_format ctx.second_half_lse_seqlen = second_half_lse_seqlen @@ -2630,7 +2635,9 @@ def backward(ctx, dout, *_args): ] p2p_comm_buffers[0][0].copy_(kv) if ctx.use_fused_attention: - fused_attn_backend = FusedAttnBackend["F16_arbitrary_seqlen"] + fused_attn_backend = ( + ctx.fused_attention_backend or FusedAttnBackend["F16_arbitrary_seqlen"] + ) # communicate for the 'a2a' part of 'a2a+p2p' dout = dout.view(*ctx.orig_o_shape) @@ -3186,6 +3193,7 @@ def backward(ctx, dout, *_args): attn_dbias, None, None, + None, # fused_attention_backend None, None, None, @@ -3275,6 +3283,7 @@ def forward( attn_bias, deterministic, use_fused_attention, + fused_attention_backend, return_max_logit, softcap, window_size, @@ -3441,7 +3450,7 @@ def forward( fp8_meta_kwargs["s_quantizer"] = S_quantizer fp8_meta_kwargs["o_quantizer"] = O_quantizer elif use_fused_attention: - fused_attn_backend = FusedAttnBackend["F16_arbitrary_seqlen"] + fused_attn_backend = fused_attention_backend or FusedAttnBackend["F16_arbitrary_seqlen"] orig_q_shape, _, orig_v_shape = q.shape, k.shape, v.shape orig_o_shape = orig_q_shape[:-1] + orig_v_shape[-1:] @@ -3964,6 +3973,7 @@ def forward( ctx.deterministic = deterministic ctx.softcap = softcap ctx.use_fused_attention = use_fused_attention + ctx.fused_attention_backend = fused_attention_backend ctx.use_flash_attn_3 = use_flash_attn_3 ctx.use_flash_attn_4 = use_flash_attn_4 ctx.pad_between_seqs = pad_between_seqs @@ -4249,7 +4259,9 @@ def backward(ctx, dout, *_args): softmax_lse_per_step[i], rng_states[i], ] - fused_attn_backend = FusedAttnBackend["F16_arbitrary_seqlen"] + fused_attn_backend = ( + ctx.fused_attention_backend or FusedAttnBackend["F16_arbitrary_seqlen"] + ) fp8_meta_kwargs = {} new_qkv_layout = ctx.qkv_layout do_format = ctx.o_format @@ -4548,6 +4560,7 @@ def backward(ctx, dout, *_args): None, None, None, + None, # fused_attention_backend None, None, None, @@ -4591,6 +4604,7 @@ def forward( attn_bias, deterministic, use_fused_attention, + fused_attention_backend, return_max_logit, softcap, window_size, @@ -4735,7 +4749,9 @@ def forward( fp8_meta_kwargs["o_quantizer"] = O_quantizer else: if use_fused_attention: - fused_attn_backend = FusedAttnBackend["F16_arbitrary_seqlen"] + fused_attn_backend = ( + fused_attention_backend or FusedAttnBackend["F16_arbitrary_seqlen"] + ) # q, k, v: # FP8DS/FP8CS: torch.uint8 @@ -5026,6 +5042,7 @@ def forward( ctx.softcap = softcap ctx.window_size = window_size ctx.use_fused_attention = use_fused_attention + ctx.fused_attention_backend = fused_attention_backend ctx.fp8_meta = fp8_meta ctx.is_input_fp8 = is_input_fp8 ctx.is_output_fp8 = is_output_fp8 @@ -5103,7 +5120,9 @@ def backward(ctx, dout, *_args): if isinstance(dout, QuantizedTensorStorage): dout = dout.dequantize(dtype=bwd_nominal_dtype) if ctx.use_fused_attention: - fused_attn_backend = FusedAttnBackend["F16_arbitrary_seqlen"] + fused_attn_backend = ( + ctx.fused_attention_backend or FusedAttnBackend["F16_arbitrary_seqlen"] + ) dout = dout.view(*ctx.orig_o_shape) # dout: @@ -5396,6 +5415,7 @@ def backward(ctx, dout, *_args): d_bias, None, None, + None, # fused_attention_backend None, None, None, @@ -5549,6 +5569,7 @@ def attn_forward_func_with_cp( attn_bias=None, deterministic=False, use_fused_attention=False, + fused_attention_backend=None, window_size=None, softcap=0.0, fp8=False, @@ -5736,6 +5757,7 @@ def attn_forward_func_with_cp( attn_bias, deterministic, use_fused_attention, + fused_attention_backend, return_max_logit, softcap, ] diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py new file mode 100644 index 00000000000..b49e13fd7eb --- /dev/null +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_pygraph.py @@ -0,0 +1,481 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Helpers for driving cuDNN frontend's Python graph API from PyTorch. + +Shared by flex_attention.py and frost_attention.py: importing the frontend, holding a handle per +device, describing tensors in cuDNN's BHSD form, and building and planning the SDPA forward and +backward graphs. Anything specific to one backend, such as a mask, a score_mod or which engine +to pin, is passed in by the caller. +""" + +from __future__ import annotations + +import contextlib +import importlib +import os +from typing import Any, Dict, Optional, Sequence, Tuple + +import torch + +# The cuDNN FROST SDPA engines, named once so the two opposite instructions about them cannot +# drift: frost pins one of these by name, flex bars both. cuDNN has already renamed this family +# once (collapsing the per-head-dim ``..._d512`` rows), and if these strings stopped matching, +# frost would fail loudly while flex failed silently. +FROST_FWD_PLAN_TOKEN = "sdpa_fwd_prefill_sm100" +FROST_BWD_PLAN_TOKEN = "sdpa_bwd_sm100" +FROST_PLAN_TOKENS = (FROST_FWD_PLAN_TOKEN, FROST_BWD_PLAN_TOKEN) + +# Distinguishes "the caller said nothing" from "the caller asked for no exclusions at all". +_BAR_FROST_BY_DEFAULT = object() + +_cudnn = None +_frost_engines_enabled = False +_HANDLES: Dict[Tuple[str, torch.device], Any] = {} + + +def import_cudnn_frontend(enable_frost_engines: bool = False): + """Import cuDNN Frontend, enabling the FROST engines if this caller needs them. + + The switch ranks FROST ahead of the backend engines process-wide, so it defaults off and + FROST asks explicitly. A caller needing a particular engine should verify by plan name rather + than trust the switch, which is what ``finalize_plans(require_plan_token=...)`` does. + + The enabling sits outside the import memo deliberately. Inside it, whichever backend imported + cuDNN first would decide for the process, and a flex-first import would leave every later + FROST call with no engine on offer. Enabling late works because cuDNN re-reads the environment + per graph, inside ``offered_ids()`` on every ``create_execution_plans``. + """ + global _cudnn, _frost_engines_enabled # pylint: disable=global-statement + if _cudnn is None: + try: + _cudnn = importlib.import_module("cudnn") + except ImportError as exc: + raise ImportError( + "cuDNN frontend Python package not found. " + "Install it with: pip install nvidia-cudnn-frontend" + ) from exc + + if enable_frost_engines and not _frost_engines_enabled: + os.environ.setdefault("CUDNN_FRONTEND_ENABLE_FROST_ENGINES", "1") + importlib.import_module("cudnn.sdpa") + _frost_engines_enabled = True + + return _cudnn + + +def cudnn_module(): + """The imported frontend, or None if nothing has imported it yet. + + For a caller that must inspect it without triggering an import, such as a version probe. + """ + return _cudnn + + +def frost_engines_enabled() -> bool: + """Whether this process has switched the FROST engines on.""" + return _frost_engines_enabled + + +def handle_for(device: torch.device, *, backend_name: str = "cuDNN attention"): + """A cuDNN handle for ``device``, rebound to PyTorch's current stream on every call. + + Without the rebinding cuDNN runs on its handle's own stream while the tensors and workspace + sit on PyTorch's, with nothing ordering the two. The p2p context-parallel ring issues + attention inside ``with torch.cuda.stream(cp_stream)``, and executes one cached plan from + different streams across ring steps, so this is per call rather than per handle. + + Keyed on ``(backend_name, device)``: a cuDNN handle is not thread-safe and ``set_stream`` + mutates it, so one handle shared by two backends widens an existing race for no benefit. + """ + if device.type != "cuda": + raise ValueError(f"{backend_name} only supports CUDA tensors, got device {device}.") + cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() + if device.index is None: + device = torch.device("cuda", torch.cuda.current_device()) + key = (backend_name, device) + with torch.cuda.device(device): + handle = _HANDLES.get(key) + if handle is None: + handle = cudnn.create_handle() + _HANDLES[key] = handle + cudnn.set_stream(handle=handle, stream=torch.cuda.current_stream(device).cuda_stream) + return handle + + +def io_data_type(cudnn, dtype: torch.dtype, *, backend_name: str = "cuDNN attention"): + """Map a torch dtype to the cuDNN enum these SDPA graphs are declared with. + + Takes ``cudnn`` rather than importing it, so a dtype lookup cannot flip the FROST switch. + """ + if dtype == torch.float16: + return cudnn.data_type.HALF + if dtype == torch.bfloat16: + return cudnn.data_type.BFLOAT16 + raise ValueError(f"{backend_name} only supports FP16/BF16 tensors, got {dtype}.") + + +def build_pygraph( + dtype: torch.dtype, device: torch.device, *, backend_name: str = "cuDNN attention" +): + """A cuDNN frontend graph for F16/BF16 SDPA, bound to this device's stream-current handle.""" + cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() + return cudnn.pygraph( + io_data_type=io_data_type(cudnn, dtype, backend_name=backend_name), + intermediate_data_type=cudnn.data_type.FLOAT, + compute_data_type=cudnn.data_type.FLOAT, + handle=handle_for(device, backend_name=backend_name), + ) + + +def bhsd_dim_stride( + tensor: torch.Tensor, tensor_format: str, *, backend_name: str = "cuDNN attention" +) -> Tuple[Tuple[int, ...], Tuple[int, ...]]: + """Describe an SBHD/BSHD tensor as cuDNN frontend's logical BHSD format. + + The tensor is never permuted: cuDNN takes dims and strides, so reordering the descriptors + says the same thing and costs nothing. + """ + if tensor_format == "sbhd": + return ( + (tensor.shape[1], tensor.shape[2], tensor.shape[0], tensor.shape[3]), + (tensor.stride(1), tensor.stride(2), tensor.stride(0), tensor.stride(3)), + ) + if tensor_format == "bshd": + return ( + (tensor.shape[0], tensor.shape[2], tensor.shape[1], tensor.shape[3]), + (tensor.stride(0), tensor.stride(2), tensor.stride(1), tensor.stride(3)), + ) + raise ValueError(f"{backend_name} only supports SBHD/BSHD tensor formats, got {tensor_format}.") + + +def bhsd_graph_tensor( + graph, tensor: torch.Tensor, tensor_format: str, *, backend_name: str = "cuDNN attention" +): + """Create a cuDNN graph tensor with BHSD dims and TE-layout strides.""" + dim, stride = bhsd_dim_stride(tensor, tensor_format, backend_name=backend_name) + return graph.tensor(dim=dim, stride=stride, data_type=tensor.dtype) + + +def device_key(device: torch.device) -> Tuple[Any, ...]: + """Normalize a device for a cache key, so ``cuda`` and ``cuda:0`` cannot key two entries.""" + if device.type == "cuda" and device.index is None: + return ("cuda", torch.cuda.current_device()) + return (device.type, device.index) + + +def tensor_key( + tensor: torch.Tensor, tensor_format: str, *, backend_name: str = "cuDNN attention" +) -> Tuple[Any, ...]: + """A tensor as the graph will see it: BHSD dims, BHSD strides, dtype. + + Strides are keyed because the graph is built for this exact layout, which is what lets bshd + and sbhd run without a transpose. The format itself is not: two formats giving the same + description are the same graph. + """ + dim, stride = bhsd_dim_stride(tensor, tensor_format, backend_name=backend_name) + return (tuple(dim), tuple(stride), tensor.dtype) + + +def cached_graph(cache: Dict[Any, Any], key: Optional[Any], build, *, device=None): + """Memoize a built graph. ``key=None`` means uncacheable: build, return, do not store. + + The build runs under ``device`` rather than merely with its handle, because the plans are + JIT-compiled and a compile path may read the ambient CUDA context. One build site, not one + per branch: a second is free to forget that. + """ + entry = None if key is None else cache.get(key) + if entry is None: + scope = ( + torch.cuda.device(device) + if device is not None and device.type == "cuda" + else contextlib.nullcontext() + ) + with scope: + entry = build() + if key is not None: + cache[key] = entry + return entry + + +def finalize_plans( + graph, + *, + backend_name: str = "cuDNN attention", + heuristics: Optional[Sequence[Any]] = None, + build_policy: Any = None, + require_plan_token: Optional[str] = None, + not_found_hint: Any = "", + exclude_plan_tokens: Any = _BAR_FROST_BY_DEFAULT, +) -> Tuple[int, Optional[str]]: + """Create plans, optionally pin one by name, build, and return (workspace size, plan name). + + ``require_plan_token`` makes the choice strict: a plan whose name lacks the token raises. + Unpinned, ``build_plans`` walks the ranked list and finalizes the first that builds, so a + graph the intended engine declines runs on whatever cuDNN ranked next with nothing in the + return value saying so. The token is matched as a substring because cuDNN has already + collapsed per-head-dim engine names into one row once. The pin has to precede + ``check_support``: selecting a plan sets cuDNN's ``_plan_pinned``, and only then is a decline + fatal rather than recorded. + + ``exclude_plan_tokens`` is the opposite instruction, defaulting to barring the FROST engines, + which accept a score_mod graph and then compute without it. The switch offering them is + process-wide, so declining to ask is not enough. Forgetting to exclude is silently wrong while + excluding wrongly costs a slower plan or a loud decline, so the default favours the caller + that does not want them; it is skipped when a plan is pinned, which has already named one. + + Remove the default once cuDNN is fixed: the gate that should decline these graphs reads a key + ``sdpa()`` never writes, so it never fires. The match is on a substring of the engine name, so + a rename would break it; what catches that is the end-to-end test comparing flex's output + across the switch, not anything here. Measured: deselect marks rather than removes, so the + plan list still names a barred engine afterwards and cannot be asserted on. + """ + cudnn = _cudnn if _cudnn is not None else import_cudnn_frontend() + + 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: + if exclude_plan_tokens is _BAR_FROST_BY_DEFAULT: + exclude_plan_tokens = FROST_PLAN_TOKENS + try: + graph.create_execution_plans(list(heuristics)) + # Bar the named engines before the walk, so build_plans falls through to the first + # entry that is both unbarred and buildable. Inert where they are not on offer, which + # is every process that has not enabled them. getattr so a frontend predating + # deselect_engines degrades rather than raising. + deselect = getattr(graph, "deselect_engines", None) + if exclude_plan_tokens and deselect is not None: + deselect(list(exclude_plan_tokens)) + graph.check_support() + except cudnn.cudnnGraphNotSupportedError as exc: + # Name the bar in the message: if it removed the only viable plan, the graph is not + # what was unsupported. + barred = ( + f" (barred engines: {list(exclude_plan_tokens)})" if exclude_plan_tokens else "" + ) + raise RuntimeError( + f"cuDNN {backend_name} SDPA graph is not supported: {exc}{barred}" + ) from exc + if build_policy is None: + build_policy = cudnn.build_plan_policy.HEURISTICS_CHOICE + graph.build_plans(build_policy) + return max(graph.get_workspace_size(), 1), None + + graph.create_execution_plans(list(heuristics)) + names = [graph.get_plan_name_at_index(i) for i in range(graph.get_execution_plan_count())] + hits = [i for i, n in enumerate(names) if require_plan_token in n] + if not hits: + # Callable hints are resolved only here: a caller may want to look up package versions to + # explain the failure, and that work should not happen on the success path. + hint = not_found_hint() if callable(not_found_hint) else not_found_hint + raise RuntimeError( + f"no cuDNN engine matching {require_plan_token!r} was offered." + f" Candidate plans: {names[:6]}.{(' ' + hint) if hint else ''}" + ) + graph.select_plan(hits[0]) + # The engine is pinned, so a decline here is its own verdict and cuDNN puts the reason in the + # exception. Surface it: a plan offered and then refused is the harder failure to read. + try: + graph.check_support() + graph.build_plans() + except cudnn.cudnnGraphNotSupportedError as exc: + hint = not_found_hint() if callable(not_found_hint) else not_found_hint + raise RuntimeError( + f"cuDNN engine {names[hits[0]]!r} was offered but declined this graph:" + f" {exc}{(' ' + hint) if hint else ''}" + ) from exc + return max(graph.get_workspace_size(), 1), names[hits[0]] + + +def _declare_tensor(graph, name: str, spec: Any): + """Declare one graph input. + + ``spec`` is a ``torch.Tensor``, described with ``tensor_like``, or a ``(dim, stride)`` / + ``(dim, stride, data_type)`` descriptor. Omitting ``data_type`` lets the tensor inherit the + graph's ``io_data_type``. + """ + if isinstance(spec, torch.Tensor): + return graph.tensor_like(spec) + kwargs: Dict[str, Any] = {"name": name, "dim": list(spec[0]), "stride": list(spec[1])} + if len(spec) > 2 and spec[2] is not None: + kwargs["data_type"] = spec[2] + return graph.tensor(**kwargs) + + +def _mark_output(tensor, spec: Any): + """Mark a graph output and apply whichever of dim/stride/data_type ``spec`` supplies. + + Each field is optional because the backends specify different subsets, and setting one a + caller left out would be describing the tensor for it rather than from it. + """ + tensor.set_output(True) + if spec[0] is not None: + tensor.set_dim(list(spec[0])) + if spec[1] is not None: + tensor.set_stride(list(spec[1])) + if len(spec) > 2 and spec[2] is not None: + tensor.set_data_type(spec[2]) + return tensor + + +def _declare_aux(graph, aux_tensors): + """Declare the auxiliary runtime tensors an sdpa callback reads, grouped by role.""" + return { + group: {name: graph.tensor_like(t) for name, t in tensors.items()} + for group, tensors in (aux_tensors or {}).items() + } + + +def _resolve_sdpa_kwargs(sdpa_kwargs, aux): + """Extra ``sdpa``/``sdpa_backward`` arguments, as a dict or a callable taking ``aux``. + + The callable form exists because a score_mod closes over graph tensors that cannot be built + until the graph is, so the caller gets them handed back here rather than building its own. + """ + return sdpa_kwargs(aux) if callable(sdpa_kwargs) else dict(sdpa_kwargs or {}) + + +def build_fwd( + *, + dtype: torch.dtype, + device: torch.device, + backend_name: str, + name: str, + q: Any, + k: Any, + v: Any, + out: Any, + attn_scale: float, + stats: Any = None, + aux_tensors: Optional[Dict[str, Dict[str, torch.Tensor]]] = None, + sdpa_kwargs: Any = None, + heuristics: Optional[Sequence[Any]] = None, + require_plan_token: Optional[str] = None, + not_found_hint: Any = "", + exclude_plan_tokens: Any = _BAR_FROST_BY_DEFAULT, +) -> Dict[str, Any]: + """Build and plan an SDPA forward graph. + + ``stats`` is the LSE output descriptor, or ``None`` to skip generating it. Everything + backend-specific -- a mask, a score_mod, which engine to pin or bar -- arrives through + ``sdpa_kwargs`` and the plan arguments, which are passed to :func:`finalize_plans`. + """ + graph = build_pygraph(dtype, device, backend_name=backend_name) + tq = _declare_tensor(graph, "q", q) + tk = _declare_tensor(graph, "k", k) + tv = _declare_tensor(graph, "v", v) + aux = _declare_aux(graph, aux_tensors) + tout, tstats = graph.sdpa( + name=name, + q=tq, + k=tk, + v=tv, + generate_stats=stats is not None, + attn_scale=attn_scale, + **_resolve_sdpa_kwargs(sdpa_kwargs, aux), + ) + _mark_output(tout, out) + if stats is None: + tstats = None + else: + _mark_output(tstats, stats) + workspace, plan = finalize_plans( + graph, + backend_name=backend_name, + heuristics=heuristics, + require_plan_token=require_plan_token, + not_found_hint=not_found_hint, + exclude_plan_tokens=exclude_plan_tokens, + ) + return { + "graph": graph, + "q": tq, + "k": tk, + "v": tv, + "out": tout, + "stats": tstats, + "aux": aux, + "workspace": workspace, + "plan": plan, + } + + +def build_bwd( + *, + dtype: torch.dtype, + device: torch.device, + backend_name: str, + name: str, + q: Any, + k: Any, + v: Any, + o: Any, + do: Any, + stats: Any, + dq: Any, + dk: Any, + dv: Any, + attn_scale: float, + deterministic: bool = False, + aux_tensors: Optional[Dict[str, Dict[str, torch.Tensor]]] = None, + sdpa_kwargs: Any = None, + heuristics: Optional[Sequence[Any]] = None, + require_plan_token: Optional[str] = None, + not_found_hint: Any = "", + exclude_plan_tokens: Any = _BAR_FROST_BY_DEFAULT, +) -> Dict[str, Any]: + """Build and plan an SDPA backward graph. The counterpart of :func:`build_fwd`.""" + graph = build_pygraph(dtype, device, backend_name=backend_name) + handles = { + n: _declare_tensor(graph, n, spec) + for n, spec in (("q", q), ("k", k), ("v", v), ("o", o), ("do", do), ("stats", stats)) + } + aux = _declare_aux(graph, aux_tensors) + tdq, tdk, tdv = graph.sdpa_backward( + name=name, + q=handles["q"], + k=handles["k"], + v=handles["v"], + o=handles["o"], + dO=handles["do"], + stats=handles["stats"], + attn_scale=attn_scale, + use_deterministic_algorithm=deterministic, + **_resolve_sdpa_kwargs(sdpa_kwargs, aux), + ) + for handle, spec in ((tdq, dq), (tdk, dk), (tdv, dv)): + _mark_output(handle, spec) + workspace, plan = finalize_plans( + graph, + backend_name=backend_name, + heuristics=heuristics, + require_plan_token=require_plan_token, + not_found_hint=not_found_hint, + exclude_plan_tokens=exclude_plan_tokens, + ) + handles.update({"dq": tdq, "dk": tdk, "dv": tdv}) + return {"graph": graph, "aux": aux, "workspace": workspace, "plan": plan, **handles} + + +def execute_graph( + graph, + variant_pack: Dict[Any, torch.Tensor], + workspace_size: int, + device: torch.device, + *, + backend_name: str = "cuDNN attention", +): + """Allocate a built graph's workspace and run it on the device's current stream. + + The handle is resolved here rather than by the caller so it is always rebound immediately + before ``execute``, which is what keeps cuDNN on the same stream as the tensors. + """ + if device.type == "cuda" and device.index is None: + device = torch.device("cuda", torch.cuda.current_device()) + workspace = torch.empty(workspace_size, device=device, dtype=torch.uint8) + graph.execute(variant_pack, workspace, handle=handle_for(device, backend_name=backend_name)) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index b9593b42d9b..74911250db5 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -5,49 +5,32 @@ """cuDNN-backed Flex Attention helpers.""" from dataclasses import dataclass -import importlib import inspect from typing import Any, Callable, Dict, Optional, Tuple import torch -_cudnn_score_mod_handles: Dict[torch.device, Any] = {} +from . import cudnn_pygraph + +_BACKEND_NAME = "Flex Attention" _cudnn_score_mod_graph_cache: Dict[Tuple[Any, ...], Any] = {} _SCORE_MOD_UNCACHEABLE = object() def _import_cudnn_frontend(): - """Import the cuDNN frontend Python package.""" - try: - return importlib.import_module("cudnn") - except ImportError as exc: - raise ImportError( - "cuDNN frontend Python package not found. " - "Install it with: pip install nvidia-cudnn-frontend" - ) from exc + """Import the cuDNN frontend Python package. + + Never enables the FROST engines: that switch is process-wide, so a backend that does not want + them must not ask. See ``cudnn_pygraph.import_cudnn_frontend``. + """ + return cudnn_pygraph.import_cudnn_frontend() def _bhsd_dim_stride( tensor: torch.Tensor, tensor_format: str ) -> Tuple[Tuple[int, ...], Tuple[int, ...]]: """Describe an SBHD/BSHD tensor as cuDNN frontend's logical BHSD format.""" - if tensor_format == "sbhd": - return ( - (tensor.shape[1], tensor.shape[2], tensor.shape[0], tensor.shape[3]), - (tensor.stride(1), tensor.stride(2), tensor.stride(0), tensor.stride(3)), - ) - if tensor_format == "bshd": - return ( - (tensor.shape[0], tensor.shape[2], tensor.shape[1], tensor.shape[3]), - (tensor.stride(0), tensor.stride(2), tensor.stride(1), tensor.stride(3)), - ) - raise ValueError(f"Flex Attention only supports SBHD/BSHD tensor formats, got {tensor_format}.") - - -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_dim_stride(tensor, tensor_format, backend_name=_BACKEND_NAME) # score_mod graph cache helpers. @@ -140,12 +123,7 @@ def _score_mod_callback_cache_key(callback: Optional[Callable]) -> Any: def _score_mod_device_key(device: torch.device) -> Tuple[Any, ...]: """Normalize a tensor device for graph cache keys.""" - if device.type == "cuda": - index = device.index - if index is None: - index = torch.cuda.current_device() - return (device.type, index) - return (device.type, device.index) + return cudnn_pygraph.device_key(device) def _score_mod_tensor_metadata(tensor: torch.Tensor) -> Tuple[Any, ...]: @@ -169,15 +147,9 @@ def _score_mod_tensor_dict_metadata( def _score_mod_bhsd_tensor_metadata(tensor: torch.Tensor, tensor_format: str) -> Tuple[Any, ...]: """Describe an SBHD/BSHD runtime tensor as a cuDNN BHSD graph tensor.""" - dim, stride = _bhsd_dim_stride(tensor, tensor_format) - return (dim, stride, tensor.dtype, _score_mod_device_key(tensor.device)) - - -def _make_cudnn_graph_tensor_dict(graph, tensors: Optional[Dict[str, torch.Tensor]]): - """Create cuDNN graph tensors matching runtime tensors.""" - if tensors is None: - return {} - return {name: graph.tensor_like(tensor) for name, tensor in tensors.items()} + return cudnn_pygraph.tensor_key(tensor, tensor_format, backend_name=_BACKEND_NAME) + ( + cudnn_pygraph.device_key(tensor.device), + ) # cuDNN frontend score_mod graph helpers. @@ -192,42 +164,10 @@ def _wrapped_score_mod(sdpa_graph, score_tensor): return _wrapped_score_mod -def _get_cudnn_current_stream_handle(cudnn, device: torch.device): - """Return a cuDNN handle for device, bound to PyTorch's current stream.""" - if device.type != "cuda": - raise ValueError(f"Flex Attention only supports CUDA tensors, got device {device}.") - if device.index is None: - device = torch.device("cuda", torch.cuda.current_device()) - - handle = _cudnn_score_mod_handles.get(device) - with torch.cuda.device(device): - if handle is None: - handle = cudnn.create_handle() - _cudnn_score_mod_handles[device] = handle - - stream = torch.cuda.current_stream(device).cuda_stream - cudnn.set_stream(handle=handle, stream=stream) - return handle - - -def _build_cudnn_pygraph(dtype: torch.dtype, device: torch.device): - """Create a cuDNN frontend Python graph for F16/BF16 SDPA.""" - 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 +def _finalize_cudnn_graph(graph) -> int: + """Build a cuDNN frontend Python graph and return its workspace size.""" + workspace_size, _ = cudnn_pygraph.finalize_plans(graph, backend_name=_BACKEND_NAME) + return workspace_size @dataclass @@ -263,21 +203,6 @@ class _CudnnScoreModBwdGraphEntry: workspace_size: int -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) - - def _execute_cudnn_graph( graph, variant_pack: Dict[Any, torch.Tensor], @@ -285,19 +210,8 @@ def _execute_cudnn_graph( device: torch.device, ): """Execute a built cuDNN frontend Python graph.""" - cudnn = _import_cudnn_frontend() - - if device.type == "cuda" and device.index is None: - device = torch.device("cuda", torch.cuda.current_device()) - workspace = torch.empty( - workspace_size, - device=device, - dtype=torch.uint8, - ) - graph.execute( - variant_pack, - workspace, - handle=_get_cudnn_current_stream_handle(cudnn, device), + cudnn_pygraph.execute_graph( + graph, variant_pack, workspace_size, device, backend_name=_BACKEND_NAME ) @@ -378,6 +292,12 @@ def _cudnn_score_mod_bwd_cache_key( ) +def _bhsd_tensor_spec(tensor: torch.Tensor, tensor_format: str): + """Describe a tensor for ``cudnn_pygraph``: BHSD dims with TE-layout strides.""" + dim, stride = _bhsd_dim_stride(tensor, tensor_format) + return (dim, stride, tensor.dtype) + + def _build_cudnn_score_mod_fwd_graph( is_training: bool, query_layer: torch.Tensor, @@ -393,46 +313,35 @@ def _build_cudnn_score_mod_fwd_graph( ) -> _CudnnScoreModFwdGraphEntry: """Build a cached cuDNN frontend graph for score_mod fprop.""" cudnn = _import_cudnn_frontend() + if is_training: + assert stats is not None - graph = _build_cudnn_pygraph(query_layer.dtype, query_layer.device) - q = _bhsd_graph_tensor(graph, query_layer, q_format) - k = _bhsd_graph_tensor(graph, key_layer, kv_format) - v = _bhsd_graph_tensor(graph, value_layer, kv_format) - - score_mod_graph_tensors = _make_cudnn_graph_tensor_dict(graph, score_mod_tensors) - wrapped_score_mod = _wrap_score_mod(score_mod, score_mod_graph_tensors) - - output_dim, output_stride = _bhsd_dim_stride(output_layer, q_format) - output, stats_tensor = graph.sdpa( + entry = cudnn_pygraph.build_fwd( + dtype=query_layer.dtype, + device=query_layer.device, + backend_name=_BACKEND_NAME, name="te_score_mod_sdpa", - q=q, - k=k, - v=v, - generate_stats=is_training, + q=_bhsd_tensor_spec(query_layer, q_format), + k=_bhsd_tensor_spec(key_layer, kv_format), + v=_bhsd_tensor_spec(value_layer, kv_format), + out=_bhsd_dim_stride(output_layer, q_format), + stats=((stats.size(), stats.stride(), cudnn.data_type.FLOAT) if is_training else None), attn_scale=attn_scale, - use_causal_mask=False, - score_mod=wrapped_score_mod, + aux_tensors={"score_mod": score_mod_tensors or {}}, + sdpa_kwargs=lambda aux: { + "use_causal_mask": False, + "score_mod": _wrap_score_mod(score_mod, aux["score_mod"]), + }, ) - output.set_output(True).set_dim(output_dim).set_stride(output_stride) - - if is_training: - assert stats is not None - stats_tensor.set_output(True).set_dim(stats.size()).set_stride( - stats.stride() - ).set_data_type(cudnn.data_type.FLOAT) - else: - stats_tensor = None - - workspace_size = _finalize_cudnn_graph(graph) return _CudnnScoreModFwdGraphEntry( - graph=graph, - q=q, - k=k, - v=v, - output=output, - stats=stats_tensor, - score_mod_graph_tensors=score_mod_graph_tensors, - workspace_size=workspace_size, + graph=entry["graph"], + q=entry["q"], + k=entry["k"], + v=entry["v"], + output=entry["out"], + stats=entry["stats"], + score_mod_graph_tensors=entry["aux"]["score_mod"], + workspace_size=entry["workspace"], ) @@ -463,14 +372,12 @@ def _get_cudnn_score_mod_fwd_graph( output_layer, stats, ) - key = _cudnn_score_mod_fwd_cache_key(*build_args) - if key is None: - return _build_cudnn_score_mod_fwd_graph(*build_args) - entry = _cudnn_score_mod_graph_cache.get(key) - if entry is None: - entry = _build_cudnn_score_mod_fwd_graph(*build_args) - _cudnn_score_mod_graph_cache[key] = entry - return entry + return cudnn_pygraph.cached_graph( + _cudnn_score_mod_graph_cache, + _cudnn_score_mod_fwd_cache_key(*build_args), + lambda: _build_cudnn_score_mod_fwd_graph(*build_args), + device=query_layer.device, + ) def _build_cudnn_score_mod_bwd_graph( @@ -490,62 +397,52 @@ def _build_cudnn_score_mod_bwd_graph( deterministic: bool, ) -> _CudnnScoreModBwdGraphEntry: """Build a cached cuDNN frontend graph for score_mod bprop.""" - graph = _build_cudnn_pygraph(query_layer.dtype, query_layer.device) - q = _bhsd_graph_tensor(graph, query_layer, q_format) - k = _bhsd_graph_tensor(graph, key_layer, kv_format) - v = _bhsd_graph_tensor(graph, value_layer, kv_format) - output = _bhsd_graph_tensor(graph, output_layer, q_format) - d_output = _bhsd_graph_tensor(graph, d_out, q_format) - stats_tensor = graph.tensor_like(stats) - - score_mod_graph_tensors = _make_cudnn_graph_tensor_dict(graph, score_mod_tensors) - score_mod_bprop_graph_tensors = ( - _make_cudnn_graph_tensor_dict(graph, score_mod_bprop_tensors) - if score_mod_bprop is not None - else {} - ) - wrapped_score_mod = _wrap_score_mod(score_mod, score_mod_graph_tensors) - wrapped_score_mod_bprop = _wrap_score_mod(score_mod_bprop, score_mod_bprop_graph_tensors) - dq_layer = torch.empty_like(query_layer) dk_layer = torch.empty_like(key_layer) dv_layer = torch.empty_like(value_layer) - dq_dim, dq_stride = _bhsd_dim_stride(dq_layer, q_format) - dk_dim, dk_stride = _bhsd_dim_stride(dk_layer, kv_format) - dv_dim, dv_stride = _bhsd_dim_stride(dv_layer, kv_format) - dq, dk, dv = graph.sdpa_backward( + + entry = cudnn_pygraph.build_bwd( + dtype=query_layer.dtype, + device=query_layer.device, + backend_name=_BACKEND_NAME, name="te_score_mod_sdpa_backward", - q=q, - k=k, - v=v, - o=output, - dO=d_output, - stats=stats_tensor, + q=_bhsd_tensor_spec(query_layer, q_format), + k=_bhsd_tensor_spec(key_layer, kv_format), + v=_bhsd_tensor_spec(value_layer, kv_format), + o=_bhsd_tensor_spec(output_layer, q_format), + do=_bhsd_tensor_spec(d_out, q_format), + stats=stats, + dq=_bhsd_dim_stride(dq_layer, q_format), + dk=_bhsd_dim_stride(dk_layer, kv_format), + dv=_bhsd_dim_stride(dv_layer, kv_format), attn_scale=attn_scale, - use_causal_mask=False, - score_mod=wrapped_score_mod, - score_mod_bprop=wrapped_score_mod_bprop, - use_deterministic_algorithm=deterministic, + deterministic=deterministic, + aux_tensors={ + "score_mod": score_mod_tensors or {}, + "score_mod_bprop": ( + score_mod_bprop_tensors or {} if score_mod_bprop is not None else {} + ), + }, + sdpa_kwargs=lambda aux: { + "use_causal_mask": False, + "score_mod": _wrap_score_mod(score_mod, aux["score_mod"]), + "score_mod_bprop": _wrap_score_mod(score_mod_bprop, aux["score_mod_bprop"]), + }, ) - dq.set_output(True).set_dim(dq_dim).set_stride(dq_stride) - dk.set_output(True).set_dim(dk_dim).set_stride(dk_stride) - dv.set_output(True).set_dim(dv_dim).set_stride(dv_stride) - - workspace_size = _finalize_cudnn_graph(graph) return _CudnnScoreModBwdGraphEntry( - graph=graph, - q=q, - k=k, - v=v, - output=output, - d_output=d_output, - stats=stats_tensor, - dq=dq, - dk=dk, - dv=dv, - score_mod_graph_tensors=score_mod_graph_tensors, - score_mod_bprop_graph_tensors=score_mod_bprop_graph_tensors, - workspace_size=workspace_size, + graph=entry["graph"], + q=entry["q"], + k=entry["k"], + v=entry["v"], + output=entry["o"], + d_output=entry["do"], + stats=entry["stats"], + dq=entry["dq"], + dk=entry["dk"], + dv=entry["dv"], + score_mod_graph_tensors=entry["aux"]["score_mod"], + score_mod_bprop_graph_tensors=entry["aux"]["score_mod_bprop"], + workspace_size=entry["workspace"], ) @@ -582,14 +479,12 @@ def _get_cudnn_score_mod_bwd_graph( score_mod_bprop_tensors, deterministic, ) - key = _cudnn_score_mod_bwd_cache_key(*build_args) - if key is None: - return _build_cudnn_score_mod_bwd_graph(*build_args) - entry = _cudnn_score_mod_graph_cache.get(key) - if entry is None: - entry = _build_cudnn_score_mod_bwd_graph(*build_args) - _cudnn_score_mod_graph_cache[key] = entry - return entry + return cudnn_pygraph.cached_graph( + _cudnn_score_mod_graph_cache, + _cudnn_score_mod_bwd_cache_key(*build_args), + lambda: _build_cudnn_score_mod_bwd_graph(*build_args), + device=query_layer.device, + ) class FusedAttentionWithScoreModFunc(torch.autograd.Function): diff --git a/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py new file mode 100644 index 00000000000..081d4b9f49e --- /dev/null +++ b/transformer_engine/pytorch/attention/dot_product_attention/frost_attention.py @@ -0,0 +1,832 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""cuDNN FROST attention. This feature is **experimental and subject to change**. + +FROST runs through cudnn-frontend's Python API and is registered behind the +CUDNN_FRONTEND_ENABLE_FROST_ENGINES=1 flag. Different from FusedAttnBackend.F16_arbitrary_seqlen +and FusedAttnBackend.FP8, FusedAttnBackend.FROST is Python only and requires the Python +installation of cudnn-frontend, not just its C++ header files. +""" + +from __future__ import annotations + +import os +from importlib.metadata import PackageNotFoundError, version as get_pkg_version +from typing import Any, Dict, Optional, Sequence, Tuple + +import torch +from packaging.version import InvalidVersion, Version as PkgVersion + +from . import cudnn_pygraph + +__all__ = [ + "is_frost_attention_available", + "is_frost_attention_supported", + "fused_attn_fwd", + "fused_attn_bwd", +] + + +# cudnn-frontend declares cutlass-dsl >= 4.6.2 but FROST enforces >= 4.7.0 at plan-build time. +# Below that floor every FROST engine declines silently and backend plans come back instead, so +# the selected plan is checked by NAME rather than trusting that the engine was used. +# Defined in cudnn_pygraph so the pin here and the bar in flex_attention cannot name different +# engines: one of them failing loudly and the other silently is exactly the drift to avoid. +_FROST_FWD_PLAN_TOKEN = cudnn_pygraph.FROST_FWD_PLAN_TOKEN +_FROST_BWD_PLAN_TOKEN = cudnn_pygraph.FROST_BWD_PLAN_TOKEN +_MIN_CUTLASS_DSL = PkgVersion("4.7.0") + +# 1.29.0 is the first release carrying the head_dim=512 backward. 1.28.0 ships the forward only, +# so without this check training would build a forward plan and raise on the first backward. +_MIN_CUDNN_FRONTEND = PkgVersion("1.29.0") + +_SUPPORTED_ARCHS = ((10, 0), (10, 3)) +# These mirror sdpa_bwd_sm100's declared envelope: d_envelope_floor=256, d={512}, +# d_pad_multiple=8. +_MAX_HEAD_DIM = 512 +_MIN_HEAD_DIM = 257 # below this the existing cuDNN/flash backends already serve the shape +# The engine pads head_dim to a multiple of 8, so 260 is in range but not servable. Declined +# here rather than failing later at plan selection. +_HEAD_DIM_MULTIPLE = 8 + +_BACKEND_NAME = "FrostAttention" +_availability: Optional[Tuple[bool, str]] = None +_PLAN_CACHE: dict = {} + + +def _import_cudnn_frontend(enable_frost_engines: bool = True): + """Import cuDNN Frontend with the FROST engines on, which is what this backend needs. + + The shared default is off; every use site here wants them. See + ``cudnn_pygraph.import_cudnn_frontend``. + """ + return cudnn_pygraph.import_cudnn_frontend(enable_frost_engines=enable_frost_engines) + + +def _diagonal_band_kwargs(cudnn, attn_mask_type: str, window: Tuple[int, int]) -> Dict[str, Any]: + """cuDNN sdpa kwargs for a TE (mask type, window): a diagonal alignment plus a band. + + Note the off-by-one. cuDNN's left bound counts the diagonal itself and TE's window_size does + not, so a window of w becomes a left bound of w + 1. Passing it through unconverted silently + drops one token of context per layer, which no shape-level test would catch. + + These kwargs are mutually exclusive with score_mod. cuDNN enforces that in the backward node + only ("Attention score mod enabled and hence other subgraphs are disabled"); its forward node + composes the two without complaint. Callers must still refuse the pair on both sides, because + forward and backward have to carry the same mask or the gradients belong to a different + attention than the output does. + """ + left, right = window + opts: Dict[str, Any] = {} + if attn_mask_type in ("causal", "causal_bottom_right") or right == 0: + opts["diagonal_alignment"] = ( + cudnn.diagonal_alignment.BOTTOM_RIGHT + if attn_mask_type == "causal_bottom_right" + else cudnn.diagonal_alignment.TOP_LEFT + ) + opts["diagonal_band_right_bound"] = 0 + if left != -1: + opts["diagonal_band_left_bound"] = left + 1 + return opts + + +def _pkg_version(name: str, module=None) -> Tuple[Optional[PkgVersion], Optional[str]]: + """A package's version, or None when it is absent or unparseable, with the raw string.""" + raw = None + 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 _frost_plan_args once a plan is built. + return _no("CUDNN_FRONTEND_ENABLE_FROST_ENGINES=0 disables the FROST engines") + if torch.cuda.get_device_capability() not in _SUPPORTED_ARCHS: + major, minor = torch.cuda.get_device_capability() + return _no(f"cuDNN FROST head_dim>256 kernels are SM100/SM103 only; found sm{major}{minor}") + try: + # Without the engines: this only reads a version, and enabling reorders plan selection + # process-wide even when the checks below decline. The use sites enable it. + _import_cudnn_frontend(enable_frost_engines=False) + except ImportError as exc: + return _no(f"nvidia-cudnn-frontend not importable: {exc}") + + # Decline only on positive evidence: a version below a floor, or a package absent outright. + # An unparseable version defers to _frost_plan_args, which checks the plan by name. + frontend, frontend_raw = _pkg_version("nvidia-cudnn-frontend", cudnn_pygraph.cudnn_module()) + if frontend is not None and frontend < _MIN_CUDNN_FRONTEND: + return _no( + 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 + + +# cuDNN expresses causal, bottom-right and sliding-window masking as one mechanism, a diagonal +# alignment plus a two-sided band, which is what _mask_options builds. Both alignments are needed: +# the p2p ring produces square tiles where they coincide, while all_gather trims KV so they differ. +_SUPPORTED_MASKS = ("no_mask", "causal", "causal_bottom_right") + +# Sliding window as TE spells it: (left, right), -1 meaning unbounded on that side. +_NO_WINDOW = (-1, -1) + + +def _window_pair(window_size) -> Tuple[int, int]: + """Normalise window_size to a (left, right) pair, declining anything that is not one. + + Its own function because _te_mask_spec unpacks the window before _mask_spec ever sees it, so + leaving these two guards inside _mask_spec left them dead on the fused path. Raised as + NotImplementedError so the selector declines, that being the only thing + is_frost_attention_supported catches; a TypeError would escape backend selection instead. + """ + if window_size is None: + return _NO_WINDOW + try: + window = tuple(window_size) + except TypeError: + raise NotImplementedError( + f"window_size must be a (left, right) pair; got {window_size!r}" + ) from None + if len(window) != 2 or not all(isinstance(bound, int) for bound in window): + raise NotImplementedError( + f"window_size must be a pair of ints (left, right); got {window!r}" + ) + return window + + +def _mask_spec(attn_mask_type: str, window_size=None): + """Validate a TE mask type and window, returning the hashable spec the plan is keyed on.""" + if attn_mask_type not in _SUPPORTED_MASKS: + raise NotImplementedError( + f"FROST attention supports attn_mask_type in {str(_SUPPORTED_MASKS)}; got" + f" {attn_mask_type!r}" + ) + window = _window_pair(window_size) + if window[0] < -1: + # cuDNN's left bound must be >= 1, so a left of -2 would build diagonal_band_left_bound=-1 + # and fail at plan build rather than declining here. + 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 _diagonal_band_kwargs(cudnn, attn_mask_type, window) + + +_SUPPORTED_QKV_FORMATS = ("bshd", "sbhd") + + +def _qkv_format_from_layout(qkv_layout: str) -> str: + """The single qkv_format a TE qkv_layout names, e.g. 'bshd_bshd_bshd' -> 'bshd'.""" + formats = { + "".join(c for c in part if c.isalpha()) + for part in qkv_layout.replace("paged_kv_", "").split("_") + } + if len(formats) != 1: + raise NotImplementedError( + f"FROST attention needs q, k and v in one format; got qkv_layout {qkv_layout!r}" + ) + qkv_format = formats.pop() + # Carried over from the permute helper this replaced, so the thd decline keeps its reason. + if qkv_format not in _SUPPORTED_QKV_FORMATS: + raise NotImplementedError( + f"FROST attention supports qkv_format in {_SUPPORTED_QKV_FORMATS}; got" + f" {qkv_format!r}. thd needs varlen support that is not implemented here." + ) + return qkv_format + + +def _te_mask_spec(attn_mask_type: str, window_size, bottom_right_diagonal: bool): + """Fold TE's (mask type, window, diagonal anchor) into the spec the plan is keyed on. + + TE carries the anchor in its own flag, so normalise it into the mask type before building the + band: diagonal_band_kwargs reads the anchor off the name, and taking it from the name alone + would quietly give a top-left band where the caller asked for bottom-right. + """ + if "padding" in attn_mask_type: + raise NotImplementedError( + f"FROST attention does not support a padding mask; got {attn_mask_type!r}" + ) + left, right = _window_pair(window_size) + if "causal" in attn_mask_type and right == -1: + right = 0 + if right == 0: + attn_mask_type = "causal_bottom_right" if bottom_right_diagonal else "causal" + else: + attn_mask_type = "no_mask" + return _mask_spec(attn_mask_type, (left, right)) + + +def _bottom_right_diagonal(attn_mask_type: str, bottom_right_diagonal) -> bool: + """Resolve the anchor flag the way cpp_extensions.fused_attn does. + + ``None`` means read it off the mask name. Only a direct call sees it unresolved. + """ + if bottom_right_diagonal is None: + return attn_mask_type in {"causal_bottom_right", "padding_causal_bottom_right"} + return bool(bottom_right_diagonal) + + +def _name_for(table, value, default=None): + """Reverse a cpp_extensions str-to-enum table.""" + for name, enum_value in table.items(): + if enum_value == value: + return name + return default + + +def is_frost_attention_supported(params) -> Tuple[int, str]: + """Whether this fused-attention config should run on the FROST sub-backend. + + Returns (sub-backend value, reject message) like tex.get_fused_attn_backend, so + get_attention_backend can fall through to it when the C++ backends decline. + + Deliberately does not probe availability: that imports cuDNN Frontend with the FROST engines + enabled, which changes the engine pool for every cuDNN consumer in the process, and this runs + for every attention config. get_attention_backend checks availability once at the end. + """ + # pylint: disable-next=import-outside-toplevel + from ...cpp_extensions.fused_attn import ( + AttnBiasType, + AttnMaskType, + FusedAttnBackend, + QKVFormat, + QKVLayout, + SoftmaxType, + TORCH_DType, + ) + + no_backend = int(FusedAttnBackend.No_Backend) + + if int(os.environ.get("NVTE_FROST_ATTN", "1")) == 0: + return no_backend, "FROST is disabled by NVTE_FROST_ATTN=0" + + # Checked per dimension: v has its own graph node and its own cache-key entry, so the range + # applies to each rather than to one standing in for both. + for name, head_dim in ( + ("head_dim_qk", params.head_dim_qk), + ("head_dim_v", params.head_dim_v), + ): + if not _MIN_HEAD_DIM <= head_dim <= _MAX_HEAD_DIM: + return no_backend, f"FROST covers head_dim in (256, 512]; got {name}={head_dim}" + if head_dim % _HEAD_DIM_MULTIPLE != 0: + return ( + no_backend, + f"FROST needs {name} to be a multiple of {_HEAD_DIM_MULTIPLE}; got {head_dim}", + ) + # The rule is the backward engine's, not this range's. sdpa_bwd_sm100 is the only f16 FROST + # backward on SM100/SM103 and leaves dqk_ge_dv unset, which cuDNN reads as requiring + # d_qk == d_v at every head dim. The forward serves 192/128 here, so only the backward binds. + if params.head_dim_qk != params.head_dim_v: + # Declined for both directions, since is_training follows module.training and eval() + # leaves autograd on, so an eval-mode call is no promise that no backward follows. + return ( + no_backend, + ( + "FROST requires symmetric head_dim in (256, 512]; got" + f" {params.head_dim_qk}/{params.head_dim_v}" + ), + ) + + qkv_dtype = TORCH_DType.get(params.qkv_dtype) + if qkv_dtype not in (torch.bfloat16, torch.float16): + return no_backend, f"FROST supports bf16/fp16; got {params.qkv_dtype}" + if params.dropout != 0.0: + return no_backend, "FROST does not support dropout" + if _name_for(AttnBiasType, params.bias_type) != "no_bias": + return no_backend, "FROST does not support attention bias" + if _name_for(SoftmaxType, params.softmax_type) != "vanilla": + return no_backend, "FROST only supports vanilla softmax" + if params.num_pages_k != 0 or params.num_pages_v != 0: + return no_backend, "FROST does not support paged KV" + if params.return_max_logit: + return no_backend, "FROST does not return max_logit" + if params.cuda_graph: + return no_backend, "FROST graphs are built lazily and cannot be captured" + if params.deterministic and params.is_training: + # The backward uses an atomic dQ accumulation whose order is not fixed, so repeat runs + # differ in the last bits. Nothing selects a deterministic variant, so decline instead. + return no_backend, "FROST does not have a deterministic backward" + + qkv_layout = _name_for(QKVLayout, params.qkv_layout) + if qkv_layout is None: + return no_backend, f"FROST got an unrecognised qkv_layout {params.qkv_layout}" + try: + qkv_format = _qkv_format_from_layout(qkv_layout) + except NotImplementedError as exc: + return no_backend, str(exc) + if qkv_format not in _SUPPORTED_QKV_FORMATS: + return ( + no_backend, + f"FROST supports qkv_format in {_SUPPORTED_QKV_FORMATS}; got {qkv_format}", + ) + # The kernels write O and dQKV with q's strides, so any format that differs from the input + # would need a copy the fused path does not make. Nothing asks for one today. + for name, value in ( + ("o_format", _name_for(QKVFormat, params.o_format)), + ("do_format", _name_for(QKVFormat, params.do_format)), + ("dqkv_layout", _name_for(QKVLayout, params.dqkv_layout)), + ): + if value is None: + continue + value = _qkv_format_from_layout(value) if name == "dqkv_layout" else value + if value != qkv_format: + return no_backend, f"FROST needs {name} to match qkv_format; got {value}/{qkv_format}" + + attn_mask_type = _name_for(AttnMaskType, params.attn_mask_type) + if attn_mask_type is None: + return no_backend, f"FROST got an unrecognised attn_mask_type {params.attn_mask_type}" + try: + mask_for_band, _ = _te_mask_spec( + attn_mask_type, + (params.window_size_left, params.window_size_right), + params.bottom_right_diagonal, + ) + except NotImplementedError as exc: + return no_backend, str(exc) + if ( + mask_for_band == "causal" + and "causal" not in attn_mask_type + and params.max_seqlen_q != params.max_seqlen_kv + ): + # Such a window takes its anchor only from bottom_right_diagonal, which defaults to + # top-left, while the all-gather ring measures its window bottom-right. Decline rather + # than guess which was meant. + return ( + no_backend, + ( + "FROST declines a right-bounded window on a non-causal mask with max_seqlen_q !=" + " max_seqlen_kv, where the diagonal anchor is ambiguous" + ), + ) + + return int(FusedAttnBackend.FROST), "" + + +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. + + ``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 agree with k on batch, heads and sequence length. + + head_dim is free, since v has its own graph node and its own cache-key entry. The other + three index the same KV positions as k by definition. + """ + if tuple(k.shape[:3]) != tuple(v.shape[:3]): + raise ValueError( + f"k and v must agree on batch, heads and seqlen; got {k.shape} and {v.shape}" + ) + + +def _head_dim_strides(shape: Sequence[int], ref_strides: Sequence[int]) -> list: + """Dense strides for ``shape`` in the memory order ``ref_strides`` describes.""" + order = sorted(range(len(shape)), key=lambda i: ref_strides[i], reverse=True) + strides = [0] * len(shape) + acc = 1 + for i in reversed(order): + strides[i] = acc + acc *= shape[i] + return strides + + +def _o_shape_stride(shape, d_v, ref_strides): + """Shape and strides for an O-shaped tensor: ``shape``'s layout carrying v's head_dim. + + Works in either space, since the head dim is last in both TE's bshd/sbhd and cuDNN's BHSD and + the rule only reorders by stride magnitude. Equal head dims keep the reference strides, which + preserves a caller's non-dense view. + """ + out = list(shape[:3]) + [d_v] + return out, (list(ref_strides) if d_v == shape[3] else _head_dim_strides(out, ref_strides)) + + +def _frost_plan_args(token: str, what: str) -> dict: + """Plan arguments that pin an engine whose name proves FROST was chosen. + + A too-old nvidia-cutlass-dsl makes the FROST engines decline silently, and in the forward an + ordinary engine may then build and compute something else. The pin turns that into a named + error at the first forward rather than a wrong number. + """ + + # 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_pygraph.cudnn_module())[1] or 'unknown'}" + f" (floor {_MIN_CUDNN_FRONTEND})," + f" nvidia-cutlass-dsl={_pkg_version('nvidia-cutlass-dsl')[1] or 'unknown'}" + f" (floor {_MIN_CUTLASS_DSL})." + ) + + cudnn = _import_cudnn_frontend() + return { + "heuristics": [cudnn.heur_mode.A], + "require_plan_token": token, + "not_found_hint": hint, + } + + +def _build_fwd(key, device) -> dict: + """Build (and JIT-compile) a forward graph. Expensive; always reached through the cache.""" + cudnn = _import_cudnn_frontend() + # deterministic is unused here: it selects a backward algorithm. Callers pass False for the + # forward so the two never split the forward cache. + _dev, (shq, qs, dtype), (shk, ks, _), (shv, vs, _), mask, scale, _deterministic = key + b, hq, sq = shq[0], shq[1], shq[2] + sho, o_stride = _o_shape_stride(shq, shv[3], qs) + return cudnn_pygraph.build_fwd( + dtype=dtype, + device=device, + backend_name=_BACKEND_NAME, + name="frost_fwd", + q=(shq, qs), + k=(shk, ks), + v=(shv, vs), + out=(sho, o_stride), # out: q's layout, v's head_dim + # the CP ring needs the LSE, and it is cheap + stats=([b, hq, sq, 1], [hq * sq, sq, 1, 1], cudnn.data_type.FLOAT), + attn_scale=scale, + sdpa_kwargs=_mask_options(cudnn, mask), + **_frost_plan_args(_FROST_FWD_PLAN_TOKEN, "forward"), + ) + + +def _build_bwd(key, device) -> dict: + """Build (and JIT-compile) a backward graph. Expensive; always reached through the cache.""" + cudnn = _import_cudnn_frontend() + _dev, (shq, qs, dtype), (shk, ks, _), (shv, vs, _), mask, scale, deterministic = key + io_dt = cudnn_pygraph.io_data_type(cudnn, dtype, backend_name=_BACKEND_NAME) + b, hq, sq = shq[0], shq[1], shq[2] + sho, o_stride = _o_shape_stride(shq, shv[3], qs) + return cudnn_pygraph.build_bwd( + dtype=dtype, + device=device, + backend_name=_BACKEND_NAME, + name="frost_bwd", + q=(shq, qs), + k=(shk, ks), + v=(shv, vs), + o=(sho, o_stride), + do=(sho, o_stride), + stats=([b, hq, sq, 1], [hq * sq, sq, 1, 1], cudnn.data_type.FLOAT), + # Each grad is declared with the layout of the tensor it differentiates. + dq=(None, qs, io_dt), + dk=(None, ks, io_dt), + dv=(None, vs, io_dt), + attn_scale=scale, + deterministic=deterministic, + sdpa_kwargs=_mask_options(cudnn, mask), + **_frost_plan_args(_FROST_BWD_PLAN_TOKEN, "backward"), + ) + + +def _cached(kind: str, key, device): + """Plan cache. See module docstring: building dominates executing even once the JIT is + cached, so this is required rather than an optimisation.""" + build = _build_fwd if kind == "fwd" else _build_bwd + return cudnn_pygraph.cached_graph( + _PLAN_CACHE, (kind,) + key, lambda: build(key, device), device=device + ) + + +def _validate_qkv(q, k, v, qkv_format): + """Check the tensors the graph will bind, and return their BHSD descriptions. + + These are not stylistic guards. ``execute`` binds raw pointers, so a tensor whose shape or + dtype disagrees with the node it is bound to is reinterpreted rather than rejected. The + context-parallel ring calls the backward outside autograd, so neither direction may assume + the other ran first. + + Rank is checked here because cuDNN cannot: _sdpa_validate inspects the descriptor this + module builds, and bhsd_dim_stride reads dims 0-3, so a rank-5 tensor is described as 4D. + The plan cache key carries strides but not rank, so such a tensor can even reuse a valid 4D + plan and be bound by pointer. The head-dim stride does reach the descriptor, so + _sdpa_validate rejects a non-unit one in graph.validate() and it is not repeated here. + """ + for name, tensor in (("q", q), ("k", k), ("v", v)): + if tensor.dim() != 4: + raise ValueError(f"{name} must be 4D; got shape {tuple(tensor.shape)}") + _check_dtype(name, tensor, q.dtype) + _check_kv_match(k, v) + qd, _ = _bhsd(q, qkv_format) + kd, _ = _bhsd(k, qkv_format) + vd, _ = _bhsd(v, qkv_format) + if kd[0] != qd[0] or kd[3] != qd[3]: + # The graph declares k and v with q's batch and head_dim, so a mismatch would bind a + # differently shaped buffer to that node and read the wrong elements silently. + raise ValueError(f"k must match q in batch and head_dim; got q {qd} and k {kd} in BHSD") + if qd[1] % kd[1] != 0: + raise ValueError(f"num_heads must be divisible by num_gqa_groups; got {qd[1]} and {kd[1]}") + return qd, kd, vd + + +def _bhsd(t: torch.Tensor, qkv_format: str): + """``t`` described in cuDNN's logical BHSD, without permuting it.""" + return cudnn_pygraph.bhsd_dim_stride(t, qkv_format, backend_name=_BACKEND_NAME) + + +def _key(q, k, v, qkv_format, mask, scale, deterministic=False): + """The plan cache key, structured per tensor rather than flattened. + + The builders destructure it by name, and the per-tensor fragment is the one flex keys on. + """ + + def described(t): + return cudnn_pygraph.tensor_key(t, qkv_format, backend_name=_BACKEND_NAME) + + return ( + # Built under whichever device was current, so it must not be reused on another. Matches + # the C++ fused-attn cache, which keys on device_id. + cudnn_pygraph.device_key(q.device), + described(q), + described(k), + described(v), + mask, + float(scale), + # 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_only(**unsupported): + """Raise if any feature the selector should have declined reached the kernels anyway.""" + for name, value in unsupported.items(): + if value: + raise NotImplementedError(f"FROST attention does not support {name}") + + +def fused_attn_fwd( + is_training, + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q, + k, + v, + fake_dtype, + fused_attention_backend, + attn_bias=None, + cu_seqlens_q_padded=None, + cu_seqlens_kv_padded=None, + page_table_k=None, + page_table_v=None, + s_quantizer=None, + o_quantizer=None, + attn_scale=None, + dropout=0.0, + fast_zero_fill=True, + qkv_layout="sbh3d", + o_format="sbhd", + qkv_scale_inv_format=None, + attn_bias_type="no_bias", + attn_mask_type="padding", + softmax_type="vanilla", + window_size=(-1, -1), + bottom_right_diagonal=None, + rng_gen=None, + softmax_offset=None, + return_max_logit=False, + cuda_graph=False, +): # pylint: disable=unused-argument + """FROST forward behind the cpp_extensions.fused_attn_fwd signature. + + Mirrors that signature so FusedAttnFunc and the context-parallel ring reach these kernels + without knowing which sub-backend they got. Returns (out, aux_ctx_tensors) with + aux_ctx_tensors = [softmax_lse, rng_state]; softmax_lse is [b, h, s] fp32 natural-log + logsumexp, which is what the ring correction consumes. + + cu_seqlens and the padded variants are ignored: they carry thd offsets, and thd is declined + at selection. + """ + _frost_only( + dropout=dropout != 0.0, + attention_bias=attn_bias_type != "no_bias", + paged_kv=page_table_k is not None or page_table_v is not None, + fp8=s_quantizer is not None or o_quantizer is not None, + sink_attention=softmax_type != "vanilla", + max_logit=return_max_logit, + cuda_graph_capture=cuda_graph, + ) + qkv_format = _qkv_format_from_layout(qkv_layout) + if o_format != qkv_format: + raise NotImplementedError( + f"FROST attention needs o_format to match qkv_format; got {o_format}/{qkv_format}" + ) + # _te_mask_spec validates as it normalises, so what it returns is the spec the plan keys on. + mask = _te_mask_spec( + attn_mask_type, window_size, _bottom_right_diagonal(attn_mask_type, bottom_right_diagonal) + ) + + q, k, v = q.contiguous(), k.contiguous(), v.contiguous() + qd, _, vd = _validate_qkv(q, k, v, qkv_format) + scale = attn_scale if attn_scale is not None else qd[3] ** -0.5 + entry = _cached("fwd", _key(q, k, v, qkv_format, mask, scale), q.device) + tq, tk, tv, tout, tlse = (entry[n] for n in ("q", "k", "v", "out", "stats")) + + # Allocated per call so concurrent uses cannot alias; the cache holds only the plan. + # empty_strided, not empty_like: the latter does not preserve an arbitrary permuted stride. + # The output takes the caller's format with v's head_dim, so nothing is converted on the way + # out; the LSE is BHSD whatever the inputs were, which is what the ring correction expects. + out_shape, out_stride = _o_shape_stride(q.shape, vd[3], q.stride()) + out = torch.empty_strided(out_shape, out_stride, device=q.device, dtype=q.dtype) + lse = torch.empty(qd[0], qd[1], qd[2], 1, device=q.device, dtype=torch.float32) + cudnn_pygraph.execute_graph( + entry["graph"], + {tq: q, tk: k, tv: v, tout: out, tlse: lse}, + entry["workspace"], + q.device, + backend_name=_BACKEND_NAME, + ) + # A real tensor, not None: it is saved for backward and handed to the activation offload + # hooks, neither of which accepts None. FROST has no dropout, so nothing reads it. + rng_state = torch.empty(2, dtype=torch.int64, device=q.device) + return out, [lse.squeeze(-1), rng_state] + + +def fused_attn_bwd( + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q, + k, + v, + o, + d_o, + fake_dtype, + aux_ctx_tensors, + fused_attention_backend, + cu_seqlens_q_padded=None, + cu_seqlens_kv_padded=None, + s_quantizer=None, + dp_quantizer=None, + dqkv_quantizer=None, + attn_scale=None, + dropout=0.0, + fast_zero_fill=True, + qkv_layout="sbh3d", + o_format="sbhd", + do_format="sbhd", + dqkv_layout="sbh3d", + qkv_scale_inv_format=None, + do_scale_inv_format=None, + attn_bias_type="no_bias", + attn_mask_type="padding", + softmax_type="vanilla", + window_size=(-1, -1), + bottom_right_diagonal=None, + deterministic=False, + cuda_graph=False, +): # pylint: disable=unused-argument + """FROST backward behind the cpp_extensions.fused_attn_bwd signature. + + Returns (dq, dk, dv, dbias) with dbias always None, matching what the fused path returns for + a no_bias config. + """ + _frost_only( + dropout=dropout != 0.0, + attention_bias=attn_bias_type != "no_bias", + fp8=s_quantizer is not None or dqkv_quantizer is not None, + sink_attention=softmax_type != "vanilla", + cuda_graph_capture=cuda_graph, + ) + qkv_format = _qkv_format_from_layout(qkv_layout) + # o and dO used to carry their own format into the permute; they now share qkv_format, so a + # divergence would silently describe them with the wrong strides. The selector already + # declines it, but this is reached directly too. + # dqkv_layout joins them: the grads go back through qkv_format, so a divergence would return + # them in a layout the caller did not ask for, with nothing raised. + for name, fmt in ( + ("o_format", o_format), + ("do_format", do_format), + ("dqkv_layout", _qkv_format_from_layout(dqkv_layout)), + ): + if fmt != qkv_format: + raise NotImplementedError( + f"FROST attention needs {name} to match qkv_format; got {fmt}/{qkv_format}" + ) + mask = _te_mask_spec( + attn_mask_type, window_size, _bottom_right_diagonal(attn_mask_type, bottom_right_diagonal) + ) + softmax_lse = aux_ctx_tensors[0] + + q, k, v = q.contiguous(), k.contiguous(), v.contiguous() + o, d_o = o.contiguous(), d_o.contiguous() + qd, _, vd = _validate_qkv(q, k, v, qkv_format) + for name, tensor in (("o", o), ("d_o", d_o)): + _check_dtype(name, tensor, q.dtype) + o_shape, o_stride = _o_shape_stride(q.shape, vd[3], q.stride()) + for name, tensor in (("o", o), ("d_o", d_o)): + if list(tensor.shape) != o_shape: + raise ValueError(f"{name} must be shaped {o_shape}; got {list(tensor.shape)}") + if softmax_lse.dtype != torch.float32: + raise ValueError(f"softmax_lse must be fp32; got {softmax_lse.dtype}") + # Compared against the BHSD description, not q's own shape: the LSE is always [b, h, s] + # whatever format the tensors arrived in. + # The whole shape, not just the leading dims: the stats node is declared [b, h, s, 1], so a + # trailing dim of any other size would be bound by pointer and read with strides that do not + # describe it. Same hole the rank check on q/k/v closes. + lse_shape = tuple(softmax_lse.shape) + if lse_shape not in (tuple(qd[:3]), tuple(qd[:3]) + (1,)): + raise ValueError( + f"softmax_lse must be [b, h, s] or [b, h, s, 1] matching q; got {lse_shape}" + ) + + scale = attn_scale if attn_scale is not None else qd[3] ** -0.5 + entry = _cached("bwd", _key(q, k, v, qkv_format, mask, scale, deterministic), q.device) + h = entry + + if softmax_lse.dim() == 3: + softmax_lse = softmax_lse.unsqueeze(-1) + softmax_lse = softmax_lse.contiguous() + + # The graph expects o and dO in the layout the forward wrote, and dO comes from autograd + # with strides we do not control, so restride rather than silently reading the wrong elements. + def _as(t, stride): + if list(t.stride()) == list(stride): + return t + buf = torch.empty_strided(t.shape, stride, device=t.device, dtype=t.dtype) + buf.copy_(t) + return buf + + o, d_o = _as(o, o_stride), _as(d_o, o_stride) + dq = torch.empty_strided(q.shape, q.stride(), device=q.device, dtype=q.dtype) + dk = torch.empty_strided(k.shape, k.stride(), device=k.device, dtype=k.dtype) + dv = torch.empty_strided(v.shape, v.stride(), device=v.device, dtype=v.dtype) + cudnn_pygraph.execute_graph( + entry["graph"], + { + h["q"]: q, + h["k"]: k, + h["v"]: v, + h["o"]: o, + h["do"]: d_o, + h["stats"]: softmax_lse, + h["dq"]: dq, + h["dk"]: dk, + h["dv"]: dv, + }, + entry["workspace"], + q.device, + backend_name=_BACKEND_NAME, + ) + return dq, dk, dv, None diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index fbbd899aa22..ca03b427f28 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -449,9 +449,17 @@ def _get_fused_attn_backend(**fused_attn_kwargs): graph break because it is baked into the graph as a literal, while an enum member comes out of the reconstruction corrupted (see the cast at the call site, which restores the enum).""" - fused_attention_backend, reject_message = tex.get_fused_attn_backend( - FusedAttentionParams(**fused_attn_kwargs) - ) + params = FusedAttentionParams(**fused_attn_kwargs) + fused_attention_backend, reject_message = tex.get_fused_attn_backend(params) + if fused_attention_backend == FusedAttnBackend.No_Backend: + from .frost_attention import ( # pylint: disable=import-outside-toplevel + is_frost_attention_supported, + ) + + frost_backend, frost_reject = is_frost_attention_supported(params) + if frost_backend != FusedAttnBackend.No_Backend: + return int(frost_backend), frost_reject + reject_message = f"{reject_message.rstrip('. ')}. {frost_reject}" return int(fused_attention_backend), reject_message @@ -1867,6 +1875,16 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt if use_flash_attention_4 and not FlashAttentionUtils.v4_is_installed: use_flash_attention_4 = False use_flash_attention = use_flash_attention_2 or use_flash_attention_3 or use_flash_attention_4 + if use_fused_attention and fused_attention_backend == FusedAttnBackend.FROST.value: + from .frost_attention import ( # pylint: disable=import-outside-toplevel + is_frost_attention_available, + ) + + frost_available, frost_reason = is_frost_attention_available() + if not frost_available: + logger.debug("Disabling FusedAttention: %s", frost_reason) + use_fused_attention = False + fused_attention_backend = None available_backends = [use_flash_attention, use_fused_attention, use_unfused_attention] if use_flash_attention_2: flash_attention_backend = FlashAttentionUtils.version diff --git a/transformer_engine/pytorch/cpp_extensions/fused_attn.py b/transformer_engine/pytorch/cpp_extensions/fused_attn.py index 9a33df7634b..9dd93f3ce1a 100644 --- a/transformer_engine/pytorch/cpp_extensions/fused_attn.py +++ b/transformer_engine/pytorch/cpp_extensions/fused_attn.py @@ -103,7 +103,7 @@ class FusedAttnBackend(IntEnum): """Fused attention sub-backends. This is the canonical fused-attention backend enum for - ``transformer_engine.pytorch``. It mirrors the backend + ``transformer_engine.pytorch``. It mirrors every member of the backend ``transformer_engine_torch.NVTE_Fused_Attn_Backend`` (pybind11) enum value-for-value, and instances of the two enums compare equal when they share the same integer value. Unlike the pybind enum, a plain-python @@ -119,6 +119,9 @@ class FusedAttnBackend(IntEnum): No_Backend = int(NVTE_Fused_Attn_Backend.NVTE_No_Backend) F16_arbitrary_seqlen = int(NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen) FP8 = int(NVTE_Fused_Attn_Backend.NVTE_FP8) + # Python-only: FROST runs through the cuDNN Frontend python API, so it has no + # NVTE_Fused_Attn_Backend counterpart and never reaches pybind. + FROST = 3 @classmethod def cast( @@ -153,9 +156,14 @@ def __hash__(self) -> int: return int.__hash__(self) +# Members with no C++ counterpart; excluded from the sync check below. +_PYTHON_ONLY_FUSED_ATTN_BACKENDS = frozenset({FusedAttnBackend.FROST}) + # Fail fast at import time if a new enumerator is added on the C++ side # without being mirrored above. -assert {f"NVTE_{m.name}" for m in FusedAttnBackend} == set(NVTE_Fused_Attn_Backend.__members__), ( +assert { + f"NVTE_{m.name}" for m in FusedAttnBackend if m not in _PYTHON_ONLY_FUSED_ATTN_BACKENDS +} == set(NVTE_Fused_Attn_Backend.__members__), ( "FusedAttnBackend in python is out of sync with" " transformer_engine_torch.NVTE_Fused_Attn_Backend defined on the C++ side." " Please make sure TE C++ and python are in sync." @@ -345,6 +353,45 @@ def fused_attn_fwd( # Accept the pybind enum for backward compatibility. fused_attention_backend = FusedAttnBackend.cast(fused_attention_backend) + if fused_attention_backend == FusedAttnBackend["FROST"]: + # Imported here so a process that never selects FROST never imports cuDNN Frontend. + # pylint: disable-next=import-outside-toplevel + from ..attention.dot_product_attention import frost_attention + + return frost_attention.fused_attn_fwd( + is_training, + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q, + k, + v, + fake_dtype, + fused_attention_backend, + attn_bias=attn_bias, + cu_seqlens_q_padded=cu_seqlens_q_padded, + cu_seqlens_kv_padded=cu_seqlens_kv_padded, + page_table_k=page_table_k, + page_table_v=page_table_v, + s_quantizer=s_quantizer, + o_quantizer=o_quantizer, + attn_scale=attn_scale, + dropout=dropout, + fast_zero_fill=fast_zero_fill, + qkv_layout=qkv_layout, + o_format=o_format, + qkv_scale_inv_format=qkv_scale_inv_format, + attn_bias_type=attn_bias_type, + attn_mask_type=attn_mask_type, + softmax_type=softmax_type, + window_size=window_size, + bottom_right_diagonal=bottom_right_diagonal, + rng_gen=rng_gen, + softmax_offset=softmax_offset, + return_max_logit=return_max_logit, + cuda_graph=cuda_graph, + ) if fused_attention_backend == FusedAttnBackend["No_Backend"]: raise ValueError( "Fused attention does not support this input combination:" @@ -600,6 +647,46 @@ def fused_attn_bwd( # Accept the pybind enum for backward compatibility. fused_attention_backend = FusedAttnBackend.cast(fused_attention_backend) + if fused_attention_backend == FusedAttnBackend["FROST"]: + # See the matching branch in fused_attn_fwd. + # pylint: disable-next=import-outside-toplevel + from ..attention.dot_product_attention import frost_attention + + return frost_attention.fused_attn_bwd( + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q, + k, + v, + o, + d_o, + fake_dtype, + aux_ctx_tensors, + fused_attention_backend, + cu_seqlens_q_padded=cu_seqlens_q_padded, + cu_seqlens_kv_padded=cu_seqlens_kv_padded, + s_quantizer=s_quantizer, + dp_quantizer=dp_quantizer, + dqkv_quantizer=dqkv_quantizer, + attn_scale=attn_scale, + dropout=dropout, + fast_zero_fill=fast_zero_fill, + qkv_layout=qkv_layout, + o_format=o_format, + do_format=do_format, + dqkv_layout=dqkv_layout, + qkv_scale_inv_format=qkv_scale_inv_format, + do_scale_inv_format=do_scale_inv_format, + attn_bias_type=attn_bias_type, + attn_mask_type=attn_mask_type, + softmax_type=softmax_type, + window_size=window_size, + bottom_right_diagonal=bottom_right_diagonal, + deterministic=deterministic, + cuda_graph=cuda_graph, + ) if fused_attention_backend == FusedAttnBackend["No_Backend"]: raise ValueError( "Fused attention backward does not support this input combination:"