Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
69 changes: 69 additions & 0 deletions tests/pytorch/attention/test_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
28 changes: 24 additions & 4 deletions tests/pytorch/attention/test_cp_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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],
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading