diff --git a/tests/pytorch/attention/run_attention_with_cp.py b/tests/pytorch/attention/run_attention_with_cp.py index 7c6cdefd15..fa2ac1d299 100644 --- a/tests/pytorch/attention/run_attention_with_cp.py +++ b/tests/pytorch/attention/run_attention_with_cp.py @@ -517,7 +517,7 @@ def run_dpa_with_cp( torch.cuda.Stream(), cp_comm_type, ) - if config.softmax_type != "vanilla": + if is_training and config.softmax_type != "vanilla": core_attn.softmax_offset.grad.zero_() if dtype == "fp8": core_attn.fp8_initialized = False @@ -690,7 +690,6 @@ def run_dpa_with_cp( ) else: out = out.index_select(0, seq_idx_q).contiguous() - out_ = out_ atol, rtol, rmse_tol = get_tols(config, dtype) tensors_cp = [out_, dq_, dk_, dv_, dbias_, d_softmax_offset_, max_logit_] diff --git a/tests/pytorch/attention/test_softmax_offset_inference.py b/tests/pytorch/attention/test_softmax_offset_inference.py new file mode 100644 index 0000000000..103cb29b40 --- /dev/null +++ b/tests/pytorch/attention/test_softmax_offset_inference.py @@ -0,0 +1,20 @@ +import pytest +import torch +from transformer_engine.pytorch import DotProductAttention + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +def test_softmax_offset_grad_none_in_eval(): + """Regression test: eval mode leaves softmax_offset.grad as None. + + The context-parallel test helper previously crashed here by calling + core_attn.softmax_offset.grad.zero_() unconditionally for non-vanilla + softmax. In eval mode requires_grad is False and no backward has run, + so .grad must stay None. + """ + core_attn = ( + DotProductAttention(8, (64, 64), num_gqa_groups=4, softmax_type="softmax_offset") + .cuda() + .eval() + ) + assert not core_attn.softmax_offset.requires_grad