From 69b4ffb1a8c4ea00599d4d611e1102534ecc040b Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Mon, 17 Aug 2026 14:20:36 -0700 Subject: [PATCH 1/9] feat: thread tanh logit softcapping through FlashAttention (FA2, opt-in FA3) Add a user `softcap` value (tanh logit softcapping, `softcap*tanh(x/softcap)`) to DotProductAttention so models like Gemma2 can run on the fused flash path instead of an unfused/FlexAttention kernel. - Add `softcap` to DotProductAttention (init+forward) and AttentionParams; thread it into the FA2 non-CP kwargs and all three context-parallel autograd functions (forward + ctx-saved backward). softcap=0.0 reproduces prior behavior. - get_attention_backend: when softcap != 0, disable FusedAttention/unfused and steer to FA2 -- disable FA3/FA4, and disable FA2 < 2.6.0 -- so the cap is never silently dropped (FA2 < 2.6.0) or hit at runtime as NotImplementedError (FA3/FA4). Also disable FA3 under context parallelism (its CP path hard-rejects nonzero softcap) so CP+softcap steers to FA2, which supports it, instead of crashing. - FA3 softcap opt-in: NVTE_FA3_SOFTCAP=1, Hopper (sm90) hd<=256, non-CP only, gated on a fail-closed signature probe (fa3_supports_softcap). Forward threads softcap into fa_3_optional_forward_kwargs; the existing Hopper autograd function carries it into backward automatically. Default off; unchanged behavior steers to FA2. - ONNX export: fail loudly (assert) rather than silently drop softcap -- export unconditionally force-selects UnfusedDotProductAttention, which has no softcap support, so this previously exported models with softcapping silently omitted. - Tests: test_softcap.py (FA2 fwd/bwd parity vs pure-PyTorch reference), wired into qa/L0_pytorch_unittest/test.sh. FA4 softcap opt-in is deliberately NOT included here -- see follow-up PR. On Blackwell (SM100), FA4's dedicated head_dim=256 forward kernel has no score_mod support at all (kernel constructor asserts `score_mod is None`), so there is currently no FA4 kernel path this could opt into; adding the scaffolding now would just be inert code with nothing to exercise. Addresses review findings: CP+FA3 softcap selection crash, ONNX silent drop, and the missing CI wiring for test_softcap.py. Signed-off-by: Nitin Vegesna Co-Authored-By: Claude Opus 4.8 (1M context) --- qa/L0_pytorch_unittest/test.sh | 1 + tests/pytorch/attention/test_softcap.py | 149 ++++++++++++++++++ .../dot_product_attention/backends.py | 38 +++++ .../dot_product_attention/context_parallel.py | 23 ++- .../dot_product_attention.py | 26 +++ .../attention/dot_product_attention/utils.py | 47 ++++++ 6 files changed, 278 insertions(+), 6 deletions(-) create mode 100644 tests/pytorch/attention/test_softcap.py diff --git a/qa/L0_pytorch_unittest/test.sh b/qa/L0_pytorch_unittest/test.sh index 14a5f4fe3d..482642e5e1 100644 --- a/qa/L0_pytorch_unittest/test.sh +++ b/qa/L0_pytorch_unittest/test.sh @@ -67,6 +67,7 @@ NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE=1 NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 python python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_linear_mxfp8_attention.xml $TE_PATH/tests/pytorch/attention/test_linear_mxfp8_attention.py || test_fail "test_linear_mxfp8_attention.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_fused_mla_q_uproj.xml $TE_PATH/tests/pytorch/attention/test_fused_mla_q_uproj.py || test_fail "test_fused_mla_q_uproj.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_kv_cache.xml $TE_PATH/tests/pytorch/attention/test_kv_cache.py || test_fail "test_kv_cache.py" +python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_softcap.xml $TE_PATH/tests/pytorch/attention/test_softcap.py || test_fail "test_softcap.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_cu_seqlens_cache.xml $TE_PATH/tests/pytorch/attention/test_cu_seqlens_cache.py || test_fail "test_cu_seqlens_cache.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_hf_integration.xml $TE_PATH/tests/pytorch/test_hf_integration.py || test_fail "test_hf_integration.py" export NVTE_TEST_CHECKPOINT_ARTIFACT_PATH=$TE_PATH/artifacts/tests/pytorch/test_checkpoint diff --git a/tests/pytorch/attention/test_softcap.py b/tests/pytorch/attention/test_softcap.py new file mode 100644 index 0000000000..b0b94bb4e8 --- /dev/null +++ b/tests/pytorch/attention/test_softcap.py @@ -0,0 +1,149 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""Isolation numerics test for tanh logit softcapping in DotProductAttention. + +The reference implements softcapping in pure PyTorch: + + scores = (Q @ K^T) * scale + scores = softcap * tanh(scores / softcap) # only when softcap != 0.0 + scores = scores + mask + attn = softmax(scores) + out = attn @ V + +and is compared against ``DotProductAttention(..., softcap=...)`` forced onto the +FlashAttention backend, for both the forward output and the input gradients +(dQ/dK/dV obtained via autograd). +""" + +import sys +import pathlib + +import pytest +import torch +from packaging.version import Version as PkgVersion + +from transformer_engine.pytorch import DotProductAttention +from transformer_engine.pytorch.attention.dot_product_attention import _attention_backends + +_current_file = pathlib.Path(__file__).resolve() +sys.path = [str(_current_file.parent.parent)] + sys.path +from utils import reset_rng_states # pylint: disable=wrong-import-position + + +def _flash_attn_2_6_available() -> bool: + """Whether flash-attn >= 2.6.0 (the first version exposing ``softcap``) is installed.""" + try: + import flash_attn # pylint: disable=import-outside-toplevel + except ImportError: + return False + return PkgVersion(flash_attn.__version__) >= PkgVersion("2.6.0") + + +# Softcapping through DotProductAttention is only wired through the FlashAttention 2 +# backend (>= 2.6.0), and requires CUDA tensors. +pytestmark = [ + pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required."), + pytest.mark.skipif( + not _flash_attn_2_6_available(), reason="flash-attn >= 2.6.0 is required." + ), +] + + +def _force_flash_backend() -> None: + """Force DotProductAttention to select the FlashAttention backend.""" + import os # pylint: disable=import-outside-toplevel + + os.environ["NVTE_FLASH_ATTN"] = "1" + os.environ["NVTE_FUSED_ATTN"] = "0" + os.environ["NVTE_UNFUSED_ATTN"] = "0" + _attention_backends["backend_selection_requires_update"] = True + + +def _reference_attention(q, k, v, scale, softcap, causal): + """Pure-PyTorch reference for softcapped scaled dot product attention. + + q, k, v are in ``bshd`` layout. GQA is supported: ``k``/``v`` may have fewer + heads than ``q``. + """ + # bshd -> bhsd + qt = q.transpose(1, 2).float() + kt = k.transpose(1, 2).float() + vt = v.transpose(1, 2).float() + + num_heads = qt.shape[1] + num_gqa_groups = kt.shape[1] + if num_heads != num_gqa_groups: + assert num_heads % num_gqa_groups == 0 + repeats = num_heads // num_gqa_groups + kt = kt.repeat_interleave(repeats, dim=1) + vt = vt.repeat_interleave(repeats, dim=1) + + scores = torch.matmul(qt, kt.transpose(-2, -1)) * scale + if softcap != 0.0: + scores = softcap * torch.tanh(scores / softcap) + if causal: + sq, skv = scores.shape[-2], scores.shape[-1] + mask = torch.triu( + torch.ones(sq, skv, dtype=torch.bool, device=scores.device), + diagonal=1 + skv - sq, + ) + scores = scores.masked_fill(mask, float("-inf")) + attn = torch.softmax(scores, dim=-1) + out = torch.matmul(attn, vt) + # bhsd -> bshd + return out.transpose(1, 2) + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +@pytest.mark.parametrize("softcap", [0.0, 50.0]) +@pytest.mark.parametrize("num_gqa_groups", [4, 2]) +@pytest.mark.parametrize("causal", [False, True]) +def test_softcap_numerics(dtype, softcap, num_gqa_groups, causal): + """FlashAttention softcap forward + grads match a pure-PyTorch reference. + + ``softcap == 0.0`` additionally proves that softcapping is a no-op relative to + the plain (no-softcap) reference, i.e. today's behavior is reproduced exactly. + """ + reset_rng_states() + + batch_size = 2 + max_seqlen = 32 + num_heads = 4 + head_dim = 64 + scale = 1.0 / (head_dim**0.5) + + q_shape = (batch_size, max_seqlen, num_heads, head_dim) + kv_shape = (batch_size, max_seqlen, num_gqa_groups, head_dim) + + q = (0.5 * torch.randn(q_shape, dtype=dtype, device="cuda")).requires_grad_() + k = (0.5 * torch.randn(kv_shape, dtype=dtype, device="cuda")).requires_grad_() + v = (0.5 * torch.randn(kv_shape, dtype=dtype, device="cuda")).requires_grad_() + q_ref, k_ref, v_ref = [x.detach().clone().requires_grad_() for x in (q, k, v)] + + grad_output = torch.randn(q_shape, dtype=dtype, device="cuda") + + _force_flash_backend() + dpa = DotProductAttention( + num_heads, + head_dim, + num_gqa_groups=num_gqa_groups, + qkv_format="bshd", + attn_mask_type="causal" if causal else "no_mask", + softmax_scale=scale, + softcap=softcap, + layer_number=1, + ).to(dtype=dtype, device="cuda") + + out = dpa(q, k, v) + out.backward(grad_output) + + out_ref = _reference_attention(q_ref, k_ref, v_ref, scale, softcap, causal) + out_ref.backward(grad_output.float()) + + atol, rtol = (2e-2, 2e-2) if dtype == torch.float16 else (3.5e-2, 3.5e-2) + + torch.testing.assert_close(out.float(), out_ref.float(), atol=atol, rtol=rtol) + torch.testing.assert_close(q.grad.float(), q_ref.grad.float(), atol=atol, rtol=rtol) + torch.testing.assert_close(k.grad.float(), k_ref.grad.float(), atol=atol, rtol=rtol) + torch.testing.assert_close(v.grad.float(), v_ref.grad.float(), atol=atol, rtol=rtol) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index 8a219a6a4d..12a2ed1492 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -7,6 +7,7 @@ from contextlib import nullcontext from importlib.metadata import version as get_pkg_version from importlib.metadata import PackageNotFoundError +import inspect import os from typing import Any, Callable, Dict, List, Optional, Tuple, Union import warnings @@ -165,6 +166,19 @@ fa_utils.set_flash_attention_3_params() + # Probe whether this FA3 build exposes a `softcap` parameter on BOTH entry points. FA3's Hopper + # (sm90) kernels DO implement tanh logit softcapping in fwd AND bwd (dedicated + # flash_{fwd,bwd}_hdim256_bf16_softcap_sm90 instantiations, off only behind a compile-time + # DISABLE_SOFTCAP flag), so this is a mature path. Still fail-closed and additionally + # gated on opt-in (NVTE_FA3_SOFTCAP) + head_dim <= 256 in get_attention_backend. + try: + fa_utils.fa3_supports_softcap = ( + "softcap" in inspect.signature(flash_attn_func_v3).parameters + and "softcap" in inspect.signature(flash_attn_varlen_func_v3).parameters + ) + except (ValueError, TypeError): + fa_utils.fa3_supports_softcap = False + # Try to import Flash Attention v4 try: fa_utils.fa4_version = PkgVersion(get_pkg_version("flash-attn-4")) @@ -885,6 +899,7 @@ def forward( max_seqlen_kv: Optional[int] = None, attn_mask_type: str = "causal", window_size: Optional[Tuple[int, int]] = None, + softcap: float = 0.0, alibi_slopes: Optional[torch.Tensor] = None, cp_group: Optional[Union[dist_group_type, List[dist_group_type]]] = None, cp_global_ranks: List[int] = None, @@ -1100,6 +1115,11 @@ def forward( assert ( alibi_slopes is None ), "Alibi slope bias addition is not supported with context parallelism." + if use_flash_attn_3 and softcap != 0.0: + raise NotImplementedError( + "softcap is not supported by the FlashAttention 3 backend in context " + "parallel. Please use FlashAttention 2 (>= 2.6.0) for softcap support." + ) with self.attention_dropout_ctx(): output = attn_forward_func_with_cp( self.training, @@ -1130,6 +1150,7 @@ def forward( attn_mask_type=attn_mask_type, deterministic=self.deterministic, window_size=window_size, + softcap=softcap, quantizers=quantizers, pad_between_seqs=pad_between_seqs, use_flash_attn_3=use_flash_attn_3, @@ -1215,6 +1236,8 @@ def forward( fa_optional_forward_kwargs["alibi_slopes"] = alibi_slopes if fa_utils.v2_4_1_plus: fa_optional_forward_kwargs["deterministic"] = self.deterministic + if fa_utils.v2_6_0_plus: + fa_optional_forward_kwargs["softcap"] = softcap if inference_params is not None: # use block_table kwarg to support thd_2bshd for non-paged fa_optional_forward_kwargs["block_table"] = ( @@ -1235,9 +1258,24 @@ def forward( **fa_optional_forward_kwargs, ) else: + # Fail-loud net: get_attention_backend only keeps FA3 for softcap on a + # softcap-capable build (signature probe) + opt-in (NVTE_FA3_SOFTCAP) + Hopper + # (FA3 is sm90-only upstream) + head_dim <= 256. If FA3 is still reached with + # softcap while the build lacks support (force-selected / regressed path), raise + # rather than silently drop the cap. The non-CP FA3 entry points + # (flash_attn_func_v3 / flash_attn_varlen_func_v3) are self-contained autograd + # functions, so threading `softcap` into the forward call also drives the + # matching FA3 softcap backward kernel. (CP + FA3 + softcap stays blocked above.) + if softcap != 0.0 and not fa_utils.fa3_supports_softcap: + raise NotImplementedError( + "softcap is not supported by the installed FlashAttention 3 build. " + "Please use FlashAttention 2 (>= 2.6.0) for softcap support." + ) fa_3_optional_forward_kwargs = {} fa_3_optional_forward_kwargs["window_size"] = window_size fa_3_optional_forward_kwargs["num_splits"] = num_splits + if softcap != 0.0 and fa_utils.fa3_supports_softcap: + fa_3_optional_forward_kwargs["softcap"] = softcap if pad_between_seqs: fa_3_optional_forward_kwargs["seqused_q"] = ( cu_seqlens_q[1:] - cu_seqlens_q[:-1] 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 ea89ca97eb..3484d4e9cd 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1387,6 +1387,7 @@ def forward( deterministic, use_fused_attention, return_max_logit, + softcap, fp8, fp8_meta, cp_group, @@ -1664,7 +1665,7 @@ def forward( if fa_utils.v2_5_7_plus and qkv_format == "thd": fa_forward_kwargs["block_table"] = None if fa_utils.v2_6_0_plus: - fa_forward_kwargs["softcap"] = 0.0 + fa_forward_kwargs["softcap"] = softcap # set up inputs for forward q_inputs = [None, None] @@ -2156,6 +2157,7 @@ def forward( ctx.attn_bias_type = attn_bias_type ctx.attn_bias_shape = None if attn_bias is None else attn_bias.shape ctx.deterministic = deterministic + ctx.softcap = softcap ctx.use_fused_attention = use_fused_attention ctx.pad_between_seqs = pad_between_seqs ctx.softmax_lse_in_packed_format = softmax_lse_in_packed_format @@ -2454,7 +2456,7 @@ def backward(ctx, dout, *_args): if fa_utils.v2_4_1_plus: fa_backward_kwargs["deterministic"] = ctx.deterministic if fa_utils.v2_6_0_plus: - fa_backward_kwargs["softcap"] = 0.0 + fa_backward_kwargs["softcap"] = ctx.softcap send_recv_reqs = [] for i in range(cp_size): @@ -2970,6 +2972,7 @@ def backward(ctx, dout, *_args): None, None, None, + None, ) @@ -3047,6 +3050,7 @@ def forward( deterministic, use_fused_attention, return_max_logit, + softcap, window_size, cp_group, cp_stream, @@ -3128,7 +3132,7 @@ def forward( if fa_utils.v2_5_7_plus and qkv_format == "thd": fa_forward_kwargs["block_table"] = None if fa_utils.v2_6_0_plus: - fa_forward_kwargs["softcap"] = 0.0 + fa_forward_kwargs["softcap"] = softcap qkv_layout = qkv_format + "_" + qkv_format + "_" + qkv_format @@ -3644,6 +3648,7 @@ def forward( ctx.attn_bias_type = attn_bias_type ctx.attn_mask_type = attn_mask_type ctx.deterministic = deterministic + ctx.softcap = softcap ctx.use_fused_attention = use_fused_attention ctx.use_flash_attn_3 = use_flash_attn_3 ctx.pad_between_seqs = pad_between_seqs @@ -3840,7 +3845,7 @@ def backward(ctx, dout, *_args): if fa_utils.v2_4_1_plus: fa_backward_kwargs["deterministic"] = ctx.deterministic if fa_utils.v2_6_0_plus: - fa_backward_kwargs["softcap"] = 0.0 + fa_backward_kwargs["softcap"] = ctx.softcap local_seq_chunk_ids = [rank, 2 * cp_size - rank - 1] for i in range(len(local_seq_chunk_ids) + 1): @@ -4164,6 +4169,7 @@ def backward(ctx, dout, *_args): None, None, None, + None, ) @@ -4195,6 +4201,7 @@ def forward( deterministic, use_fused_attention, return_max_logit, + softcap, window_size, fp8, fp8_meta, @@ -4284,7 +4291,7 @@ def forward( if fa_utils.v2_5_7_plus and qkv_format == "thd": fa_forward_kwargs["block_table"] = None if fa_utils.v2_6_0_plus: - fa_forward_kwargs["softcap"] = 0.0 + fa_forward_kwargs["softcap"] = softcap assert isinstance(k, q.__class__) and isinstance( v, q.__class__ @@ -4585,6 +4592,7 @@ def forward( ctx.attn_mask_type = attn_mask_type ctx.attn_bias_type = attn_bias_type ctx.deterministic = deterministic + ctx.softcap = softcap ctx.window_size = window_size ctx.use_fused_attention = use_fused_attention ctx.fp8_meta = fp8_meta @@ -4725,7 +4733,7 @@ def backward(ctx, dout, *_args): if fa_utils.v2_4_1_plus: fa_backward_kwargs["deterministic"] = ctx.deterministic if fa_utils.v2_6_0_plus: - fa_backward_kwargs["softcap"] = 0.0 + fa_backward_kwargs["softcap"] = ctx.softcap dq_fp8, dk_fp8, dv_fp8 = None, None, None if ctx.use_fused_attention: @@ -4916,6 +4924,7 @@ def backward(ctx, dout, *_args): None, None, None, + None, d_softmax_offset, None, ) @@ -4945,6 +4954,7 @@ def attn_forward_func_with_cp( deterministic=False, use_fused_attention=False, window_size=None, + softcap=0.0, fp8=False, fp8_meta=None, quantizers=None, @@ -5091,6 +5101,7 @@ def attn_forward_func_with_cp( deterministic, use_fused_attention, return_max_logit, + softcap, ] if cp_comm_type in ["p2p", "a2a+p2p"]: diff --git a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py index d5adbbcadf..fad2c51951 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py @@ -578,6 +578,12 @@ def nvfp4_linear_mxfp8_dpa_factory(role): or bottom right (`True`) corner of the softmax matrix in the encoder. If `None`, it will be set to `False` for `attn_mask_type` = {'causal', 'padding_causal'} and `True` for other mask types. + softcap : float, default = 0.0 + tanh logit softcapping value applied to the attention scores as + ``softcap * tanh(scores / softcap)``. A value of ``0.0`` disables + softcapping. Softcapping is only supported by the FlashAttention + backend. Similar to :attr:`window_size`, ``softcap`` can be + overridden by :attr:`softcap` in ``forward`` as well. attention_type : str, default = "self" type of attention, either ``"self"`` and ``"cross"``. layer_number : int, default = None @@ -677,6 +683,7 @@ def __init__( attn_mask_type: str = "causal", window_size: Optional[Tuple[int, int]] = None, bottom_right_diagonal: Optional[bool] = None, + softcap: float = 0.0, sequence_parallel: bool = False, tp_size: int = 1, get_rng_state_tracker: Optional[Callable] = None, @@ -713,6 +720,7 @@ def __init__( self.attn_mask_type = attn_mask_type self.window_size = dpa_utils.check_set_window_size(attn_mask_type, window_size) self.bottom_right_diagonal = bottom_right_diagonal + self.softcap = softcap if tp_group is None: self.tp_size = tp_size if tp_size == 1: @@ -1393,6 +1401,7 @@ def forward( attn_mask_type: Optional[str] = None, window_size: Optional[Tuple[int, int]] = None, bottom_right_diagonal: Optional[bool] = None, + softcap: Optional[float] = None, checkpoint_core_attention: bool = False, core_attention_bias_type: str = "no_bias", core_attention_bias: Optional[torch.Tensor] = None, @@ -1563,6 +1572,11 @@ def forward( causal masks are aligned to the bottom right corner. window_size: Optional[Tuple[int, int]], default = None Sliding window size for local attention. + softcap: Optional[float], default = None + tanh logit softcapping value applied to the attention scores as + ``softcap * tanh(scores / softcap)``. A value of ``0.0`` disables + softcapping. When `None`, the value passed to the constructor is used. + Softcapping is only supported by the FlashAttention backend. bottom_right_diagonal: Optional[bool], default = None Align sliding window and ALiBi diagonal to the top left (`False`) or bottom right (`True`) corner of the softmax matrix in the encoder. @@ -1743,6 +1757,8 @@ def forward( if window_size is None: window_size = self.window_size window_size = dpa_utils.check_set_window_size(attn_mask_type, window_size) + if softcap is None: + softcap = self.softcap if bottom_right_diagonal is None: bottom_right_diagonal = self.bottom_right_diagonal if attn_mask_type in {"causal", "padding_causal"}: @@ -2025,6 +2041,14 @@ def forward( else: pad_between_seqs = False + # softcap is not supported by UnfusedDotProductAttention, which is the backend ONNX + # export unconditionally force-selects further down (bypassing get_attention_backend's + # softcap-aware filter). Fail loudly rather than silently export a model that omits + # softcapping. + assert ( + softcap == 0.0 or not is_in_onnx_export_mode() + ), "Attention logit softcapping (softcap != 0.0) is not supported with ONNX export!" + # Validate experimental Flex Attention API inputs that backend selection # cannot represent. if score_mod is None: @@ -2074,6 +2098,7 @@ def forward( attn_mask_type=attn_mask_type, window_size=window_size, bottom_right_diagonal=bottom_right_diagonal, + softcap=softcap, alibi_slopes_shape=alibi_slopes.shape if alibi_slopes is not None else None, core_attention_bias_type=core_attention_bias_type, core_attention_bias_shape=core_attention_bias_shape, @@ -2205,6 +2230,7 @@ def forward( cu_seqlens_kv=cu_seqlens_kv, attn_mask_type=attn_mask_type, window_size=window_size, + softcap=softcap, alibi_slopes=alibi_slopes, cp_group=self.cp_group, cp_global_ranks=self.cp_global_ranks, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index ba049c9aef..aeab47007e 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -148,6 +148,11 @@ class FlashAttentionUtils: v4_is_installed = False fa4_version = PkgVersion("0") use_v4 = False + # True only if the installed FA3 build exposes a `softcap` parameter (signature probe in + # backends.py, fail-closed default False). Necessary-but-not-sufficient: FA3 softcap is also + # gated on opt-in (NVTE_FA3_SOFTCAP=1) and head_dim <= 256 in get_attention_backend. FA3 is + # already restricted to Hopper (sm90) upstream, where its softcap fwd+bwd kernels are mature. + fa3_supports_softcap = False v4_installation_steps = """\ pip install flash-attn-4==4.0.0b11 nvidia-cutlass-dsl[cu13]""" v4_warning_printed = False @@ -229,6 +234,9 @@ class AttentionParams: bottom_right_diagonal: bool, default = `None` Whether to align sliding window and ALiBi diagonal to the bottom right corner of the softmax matrix. + softcap : float, default = 0.0 + Tanh logit softcapping value applied to the attention scores. A value of + ``0.0`` disables softcapping. Only supported by the FlashAttention backend. alibi_slopes_shape : Optional[Union[torch.Size, List]], default = None Tensor shape of :attr:`alibi_slopes` in `DotProductAttention`. core_attention_bias_type : str, default = no_bias @@ -289,6 +297,7 @@ class AttentionParams: attn_mask_type: str = "no_mask" window_size: Union[Tuple[int, int], None] = None bottom_right_diagonal: bool = True + softcap: float = 0.0 alibi_slopes_shape: Union[torch.Size, List, None] = None core_attention_bias_type: str = "no_bias" core_attention_bias_shape: str = "1hss" @@ -433,6 +442,7 @@ def get_attention_backend( attn_mask_type = attention_params.attn_mask_type window_size = attention_params.window_size bottom_right_diagonal = attention_params.bottom_right_diagonal + softcap = attention_params.softcap alibi_slopes_shape = attention_params.alibi_slopes_shape core_attention_bias_type = attention_params.core_attention_bias_type core_attention_bias_shape = attention_params.core_attention_bias_shape @@ -764,6 +774,43 @@ def _disable_all_flash_attention() -> None: use_unfused_attention = False logger.debug("Disabling all backends for max_logit with FP8 attention") + # Filter: softcap + # The scalar `softcap` kwarg (tanh logit softcapping) is plumbed to the FlashAttention 2 + # backend (>= 2.6.0) by default, and to FA3 only behind an explicit opt-in gate below. + # FusedAttention/unfused don't take the scalar kwarg (cuDNN can softcap via score_mod, but that + # path is not used here). Steer selection to FA2 rather than (a) hitting a runtime + # NotImplementedError when an unwired backend is selected, or (b) silently dropping the cap. + if softcap != 0.0: + if use_fused_attention: + logger.debug("Disabling FusedAttention as it does not support softcap") + use_fused_attention = False + if use_unfused_attention: + logger.debug("Disabling UnfusedDotProductAttention as it does not support softcap") + use_unfused_attention = False + if use_flash_attention_3 and not ( + FlashAttentionUtils.fa3_supports_softcap + and os.getenv("NVTE_FA3_SOFTCAP", "0") == "1" + and max(head_dim_qk, head_dim_v) <= 256 + and not context_parallel + ): + # FA3 softcap is opt-in (NVTE_FA3_SOFTCAP=1) and requires a softcap-capable FA3 build, + # head_dim <= 256 (the range FA3's sm90 softcap kernels are instantiated for), and no + # context parallelism -- FA3's CP path hard-rejects nonzero softcap (backends.py), so + # selecting it here would just crash at dispatch instead of steering to FA2, which does + # support CP+softcap via context_parallel.py's autograd threading. FA3 is already + # Hopper-only upstream. FA3's non-CP softcap fwd+bwd is mature, so no arch/beta caveat is + # needed beyond the build probe; keep it opt-in to preserve FA2 as the default (unchanged + # behavior) and allow a clean FA2-vs-FA3 comparison. When all conditions hold, FA3 + # survives and the softcap kwarg is threaded in backends.py. + logger.debug( + "Disabling FlashAttention 3 for softcap (requires softcap-capable FA3 build, " + "NVTE_FA3_SOFTCAP=1, head_dim <= 256, and no context parallelism)" + ) + use_flash_attention_3 = False + if use_flash_attention_2 and not FlashAttentionUtils.v2_6_0_plus: + logger.debug("Disabling FlashAttention 2 for softcap (requires flash-attn >= 2.6.0)") + use_flash_attention_2 = False + # Filter: score_mod if has_score_mod_bprop and not has_score_mod: logger.debug("Disabling all backends because score_mod_bprop requires score_mod") From 9839ae08ac71d6bfa7bbda5752eb62334685fd21 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Mon, 17 Aug 2026 14:34:54 -0700 Subject: [PATCH 2/9] fix: use raise instead of assert for the ONNX+softcap guard python -O / PYTHONOPTIMIZE strips assert statements, which would silently reopen the ONNX export softcap-drop bug the previous commit fixed (ONNX mode would again force-select UnfusedDotProductAttention with softcap silently omitted, with no error). Switch to an explicit if/raise ValueError, which survives optimized execution. Signed-off-by: Nitin Vegesna Co-Authored-By: Claude Opus 4.8 (1M context) --- .../dot_product_attention/dot_product_attention.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py index fad2c51951..8ca8fbef3a 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py @@ -2044,10 +2044,13 @@ def forward( # softcap is not supported by UnfusedDotProductAttention, which is the backend ONNX # export unconditionally force-selects further down (bypassing get_attention_backend's # softcap-aware filter). Fail loudly rather than silently export a model that omits - # softcapping. - assert ( - softcap == 0.0 or not is_in_onnx_export_mode() - ), "Attention logit softcapping (softcap != 0.0) is not supported with ONNX export!" + # softcapping. Uses an explicit raise (not assert) so the check survives python -O / + # PYTHONOPTIMIZE, which strips asserts and would otherwise silently re-open this gap. + if softcap != 0.0 and is_in_onnx_export_mode(): + raise ValueError( + "Attention logit softcapping (softcap != 0.0) is not supported with " + "ONNX export!" + ) # Validate experimental Flex Attention API inputs that backend selection # cannot represent. From eb215b51decb54891895d92c516360e60ce377fc Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Mon, 17 Aug 2026 14:35:45 -0700 Subject: [PATCH 3/9] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci (reapplied after a force-push rebase clobbered pre-commit.ci's original 19a21eb7 commit; same content, restored by hand) Signed-off-by: Nitin Vegesna --- tests/pytorch/attention/test_softcap.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/tests/pytorch/attention/test_softcap.py b/tests/pytorch/attention/test_softcap.py index b0b94bb4e8..78a8cf2829 100644 --- a/tests/pytorch/attention/test_softcap.py +++ b/tests/pytorch/attention/test_softcap.py @@ -44,9 +44,7 @@ def _flash_attn_2_6_available() -> bool: # backend (>= 2.6.0), and requires CUDA tensors. pytestmark = [ pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required."), - pytest.mark.skipif( - not _flash_attn_2_6_available(), reason="flash-attn >= 2.6.0 is required." - ), + pytest.mark.skipif(not _flash_attn_2_6_available(), reason="flash-attn >= 2.6.0 is required."), ] From 2475521e4a2b284bd7d44ff98f18f40688c3811b Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 27 Aug 2026 09:22:15 -0700 Subject: [PATCH 4/9] refactor: move softcap reference into UnfusedDotProductAttention and fold test into test_attention.py UnfusedDotProductAttention now applies softcap * tanh(scores / softcap) to the already-scaled logits, matching how FlashAttention folds softmax_scale into its tanh argument, so it can serve as the softcap reference backend. Backend selection therefore no longer disqualifies unfused attention for softcap, and the ONNX-export guard is dropped since the export path force-selects unfused and torch.tanh is exportable. test_softcap.py is replaced by a model_configs_softcap dict and test_dpa_softcap in test_attention.py, which reuses test_dot_product_attention for backend sweeping. Signed-off-by: Nitin Vegesna --- qa/L0_pytorch_unittest/test.sh | 1 - tests/pytorch/attention/test_attention.py | 26 ++++ tests/pytorch/attention/test_softcap.py | 147 ------------------ tests/pytorch/utils.py | 3 + .../dot_product_attention/backends.py | 8 + .../dot_product_attention.py | 21 +-- .../attention/dot_product_attention/utils.py | 14 +- 7 files changed, 50 insertions(+), 170 deletions(-) delete mode 100644 tests/pytorch/attention/test_softcap.py diff --git a/qa/L0_pytorch_unittest/test.sh b/qa/L0_pytorch_unittest/test.sh index 482642e5e1..14a5f4fe3d 100644 --- a/qa/L0_pytorch_unittest/test.sh +++ b/qa/L0_pytorch_unittest/test.sh @@ -67,7 +67,6 @@ NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE=1 NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 python python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_linear_mxfp8_attention.xml $TE_PATH/tests/pytorch/attention/test_linear_mxfp8_attention.py || test_fail "test_linear_mxfp8_attention.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_fused_mla_q_uproj.xml $TE_PATH/tests/pytorch/attention/test_fused_mla_q_uproj.py || test_fail "test_fused_mla_q_uproj.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_kv_cache.xml $TE_PATH/tests/pytorch/attention/test_kv_cache.py || test_fail "test_kv_cache.py" -python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_softcap.xml $TE_PATH/tests/pytorch/attention/test_softcap.py || test_fail "test_softcap.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_cu_seqlens_cache.xml $TE_PATH/tests/pytorch/attention/test_cu_seqlens_cache.py || test_fail "test_cu_seqlens_cache.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_hf_integration.xml $TE_PATH/tests/pytorch/test_hf_integration.py || test_fail "test_hf_integration.py" export NVTE_TEST_CHECKPOINT_ARTIFACT_PATH=$TE_PATH/artifacts/tests/pytorch/test_checkpoint diff --git a/tests/pytorch/attention/test_attention.py b/tests/pytorch/attention/test_attention.py index bfd2cdf9fd..8e8e557405 100644 --- a/tests/pytorch/attention/test_attention.py +++ b/tests/pytorch/attention/test_attention.py @@ -645,6 +645,31 @@ def test_dpa_softmax_thd(dtype, model_configs, model): test_dot_product_attention(dtype, model_configs, model, True, "thd_thd_thd", False, False) +model_configs_softcap = { + # test: ModelConfig(b, sq, hq, dqk) + "softcap_1_0": ModelConfig(4, 128, 16, 64, softcap=50.0), + "softcap_1_1": ModelConfig(4, 128, 16, 64, num_gqa_groups=4, softcap=50.0), + "softcap_2_0": ModelConfig(2, 512, 16, 64, attn_mask_type="causal", softcap=50.0), + "softcap_2_1": ModelConfig(2, 512, 24, 128, attn_mask_type="padding_causal", softcap=50.0), + # 0.01 is on the order of the logits these inputs produce, so tanh runs in its nonlinear + # region instead of acting as a no-op, and a misapplied softmax_scale or a missing outer + # softcap factor changes the output. + "softcap_3_0": ModelConfig(4, 128, 16, 64, softcap=0.01), + "softcap_3_1": ModelConfig(2, 512, 16, 64, attn_mask_type="causal", softcap=0.01), +} + + +@pytest.mark.skipif( + not FlashAttentionUtils.v2_6_0_plus, reason="flash-attn 2.6.0+ is required for softcap." +) +@pytest.mark.parametrize("dtype", param_types) +@pytest.mark.parametrize("model_configs", [model_configs_softcap]) +@pytest.mark.parametrize("model", model_configs_softcap.keys()) +def test_dpa_softcap(dtype, model_configs, model): + """Test DotProductAttention module with tanh logit softcapping""" + test_dot_product_attention(dtype, model_configs, model, False, "bshd_bshd_bshd", False, False) + + model_configs_mla = { # test: ModelConfig(b, sq, hq, dqk) "mla_1_0": ModelConfig(8, 128, 16, 64, head_dim_v=128), @@ -1447,6 +1472,7 @@ def get_dummy_cuda_rng_tracker() -> CudaRNGStatesTracker: attention_type=config.attn_type, softmax_type=config.softmax_type, return_max_logit=config.return_max_logit, + softcap=config.softcap, ).to(dtype=dtype, device="cuda") if not is_training: block = block.eval() diff --git a/tests/pytorch/attention/test_softcap.py b/tests/pytorch/attention/test_softcap.py deleted file mode 100644 index 78a8cf2829..0000000000 --- a/tests/pytorch/attention/test_softcap.py +++ /dev/null @@ -1,147 +0,0 @@ -# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# -# See LICENSE for license information. -"""Isolation numerics test for tanh logit softcapping in DotProductAttention. - -The reference implements softcapping in pure PyTorch: - - scores = (Q @ K^T) * scale - scores = softcap * tanh(scores / softcap) # only when softcap != 0.0 - scores = scores + mask - attn = softmax(scores) - out = attn @ V - -and is compared against ``DotProductAttention(..., softcap=...)`` forced onto the -FlashAttention backend, for both the forward output and the input gradients -(dQ/dK/dV obtained via autograd). -""" - -import sys -import pathlib - -import pytest -import torch -from packaging.version import Version as PkgVersion - -from transformer_engine.pytorch import DotProductAttention -from transformer_engine.pytorch.attention.dot_product_attention import _attention_backends - -_current_file = pathlib.Path(__file__).resolve() -sys.path = [str(_current_file.parent.parent)] + sys.path -from utils import reset_rng_states # pylint: disable=wrong-import-position - - -def _flash_attn_2_6_available() -> bool: - """Whether flash-attn >= 2.6.0 (the first version exposing ``softcap``) is installed.""" - try: - import flash_attn # pylint: disable=import-outside-toplevel - except ImportError: - return False - return PkgVersion(flash_attn.__version__) >= PkgVersion("2.6.0") - - -# Softcapping through DotProductAttention is only wired through the FlashAttention 2 -# backend (>= 2.6.0), and requires CUDA tensors. -pytestmark = [ - pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required."), - pytest.mark.skipif(not _flash_attn_2_6_available(), reason="flash-attn >= 2.6.0 is required."), -] - - -def _force_flash_backend() -> None: - """Force DotProductAttention to select the FlashAttention backend.""" - import os # pylint: disable=import-outside-toplevel - - os.environ["NVTE_FLASH_ATTN"] = "1" - os.environ["NVTE_FUSED_ATTN"] = "0" - os.environ["NVTE_UNFUSED_ATTN"] = "0" - _attention_backends["backend_selection_requires_update"] = True - - -def _reference_attention(q, k, v, scale, softcap, causal): - """Pure-PyTorch reference for softcapped scaled dot product attention. - - q, k, v are in ``bshd`` layout. GQA is supported: ``k``/``v`` may have fewer - heads than ``q``. - """ - # bshd -> bhsd - qt = q.transpose(1, 2).float() - kt = k.transpose(1, 2).float() - vt = v.transpose(1, 2).float() - - num_heads = qt.shape[1] - num_gqa_groups = kt.shape[1] - if num_heads != num_gqa_groups: - assert num_heads % num_gqa_groups == 0 - repeats = num_heads // num_gqa_groups - kt = kt.repeat_interleave(repeats, dim=1) - vt = vt.repeat_interleave(repeats, dim=1) - - scores = torch.matmul(qt, kt.transpose(-2, -1)) * scale - if softcap != 0.0: - scores = softcap * torch.tanh(scores / softcap) - if causal: - sq, skv = scores.shape[-2], scores.shape[-1] - mask = torch.triu( - torch.ones(sq, skv, dtype=torch.bool, device=scores.device), - diagonal=1 + skv - sq, - ) - scores = scores.masked_fill(mask, float("-inf")) - attn = torch.softmax(scores, dim=-1) - out = torch.matmul(attn, vt) - # bhsd -> bshd - return out.transpose(1, 2) - - -@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) -@pytest.mark.parametrize("softcap", [0.0, 50.0]) -@pytest.mark.parametrize("num_gqa_groups", [4, 2]) -@pytest.mark.parametrize("causal", [False, True]) -def test_softcap_numerics(dtype, softcap, num_gqa_groups, causal): - """FlashAttention softcap forward + grads match a pure-PyTorch reference. - - ``softcap == 0.0`` additionally proves that softcapping is a no-op relative to - the plain (no-softcap) reference, i.e. today's behavior is reproduced exactly. - """ - reset_rng_states() - - batch_size = 2 - max_seqlen = 32 - num_heads = 4 - head_dim = 64 - scale = 1.0 / (head_dim**0.5) - - q_shape = (batch_size, max_seqlen, num_heads, head_dim) - kv_shape = (batch_size, max_seqlen, num_gqa_groups, head_dim) - - q = (0.5 * torch.randn(q_shape, dtype=dtype, device="cuda")).requires_grad_() - k = (0.5 * torch.randn(kv_shape, dtype=dtype, device="cuda")).requires_grad_() - v = (0.5 * torch.randn(kv_shape, dtype=dtype, device="cuda")).requires_grad_() - q_ref, k_ref, v_ref = [x.detach().clone().requires_grad_() for x in (q, k, v)] - - grad_output = torch.randn(q_shape, dtype=dtype, device="cuda") - - _force_flash_backend() - dpa = DotProductAttention( - num_heads, - head_dim, - num_gqa_groups=num_gqa_groups, - qkv_format="bshd", - attn_mask_type="causal" if causal else "no_mask", - softmax_scale=scale, - softcap=softcap, - layer_number=1, - ).to(dtype=dtype, device="cuda") - - out = dpa(q, k, v) - out.backward(grad_output) - - out_ref = _reference_attention(q_ref, k_ref, v_ref, scale, softcap, causal) - out_ref.backward(grad_output.float()) - - atol, rtol = (2e-2, 2e-2) if dtype == torch.float16 else (3.5e-2, 3.5e-2) - - torch.testing.assert_close(out.float(), out_ref.float(), atol=atol, rtol=rtol) - torch.testing.assert_close(q.grad.float(), q_ref.grad.float(), atol=atol, rtol=rtol) - torch.testing.assert_close(k.grad.float(), k_ref.grad.float(), atol=atol, rtol=rtol) - torch.testing.assert_close(v.grad.float(), v_ref.grad.float(), atol=atol, rtol=rtol) diff --git a/tests/pytorch/utils.py b/tests/pytorch/utils.py index 21601d8cdd..0002bcef2c 100644 --- a/tests/pytorch/utils.py +++ b/tests/pytorch/utils.py @@ -282,6 +282,7 @@ def __init__( alibi_type: str = "none", bias_shape: str = "1hss", window_size: Tuple[int, int] = (-1, -1), + softcap: float = 0.0, context_parallel: bool = False, cp_comm_type: str = "p2p", return_max_logit=False, @@ -312,6 +313,7 @@ def __init__( self.attn_type = "self" if (self.max_seqlen_q == self.max_seqlen_kv) else "cross" self.bias_shape = bias_shape self.window_size = check_set_window_size(self.attn_mask_type, window_size) + self.softcap = softcap self.context_parallel = context_parallel self.cp_comm_type = cp_comm_type self.return_max_logit = return_max_logit @@ -390,6 +392,7 @@ def test(): head_dim_v=config.head_dim_v, attn_mask_type=config.attn_mask_type, window_size=config.window_size, + softcap=config.softcap, alibi_slopes_shape=alibi_slopes_shape, core_attention_bias_type=config.attn_bias_type, core_attention_bias_shape=core_attention_bias_shape, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index 12a2ed1492..c74e0f1f04 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -440,6 +440,7 @@ def _forward( attention_mask: Optional[Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]] = None, window_size: Optional[Tuple[int, int]] = None, bottom_right_diagonal: Optional[bool] = None, + softcap: float = 0.0, core_attention_bias_type: str = "no_bias", core_attention_bias: Optional[torch.Tensor] = None, alibi_slopes: Optional[torch.Tensor] = None, @@ -678,6 +679,13 @@ def _forward( dtype=query_layer.dtype ) + # Cap the scaled logits -- softcap * tanh(scores * scale / softcap) -- matching how + # FlashAttention folds softmax_scale into its tanh argument. qk layer scaling defers the + # layer_number factor to the softmax below, so it is divided out of the cap here. + if softcap != 0.0: + cap = softcap / self.layer_number if apply_qk_layer_scaling else softcap + matmul_result = cap * torch.tanh(matmul_result / cap) + if fp8: # quantize and dequantize dP to emulate FP8 matmul_result, *_ = FP8EmulationFunc.apply( diff --git a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py index 8ca8fbef3a..2b51ed8f2c 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py @@ -582,8 +582,9 @@ def nvfp4_linear_mxfp8_dpa_factory(role): tanh logit softcapping value applied to the attention scores as ``softcap * tanh(scores / softcap)``. A value of ``0.0`` disables softcapping. Softcapping is only supported by the FlashAttention - backend. Similar to :attr:`window_size`, ``softcap`` can be - overridden by :attr:`softcap` in ``forward`` as well. + and UnfusedDotProductAttention backends. Similar to + :attr:`window_size`, ``softcap`` can be overridden by + :attr:`softcap` in ``forward`` as well. attention_type : str, default = "self" type of attention, either ``"self"`` and ``"cross"``. layer_number : int, default = None @@ -1576,7 +1577,8 @@ def forward( tanh logit softcapping value applied to the attention scores as ``softcap * tanh(scores / softcap)``. A value of ``0.0`` disables softcapping. When `None`, the value passed to the constructor is used. - Softcapping is only supported by the FlashAttention backend. + Softcapping is only supported by the FlashAttention and + UnfusedDotProductAttention backends. bottom_right_diagonal: Optional[bool], default = None Align sliding window and ALiBi diagonal to the top left (`False`) or bottom right (`True`) corner of the softmax matrix in the encoder. @@ -2041,17 +2043,6 @@ def forward( else: pad_between_seqs = False - # softcap is not supported by UnfusedDotProductAttention, which is the backend ONNX - # export unconditionally force-selects further down (bypassing get_attention_backend's - # softcap-aware filter). Fail loudly rather than silently export a model that omits - # softcapping. Uses an explicit raise (not assert) so the check survives python -O / - # PYTHONOPTIMIZE, which strips asserts and would otherwise silently re-open this gap. - if softcap != 0.0 and is_in_onnx_export_mode(): - raise ValueError( - "Attention logit softcapping (softcap != 0.0) is not supported with " - "ONNX export!" - ) - # Validate experimental Flex Attention API inputs that backend selection # cannot represent. if score_mod is None: @@ -2365,6 +2356,7 @@ def forward( attention_mask=attention_mask, window_size=window_size, bottom_right_diagonal=bottom_right_diagonal, + softcap=softcap, core_attention_bias_type=core_attention_bias_type, core_attention_bias=core_attention_bias, alibi_slopes=alibi_slopes, @@ -2389,6 +2381,7 @@ def forward( attention_mask=attention_mask, window_size=window_size, bottom_right_diagonal=bottom_right_diagonal, + softcap=softcap, core_attention_bias_type=core_attention_bias_type, core_attention_bias=core_attention_bias, alibi_slopes=alibi_slopes, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index aeab47007e..94c2fc4d99 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -236,7 +236,8 @@ class AttentionParams: of the softmax matrix. softcap : float, default = 0.0 Tanh logit softcapping value applied to the attention scores. A value of - ``0.0`` disables softcapping. Only supported by the FlashAttention backend. + ``0.0`` disables softcapping. Only supported by the FlashAttention and + UnfusedDotProductAttention backends. alibi_slopes_shape : Optional[Union[torch.Size, List]], default = None Tensor shape of :attr:`alibi_slopes` in `DotProductAttention`. core_attention_bias_type : str, default = no_bias @@ -776,17 +777,14 @@ def _disable_all_flash_attention() -> None: # Filter: softcap # The scalar `softcap` kwarg (tanh logit softcapping) is plumbed to the FlashAttention 2 - # backend (>= 2.6.0) by default, and to FA3 only behind an explicit opt-in gate below. - # FusedAttention/unfused don't take the scalar kwarg (cuDNN can softcap via score_mod, but that - # path is not used here). Steer selection to FA2 rather than (a) hitting a runtime - # NotImplementedError when an unwired backend is selected, or (b) silently dropping the cap. + # backend (>= 2.6.0) and to UnfusedDotProductAttention by default, and to FA3 only behind an + # explicit opt-in gate below. FusedAttention does not take the scalar kwarg (cuDNN can softcap + # via score_mod, but that path is not used here), so disable it rather than silently dropping + # the cap. if softcap != 0.0: if use_fused_attention: logger.debug("Disabling FusedAttention as it does not support softcap") use_fused_attention = False - if use_unfused_attention: - logger.debug("Disabling UnfusedDotProductAttention as it does not support softcap") - use_unfused_attention = False if use_flash_attention_3 and not ( FlashAttentionUtils.fa3_supports_softcap and os.getenv("NVTE_FA3_SOFTCAP", "0") == "1" From 3c5eb4a254b616f00a0421b7ca672d5144f70a99 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 27 Aug 2026 09:28:45 -0700 Subject: [PATCH 5/9] fix(pytorch): gate FA3 softcap on existing NVTE_FLASH_ATTN_V3 Drop the redundant NVTE_FA3_SOFTCAP opt-in. `use_flash_attention_3` already derives from NVTE_FLASH_ATTN_V3, so the existing flag governs the FA3 softcap path and NVTE_FLASH_ATTN_V3=0 disables it. Correctness stays established by the build-capability probe, head_dim <= 256, and the non-CP requirement. Signed-off-by: Nitin Vegesna --- .../dot_product_attention/backends.py | 10 +++---- .../attention/dot_product_attention/utils.py | 26 ++++++++----------- 2 files changed, 16 insertions(+), 20 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index c74e0f1f04..d2efd9048d 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -170,7 +170,7 @@ # (sm90) kernels DO implement tanh logit softcapping in fwd AND bwd (dedicated # flash_{fwd,bwd}_hdim256_bf16_softcap_sm90 instantiations, off only behind a compile-time # DISABLE_SOFTCAP flag), so this is a mature path. Still fail-closed and additionally - # gated on opt-in (NVTE_FA3_SOFTCAP) + head_dim <= 256 in get_attention_backend. + # gated on head_dim <= 256 + non-CP in get_attention_backend. try: fa_utils.fa3_supports_softcap = ( "softcap" in inspect.signature(flash_attn_func_v3).parameters @@ -1267,10 +1267,10 @@ def forward( ) else: # Fail-loud net: get_attention_backend only keeps FA3 for softcap on a - # softcap-capable build (signature probe) + opt-in (NVTE_FA3_SOFTCAP) + Hopper - # (FA3 is sm90-only upstream) + head_dim <= 256. If FA3 is still reached with - # softcap while the build lacks support (force-selected / regressed path), raise - # rather than silently drop the cap. The non-CP FA3 entry points + # softcap-capable build (signature probe) + Hopper (FA3 is sm90-only upstream) + # + head_dim <= 256. If FA3 is still reached with softcap while the build lacks + # support (force-selected / regressed path), raise rather than silently drop the + # cap. The non-CP FA3 entry points # (flash_attn_func_v3 / flash_attn_varlen_func_v3) are self-contained autograd # functions, so threading `softcap` into the forward call also drives the # matching FA3 softcap backward kernel. (CP + FA3 + softcap stays blocked above.) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 94c2fc4d99..ea0329f557 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -150,8 +150,8 @@ class FlashAttentionUtils: use_v4 = False # True only if the installed FA3 build exposes a `softcap` parameter (signature probe in # backends.py, fail-closed default False). Necessary-but-not-sufficient: FA3 softcap is also - # gated on opt-in (NVTE_FA3_SOFTCAP=1) and head_dim <= 256 in get_attention_backend. FA3 is - # already restricted to Hopper (sm90) upstream, where its softcap fwd+bwd kernels are mature. + # gated on head_dim <= 256 and non-CP in get_attention_backend. FA3 is already restricted to + # Hopper (sm90) upstream, where its softcap fwd+bwd kernels are mature. fa3_supports_softcap = False v4_installation_steps = """\ pip install flash-attn-4==4.0.0b11 nvidia-cutlass-dsl[cu13]""" @@ -777,8 +777,8 @@ def _disable_all_flash_attention() -> None: # Filter: softcap # The scalar `softcap` kwarg (tanh logit softcapping) is plumbed to the FlashAttention 2 - # backend (>= 2.6.0) and to UnfusedDotProductAttention by default, and to FA3 only behind an - # explicit opt-in gate below. FusedAttention does not take the scalar kwarg (cuDNN can softcap + # backend (>= 2.6.0) and to UnfusedDotProductAttention by default, and to FA3 subject to the + # build/shape checks below. FusedAttention does not take the scalar kwarg (cuDNN can softcap # via score_mod, but that path is not used here), so disable it rather than silently dropping # the cap. if softcap != 0.0: @@ -787,22 +787,18 @@ def _disable_all_flash_attention() -> None: use_fused_attention = False if use_flash_attention_3 and not ( FlashAttentionUtils.fa3_supports_softcap - and os.getenv("NVTE_FA3_SOFTCAP", "0") == "1" and max(head_dim_qk, head_dim_v) <= 256 and not context_parallel ): - # FA3 softcap is opt-in (NVTE_FA3_SOFTCAP=1) and requires a softcap-capable FA3 build, - # head_dim <= 256 (the range FA3's sm90 softcap kernels are instantiated for), and no - # context parallelism -- FA3's CP path hard-rejects nonzero softcap (backends.py), so - # selecting it here would just crash at dispatch instead of steering to FA2, which does - # support CP+softcap via context_parallel.py's autograd threading. FA3 is already - # Hopper-only upstream. FA3's non-CP softcap fwd+bwd is mature, so no arch/beta caveat is - # needed beyond the build probe; keep it opt-in to preserve FA2 as the default (unchanged - # behavior) and allow a clean FA2-vs-FA3 comparison. When all conditions hold, FA3 - # survives and the softcap kwarg is threaded in backends.py. + # FA3 softcap requires a softcap-capable FA3 build, head_dim <= 256 (the range FA3's + # sm90 softcap kernels are instantiated for), and no context parallelism -- FA3's CP + # path hard-rejects nonzero softcap (backends.py), so selecting it here would just + # crash at dispatch instead of steering to FA2, which does support CP+softcap via + # context_parallel.py's autograd threading. Whether FA3 is eligible at all is governed + # by NVTE_FLASH_ATTN_V3 through use_flash_attention_3. logger.debug( "Disabling FlashAttention 3 for softcap (requires softcap-capable FA3 build, " - "NVTE_FA3_SOFTCAP=1, head_dim <= 256, and no context parallelism)" + "head_dim <= 256, and no context parallelism)" ) use_flash_attention_3 = False if use_flash_attention_2 and not FlashAttentionUtils.v2_6_0_plus: From 900371a61dbc0a9573e79532738f293f5706b8bf Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 27 Aug 2026 09:45:02 -0700 Subject: [PATCH 6/9] fix(pytorch): disable FlashAttention 4 for softcap FA4 exposes no softcap kwarg and its head_dim=256 kernel asserts score_mod is None, so there is no kernel to route the cap through. The FA4 call path in backends.py passes no softcap, so an FA4 selection with a nonzero softcap silently dropped the cap instead of failing closed. NVTE_FLASH_ATTN_V4 defaults to enabled, so this was reachable on SM100+ with flash-attn v4 installed and no context parallelism. Signed-off-by: Nitin Vegesna --- .../pytorch/attention/dot_product_attention/utils.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index ea0329f557..75007e60b3 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -779,12 +779,19 @@ def _disable_all_flash_attention() -> None: # The scalar `softcap` kwarg (tanh logit softcapping) is plumbed to the FlashAttention 2 # backend (>= 2.6.0) and to UnfusedDotProductAttention by default, and to FA3 subject to the # build/shape checks below. FusedAttention does not take the scalar kwarg (cuDNN can softcap - # via score_mod, but that path is not used here), so disable it rather than silently dropping - # the cap. + # via score_mod, but that path is not used here), and FA4 has no softcap kernel to call, so + # disable both rather than silently dropping the cap. if softcap != 0.0: if use_fused_attention: logger.debug("Disabling FusedAttention as it does not support softcap") use_fused_attention = False + if use_flash_attention_4: + # FA4 exposes no softcap kwarg and its head_dim=256 kernel asserts score_mod is None, + # so there is no kernel to route the cap through, and the FA4 call path in backends.py + # passes no softcap -- selecting it here would silently drop the cap. + if FlashAttentionUtils.v4_is_installed: + logger.debug("Disabling FlashAttention 4 as it does not support softcap") + use_flash_attention_4 = False if use_flash_attention_3 and not ( FlashAttentionUtils.fa3_supports_softcap and max(head_dim_qk, head_dim_v) <= 256 From 5ecabacd816ec754bb1cc9f389e8dceb22a8a93f Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 27 Aug 2026 09:45:46 -0700 Subject: [PATCH 7/9] fix(pytorch): disable FlashAttention 2 for softcap with dropout flash-attn rejects a nonzero softcap combined with nonzero dropout at dispatch: "Softcapping does not support dropout for now" in csrc/flash_attn/flash_api.cpp, present in mha_fwd and mha_varlen_fwd from v2.6.0 (the earliest version TE allows softcap on) onwards. Backend selection did not model this, so a softcap + attention-dropout config passed selection, routed to FA2, and crashed inside flash-attn. Dropout only reaches the kernel while training, since backends.py passes `self.attention_dropout if self.training else 0.0`, so the gate is on `attention_dropout != 0.0 and is_training` to avoid blocking valid inference configs. UnfusedDotProductAttention supports both softcap and dropout and stays available, so this steers rather than hard-fails. Signed-off-by: Nitin Vegesna --- .../pytorch/attention/dot_product_attention/utils.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 75007e60b3..6c0ac25006 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -811,6 +811,12 @@ def _disable_all_flash_attention() -> None: if use_flash_attention_2 and not FlashAttentionUtils.v2_6_0_plus: logger.debug("Disabling FlashAttention 2 for softcap (requires flash-attn >= 2.6.0)") use_flash_attention_2 = False + if use_flash_attention_2 and attention_dropout != 0.0 and is_training: + # FA2 hard-rejects a nonzero softcap combined with nonzero dropout at dispatch + # ("Softcapping does not support dropout for now", flash_api.cpp). Dropout only reaches + # the kernel while training -- backends.py passes 0.0 in eval -- hence the is_training. + logger.debug("Disabling FlashAttention 2 for softcap with dropout") + use_flash_attention_2 = False # Filter: score_mod if has_score_mod_bprop and not has_score_mod: From 4cc2e9c773dc10a444768b8b05a781f8f49425fb Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 27 Aug 2026 10:16:16 -0700 Subject: [PATCH 8/9] test: restore softcap dQ/dK/dV parity in the shared DPA harness test_dot_product_attention forced is_training=False whenever FusedAttention could not train a config, so that backends only available for inference could still be compared. softcap always disables FusedAttention, so test_dpa_softcap silently degraded to a forward-only comparison and the PR's backward-parity claim -- the FA2 softcap backward kernel included -- went untested. Add fwd_only_without_fused_attn (default True, so every other caller is byte-for-byte unchanged) and opt test_dpa_softcap out, which pairs FlashAttention against UnfusedDotProductAttention with is_training=True and restores the dgrad comparison. Signed-off-by: Nitin Vegesna --- tests/pytorch/attention/test_attention.py | 18 ++++++++++++++++-- 1 file changed, 16 insertions(+), 2 deletions(-) diff --git a/tests/pytorch/attention/test_attention.py b/tests/pytorch/attention/test_attention.py index 9e3a078817..b9e446cd37 100644 --- a/tests/pytorch/attention/test_attention.py +++ b/tests/pytorch/attention/test_attention.py @@ -170,6 +170,7 @@ def test_dot_product_attention( pad_between_seqs, declarative_packed=False, is_training=True, + fwd_only_without_fused_attn=True, ): """Test DotProductAttention module""" @@ -222,7 +223,11 @@ def test_dot_product_attention( ) flash_attn_supported, fused_attn_supported, unfused_attn_supported = available_backends - if not fused_attn_supported: + # Some backends are only available in inference mode, so when FusedAttention cannot train this + # config the query is repeated forward-only to recover enough backends to compare. Callers + # whose backward-capable pair does not include FusedAttention -- softcap, where + # get_attention_backend always disables FusedAttention -- opt out to keep dgrad coverage. + if not fused_attn_supported and fwd_only_without_fused_attn: is_training = False available_backends, _, fused_attn_backends = get_available_attention_backends( config, @@ -667,7 +672,16 @@ def test_dpa_softmax_thd(dtype, model_configs, model): @pytest.mark.parametrize("model", model_configs_softcap.keys()) def test_dpa_softcap(dtype, model_configs, model): """Test DotProductAttention module with tanh logit softcapping""" - test_dot_product_attention(dtype, model_configs, model, False, "bshd_bshd_bshd", False, False) + test_dot_product_attention( + dtype, + model_configs, + model, + False, + "bshd_bshd_bshd", + False, + False, + fwd_only_without_fused_attn=False, + ) model_configs_mla = { From 13d65a10929168792ab849e21ba17fe73138dbd6 Mon Sep 17 00:00:00 2001 From: Nitin Vegesna Date: Thu, 27 Aug 2026 10:16:30 -0700 Subject: [PATCH 9/9] test: add softcap no-op and closed-form reference coverage Two gaps remained after folding test_softcap.py into test_attention.py. softcap=0.0 no-op: every model_configs_softcap entry uses a nonzero cap, so nothing asserted the backward-compatibility claim. The half that the PR actually changed is backend selection, and a filter that fired at 0.0 would silently remove FusedAttention and FA4 from other tests rather than fail one. test_dpa_softcap_zero_backend_selection asserts FusedAttention survives softcap=0.0 and is disabled by a nonzero cap. Unfused coverage and tanh's nonlinear region: test_dpa_softcap needs two TE backends, so it skips entirely without flash-attn even though UnfusedDotProductAttention now implements softcap and is the reference for everything else. It also cannot detect a dropped cap at all: 0.1 * randn inputs put the logits at O(1e-2), where the reference output moves by 9e-9 at cap=50 and 2e-4 at cap=0.01. test_dpa_softcap_vs_reference compares forward and dQ/dK/dV against a pure-PyTorch oracle one backend at a time, so it runs with unfused alone, and uses randn inputs so the cap moves the output by O(1). An assertion on that displacement keeps the test from going vacuous if the config drifts. Signed-off-by: Nitin Vegesna --- tests/pytorch/attention/test_attention.py | 169 +++++++++++++++++++++- 1 file changed, 163 insertions(+), 6 deletions(-) diff --git a/tests/pytorch/attention/test_attention.py b/tests/pytorch/attention/test_attention.py index b9e446cd37..0ad3c90e1b 100644 --- a/tests/pytorch/attention/test_attention.py +++ b/tests/pytorch/attention/test_attention.py @@ -656,17 +656,17 @@ def test_dpa_softmax_thd(dtype, model_configs, model): "softcap_1_1": ModelConfig(4, 128, 16, 64, num_gqa_groups=4, softcap=50.0), "softcap_2_0": ModelConfig(2, 512, 16, 64, attn_mask_type="causal", softcap=50.0), "softcap_2_1": ModelConfig(2, 512, 24, 128, attn_mask_type="padding_causal", softcap=50.0), - # 0.01 is on the order of the logits these inputs produce, so tanh runs in its nonlinear - # region instead of acting as a no-op, and a misapplied softmax_scale or a missing outer - # softcap factor changes the output. + # The shared harness feeds 0.1 * randn, which puts the logits at O(1e-2) whatever the head + # dim, so tanh is numerically linear at a Gemma-sized cap. A cap of 0.01 is the one regime + # these inputs can distinguish: dropping the outer softcap factor would leave logits of + # O(1) instead of O(1e-2) and move the output well past the tolerance. Softcapping in + # tanh's saturating region is covered by test_dpa_softcap_vs_reference, which uses its own + # inputs. "softcap_3_0": ModelConfig(4, 128, 16, 64, softcap=0.01), "softcap_3_1": ModelConfig(2, 512, 16, 64, attn_mask_type="causal", softcap=0.01), } -@pytest.mark.skipif( - not FlashAttentionUtils.v2_6_0_plus, reason="flash-attn 2.6.0+ is required for softcap." -) @pytest.mark.parametrize("dtype", param_types) @pytest.mark.parametrize("model_configs", [model_configs_softcap]) @pytest.mark.parametrize("model", model_configs_softcap.keys()) @@ -684,6 +684,163 @@ def test_dpa_softcap(dtype, model_configs, model): ) +@pytest.mark.skipif(get_cudnn_version() < (8, 9, 1), reason="cuDNN 8.9.1+ is required.") +@pytest.mark.parametrize("dtype", param_types_lean) +@pytest.mark.parametrize("model_configs", [model_configs_softcap]) +@pytest.mark.parametrize("model", ["softcap_1_0"]) +def test_dpa_softcap_zero_backend_selection(dtype, model_configs, model): + """Test that softcap=0.0 leaves backend selection untouched. + + The softcap filter in get_attention_backend disables FusedAttention (and FA4) whenever the + cap is nonzero. If it also fired at 0.0, those backends would silently drop out of every + other test in this file rather than failing one, so assert both halves here. + """ + config = copy.deepcopy(model_configs[model]) + query = dict( + qkv_dtype=dtype, + qkv_layout="bshd_bshd_bshd", + is_training=True, + deterministic=_deterministic, + ) + + config.softcap = 0.0 + (_, fused_off, unfused_off), _, _ = get_available_attention_backends(config, **query) + config.softcap = 50.0 + (_, fused_on, unfused_on), _, _ = get_available_attention_backends(config, **query) + + assert fused_off, "softcap=0.0 must not disable FusedAttention" + assert not fused_on, "a nonzero softcap must disable FusedAttention" + assert unfused_off and unfused_on, "UnfusedDotProductAttention must support softcap" + + +def _softcap_reference_attention( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + softmax_scale: float, + softcap: float, + causal: bool, +) -> torch.Tensor: + """Closed-form softcapped attention in bshd layout, computed in fp32. + + scores = softcap * tanh(Q @ K^T * softmax_scale / softcap), with the tanh skipped entirely + when softcap == 0.0, so this doubles as the reference for the no-op claim. GQA is supported. + """ + q, k, v = (x.transpose(1, 2).float() for x in (q, k, v)) + if q.shape[1] != k.shape[1]: + repeats = q.shape[1] // k.shape[1] + k = k.repeat_interleave(repeats, dim=1) + v = v.repeat_interleave(repeats, dim=1) + scores = torch.matmul(q, k.transpose(-2, -1)) * softmax_scale + if softcap != 0.0: + scores = softcap * torch.tanh(scores / softcap) + if causal: + max_seqlen_q, max_seqlen_kv = scores.shape[-2], scores.shape[-1] + mask = torch.triu( + torch.ones(max_seqlen_q, max_seqlen_kv, dtype=torch.bool, device=scores.device), + diagonal=1 + max_seqlen_kv - max_seqlen_q, + ) + scores = scores.masked_fill(mask, float("-inf")) + return torch.matmul(torch.softmax(scores, dim=-1), v).transpose(1, 2) + + +model_configs_softcap_reference = { + # test: ModelConfig(b, sq, hq, dqk) + "softcap_ref_1_0": ModelConfig(2, 128, 8, 64), + "softcap_ref_1_1": ModelConfig(2, 128, 8, 64, num_gqa_groups=2), + "softcap_ref_2_0": ModelConfig(2, 128, 8, 64, attn_mask_type="causal"), +} + + +@pytest.mark.parametrize("dtype", param_types) +@pytest.mark.parametrize("model_configs", [model_configs_softcap_reference]) +@pytest.mark.parametrize("model", model_configs_softcap_reference.keys()) +@pytest.mark.parametrize("softcap", [0.0, 0.5]) +@pytest.mark.parametrize("backend", ["UnfusedDotProductAttention", "FlashAttention"]) +def test_dpa_softcap_vs_reference(dtype, model_configs, model, softcap, backend): + """Test softcap forward and dQ/dK/dV against a closed-form reference, one backend at a time. + + This needs only one TE backend, so UnfusedDotProductAttention -- the reference + implementation for every other softcap test -- stays covered on machines without + flash-attn. softcap=0.0 checks against a reference that never applies tanh, which is the + numerical half of the no-op claim. + """ + config = copy.deepcopy(model_configs[model]) + config.softcap = softcap + available_backends, _, _ = get_available_attention_backends( + config, + qkv_dtype=dtype, + qkv_layout="bshd_bshd_bshd", + is_training=True, + deterministic=_deterministic, + ) + supported = dict( + zip(["FlashAttention", "FusedAttention", "UnfusedDotProductAttention"], available_backends) + ) + if not supported[backend]: + pytest.skip(f"{backend} is unavailable for this config.") + + reset_rng_states() + os.environ["NVTE_FLASH_ATTN"] = "1" if backend == "FlashAttention" else "0" + os.environ["NVTE_FUSED_ATTN"] = "0" + os.environ["NVTE_UNFUSED_ATTN"] = "1" if backend == "UnfusedDotProductAttention" else "0" + _attention_backends["backend_selection_requires_update"] = True + + causal = "causal" in config.attn_mask_type + softmax_scale = 1.0 / config.head_dim_qk**0.5 + q_shape = (config.batch_size, config.max_seqlen_q, config.num_heads, config.head_dim_qk) + k_shape = (config.batch_size, config.max_seqlen_kv, config.num_gqa_groups, config.head_dim_qk) + v_shape = (config.batch_size, config.max_seqlen_kv, config.num_gqa_groups, config.head_dim_v) + out_shape = (config.batch_size, config.max_seqlen_q, config.num_heads, config.head_dim_v) + # randn puts the logits at O(1), so a cap of 0.5 lands in tanh's saturating region and moves + # the output by O(1). The shared harness uses 0.1 * randn, where the logits are O(1e-2) and + # no cap value is distinguishable from no cap at all. + q, k, v = ( + torch.randn(shape, dtype=dtype, device="cuda").requires_grad_() + for shape in (q_shape, k_shape, v_shape) + ) + q_ref, k_ref, v_ref = (x.detach().clone().requires_grad_() for x in (q, k, v)) + # DotProductAttention merges the head and head-dim axes of its output. + d_out = torch.randn(out_shape, dtype=dtype, device="cuda") + + block = DotProductAttention( + config.num_heads, + (config.head_dim_qk, config.head_dim_v), + num_gqa_groups=config.num_gqa_groups, + qkv_format="bshd", + attn_mask_type=config.attn_mask_type, + softmax_scale=softmax_scale, + softcap=softcap, + layer_number=1, + ).to(dtype=dtype, device="cuda") + out = block(q, k, v).view(out_shape) + out.backward(d_out) + + out_ref = _softcap_reference_attention(q_ref, k_ref, v_ref, softmax_scale, softcap, causal) + out_ref.backward(d_out.float()) + + tols = dict(atol=2e-2, rtol=2e-2) + if dtype == torch.bfloat16: + tols = dict(atol=4e-2, rtol=4e-2) + + if softcap != 0.0: + # Without this the test could be vacuous: a backend that dropped softcap on the floor + # would still match a reference whose tanh is numerically the identity. + out_ref_uncapped = _softcap_reference_attention( + q_ref.detach(), k_ref.detach(), v_ref.detach(), softmax_scale, 0.0, causal + ) + cap_effect = (out_ref.detach() - out_ref_uncapped).abs().max().item() + assert cap_effect > 10 * tols["atol"], ( + f"softcap={softcap} moves the reference output by only {cap_effect:.2e}; this config" + " would pass even if the backend ignored softcap" + ) + + torch.testing.assert_close(out.float(), out_ref, **tols) + torch.testing.assert_close(q.grad.float(), q_ref.grad.float(), **tols) + torch.testing.assert_close(k.grad.float(), k_ref.grad.float(), **tols) + torch.testing.assert_close(v.grad.float(), v_ref.grad.float(), **tols) + + model_configs_mla = { # test: ModelConfig(b, sq, hq, dqk) "mla_1_0": ModelConfig(8, 128, 16, 64, head_dim_v=128),