diff --git a/tests/pytorch/attention/test_attention.py b/tests/pytorch/attention/test_attention.py index a5d2de9a4c..f045626aed 100644 --- a/tests/pytorch/attention/test_attention.py +++ b/tests/pytorch/attention/test_attention.py @@ -608,6 +608,75 @@ def test_dpa_fa4_hdim256(dtype, model_configs, model): ) +@requires_fa4 +@pytest.mark.parametrize("qkv_format", ["bshd", "thd"]) +def test_dpa_fa4_unlimited_causal_window(qkv_format, monkeypatch): + """FA4 causal attention agrees with PyTorch when TE specifies an unlimited left window.""" + for name, value in { + "NVTE_FLASH_ATTN": "1", + "NVTE_FLASH_ATTN_V2": "0", + "NVTE_FLASH_ATTN_V3": "0", + "NVTE_FLASH_ATTN_V4": "1", + "NVTE_FUSED_ATTN": "0", + "NVTE_UNFUSED_ATTN": "0", + }.items(): + monkeypatch.setenv(name, value) + _attention_backends["backend_selection_requires_update"] = True + + batch_size, seq_len, num_heads, head_dim = 2, 64, 4, 128 + torch.manual_seed(7) + q, k, v = ( + torch.randn(batch_size, seq_len, num_heads, head_dim, device="cuda", dtype=torch.bfloat16) + for _ in range(3) + ) + q_ref, k_ref, v_ref = (tensor.detach().clone().requires_grad_() for tensor in (q, k, v)) + if qkv_format == "thd": + q, k, v = (tensor.reshape(-1, num_heads, head_dim) for tensor in (q, k, v)) + cu_seqlens = torch.arange( + 0, (batch_size + 1) * seq_len, seq_len, device="cuda", dtype=torch.int32 + ) + kwargs = { + "cu_seqlens_q": cu_seqlens, + "cu_seqlens_kv": cu_seqlens, + "max_seqlen_q": seq_len, + "max_seqlen_kv": seq_len, + } + else: + kwargs = {} + q, k, v = (tensor.detach().clone().requires_grad_() for tensor in (q, k, v)) + + try: + attention = DotProductAttention( + num_heads, + head_dim, + qkv_format=qkv_format, + attn_mask_type="padding_causal" if qkv_format == "thd" else "causal", + ).to(device="cuda", dtype=torch.bfloat16) + output = attention(q, k, v, **kwargs).reshape(batch_size, seq_len, num_heads, head_dim) + assert _attention_backends["flash_attention_backend"].major == 4 + + reference = torch.nn.functional.scaled_dot_product_attention( + q_ref.permute(0, 2, 1, 3), + k_ref.permute(0, 2, 1, 3), + v_ref.permute(0, 2, 1, 3), + is_causal=True, + ).permute(0, 2, 1, 3) + torch.testing.assert_close(output, reference, atol=2e-2, rtol=2e-2) + + dout = torch.randn_like(reference) + output.backward(dout) + reference.backward(dout) + for actual, expected in zip((q, k, v), (q_ref, k_ref, v_ref)): + torch.testing.assert_close( + actual.grad.reshape(batch_size, seq_len, num_heads, head_dim), + expected.grad, + atol=2e-2, + rtol=2e-2, + ) + finally: + _attention_backends["backend_selection_requires_update"] = True + + # cuDNN FusedAttention D=256 bprop is supported on sm10x by the dedicated deterministic # SDPA bprop kernel. BSHD support starts with cuDNN FE 1.24 / BE 9.23; THD support starts # with cuDNN FE 1.26 / BE 9.25. The kernel supports d_qk == d_v == 256 only, vanilla softmax only, diff --git a/tests/pytorch/attention/test_cp_utils.py b/tests/pytorch/attention/test_cp_utils.py index 25e1fcbfe7..faa371c684 100644 --- a/tests/pytorch/attention/test_cp_utils.py +++ b/tests/pytorch/attention/test_cp_utils.py @@ -5,20 +5,25 @@ """Unit tests for context parallel utils.""" import itertools -import torch import unittest + +import torch from transformer_engine.pytorch import CPLoadBalancingStrategy from transformer_engine.pytorch.attention.dot_product_attention.context_parallel import ( _zero_thd_padding, + generate_positional_ids_for_cp, get_batch_on_this_cp_rank, get_no_load_balance_thd_causal_metadata, get_thd_partitioned_indices, + pad_thd_sequences_for_cp, restore_thd_gathered_kv, unrestore_thd_gathered_kv, - pad_thd_sequences_for_cp, - generate_positional_ids_for_cp, ) -from transformer_engine.pytorch.attention.dot_product_attention.utils import get_thd_padding_mask +from transformer_engine.pytorch.attention.dot_product_attention.utils import ( + fa4_window_size, + get_thd_padding_mask, + normalize_fa4_window_kwargs, +) try: import transformer_engine_torch as tex @@ -69,6 +74,21 @@ def test_no_load_balance_restore_uses_captured_mode(self): self.assertIs(unrestored, tokens) +class TestFA4WindowSize(unittest.TestCase): + def test_unlimited_sentinel_maps_to_none(self): + self.assertIsNone(fa4_window_size(None)) + self.assertEqual(fa4_window_size((-1, -1)), (None, None)) + self.assertEqual(fa4_window_size((-1, 0)), (None, 0)) + self.assertEqual(fa4_window_size((128, 0)), (128, 0)) + + kwargs = {"window_size_left": -1, "window_size_right": 0, "softmax_scale": 0.5} + normalize_fa4_window_kwargs(kwargs) + self.assertEqual( + kwargs, + {"window_size_left": None, "window_size_right": 0, "softmax_scale": 0.5}, + ) + + class TestSequencePadding(unittest.TestCase): def test_padding_with_custom_padding_values_sequences_shorter_than_divisibility_factor( self, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index 9a339233a4..76b687031d 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -1233,7 +1233,7 @@ def forward( fa_optional_forward_args_thd.append(max_seqlen_kv) if use_flash_attn_4: fa_4_optional_forward_kwargs = { - "window_size": window_size, + "window_size": dpa_utils.fa4_window_size(window_size), "num_splits": num_splits, } if inference_params is None: 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 9e181acbdd..f1fec9c6aa 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1171,6 +1171,7 @@ def cp_p2p_fwd_flash_attn( cu_seqlens_q_ = cu_seqlens_q_padded // 2 if use_flash_attn_4: + dpa_utils.normalize_fa4_window_kwargs(fa_forward_kwargs) fa_outputs = flash_attn_fwd( q_part, k_part, @@ -1548,6 +1549,7 @@ def cp_p2p_bwd_flash_attn( else: fa_backward_kwargs["causal"] = causal_ if use_flash_attn_4: + dpa_utils.normalize_fa4_window_kwargs(fa_backward_kwargs) dq, dk, dv = flash_attn_bwd( q_part, k_part, @@ -3766,6 +3768,7 @@ def forward( fa_forward_kwargs["window_size_left"] = window_size_per_step[i][0] fa_forward_kwargs["window_size_right"] = window_size_per_step[i][1] if use_flash_attn_4: + dpa_utils.normalize_fa4_window_kwargs(fa_forward_kwargs) fa_outputs = flash_attn_fwd( q_part, k_part, @@ -4412,6 +4415,7 @@ def backward(ctx, dout, *_args): elif not ctx.use_flash_attn_4: fa_backward_kwargs["causal"] = causal if ctx.use_flash_attn_4: + dpa_utils.normalize_fa4_window_kwargs(fa_backward_kwargs) ( dq_per_step[i], dk_per_step[i], @@ -4855,6 +4859,7 @@ def forward( fa_cu_seqlens_q = cu_seqlens_q_padded fa_cu_seqlens_kv = cu_seqlens_kv_padded if use_flash_attn_4: + dpa_utils.normalize_fa4_window_kwargs(fa_forward_kwargs) fa_outputs = flash_attn_fwd( q_part, k_part, @@ -5283,6 +5288,7 @@ def backward(ctx, dout, *_args): fa_backward_kwargs["is_causal"] = causal if ctx.use_flash_attn_4: + dpa_utils.normalize_fa4_window_kwargs(fa_backward_kwargs) dq, dk, dv = flash_attn_bwd( q, k, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 4d23e985d3..bb7c7d725b 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -74,6 +74,23 @@ _cu_seqlens_cache = {} +def fa4_window_size( + window_size: Optional[Tuple[int, int]], +) -> Optional[Tuple[Optional[int], Optional[int]]]: + """Convert TE's unlimited-window sentinel to the FA4 API convention.""" + if window_size is None: + return None + left, right = window_size + return None if left == -1 else left, None if right == -1 else right + + +def normalize_fa4_window_kwargs(kwargs: Dict[str, Any]) -> None: + """Map unlimited left and right windows before a raw FA4 call.""" + for key in ("window_size_left", "window_size_right"): + if kwargs.get(key) == -1: + kwargs[key] = None + + class AttentionLogging: """ Manage logging for attention module