diff --git a/tests/pytorch/attention/test_attention.py b/tests/pytorch/attention/test_attention.py index 9dd377c41c..0ad3c90e1b 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, @@ -645,6 +650,197 @@ 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), + # 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.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, + fwd_only_without_fused_attn=False, + ) + + +@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), @@ -1447,6 +1643,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/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 bca1f3200d..0904c1a605 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 @@ -166,6 +167,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 head_dim <= 256 + non-CP 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")) @@ -435,6 +449,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, @@ -673,6 +688,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( @@ -894,6 +916,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, @@ -1110,6 +1133,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, @@ -1140,6 +1168,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, @@ -1237,6 +1266,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"] = ( @@ -1257,9 +1288,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) + 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 78e599199f..0bca0a5f9f 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1617,6 +1617,7 @@ def forward( deterministic, use_fused_attention, return_max_logit, + softcap, fp8, fp8_meta, cp_group, @@ -1906,7 +1907,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] @@ -2399,6 +2400,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 @@ -2705,7 +2707,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): @@ -3289,6 +3291,7 @@ def forward( deterministic, use_fused_attention, return_max_logit, + softcap, window_size, cp_group, cp_stream, @@ -3382,7 +3385,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 @@ -3974,6 +3977,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.use_flash_attn_4 = use_flash_attn_4 @@ -4183,7 +4187,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] @@ -4594,6 +4598,7 @@ def forward( deterministic, use_fused_attention, return_max_logit, + softcap, window_size, fp8, fp8_meta, @@ -4695,7 +4700,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__ @@ -5020,6 +5025,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 @@ -5170,7 +5176,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: @@ -5427,6 +5433,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, @@ -5614,6 +5621,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 1e5f5552d4..88b8f42a7c 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 @@ -584,6 +584,13 @@ 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 + 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 @@ -683,6 +690,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, @@ -719,6 +727,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: @@ -1408,6 +1417,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, @@ -1578,6 +1588,12 @@ 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 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. @@ -1758,6 +1774,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"}: @@ -2089,6 +2107,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, @@ -2220,6 +2239,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, @@ -2354,6 +2374,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, @@ -2378,6 +2399,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 e59405db74..2bd3876e1b 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 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]""" v4_warning_printed = False @@ -229,6 +234,10 @@ 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 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 @@ -289,6 +298,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 +443,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 +775,49 @@ 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) 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), 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 + and not context_parallel + ): + # 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, " + "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 + 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: logger.debug("Disabling all backends because score_mod_bprop requires score_mod")