diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index e36841e44..45a6572cf 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -50,11 +50,6 @@ if TYPE_CHECKING: from sgl_kernel import merge_state_v2 -from sglang.kernels.ops.attention.flash_attention import ( - flash_attn_varlen_func, - flash_attn_with_kvcache, -) - def _should_disable_scheduler_metadata_precompute() -> bool: return bool(get_parallel().enable_prefill_cp or get_parallel().enable_dp_attention) @@ -1283,7 +1278,13 @@ class FlashAttentionBackend(AttentionBackend): aux_tensors=None, rel_bias=None, rel_bias_event=None, - ): + # Returns (output, lse) with lse in [total_q, num_heads]. + return_lse: bool = False, + ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: + lse_out = None + # Bound in __init__ so a subclass can substitute a different FA4 build. + flash_attn_with_kvcache = self.flash_attn_with_kvcache + flash_attn_varlen_func = self.flash_attn_varlen_func if score_mod is not None and self.fa_impl_ver != 4: raise RuntimeError("score_mod is only supported by the FA4 backend.") cp_active = is_cp_active(forward_batch) @@ -1563,7 +1564,7 @@ class FlashAttentionBackend(AttentionBackend): causal=False if use_cascade_attn else causal, window_size=window_size, softcap=layer.logit_cap, - return_softmax_lse=use_cascade_attn, + return_softmax_lse=use_cascade_attn or return_lse, num_splits=self.num_splits, out=_fa_out, ver=self.fa_impl_ver, @@ -1623,6 +1624,8 @@ class FlashAttentionBackend(AttentionBackend): o_expand, softmax_lse_expand.T.contiguous(), ) + elif return_lse: + o, lse_out, *_ = result else: o = result else: @@ -1823,7 +1826,12 @@ class FlashAttentionBackend(AttentionBackend): else: o = result - return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) + o = o.view(-1, layer.tp_q_head_num * layer.v_head_dim) + if return_lse: + assert lse_out is not None + # The varlen kernel emits LSE head-major [num_heads, total_q]. + return o, lse_out.transpose(0, 1).contiguous() + return o def forward_decode( self, @@ -1844,7 +1852,13 @@ class FlashAttentionBackend(AttentionBackend): aux_tensors=None, rel_bias=None, rel_bias_event=None, - ) -> torch.Tensor: + # Returns (output, lse) with lse in [total_q, num_heads]. + return_lse: bool = False, + ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: + lse_out = None + # Bound in __init__ so a subclass can substitute a different FA4 build. + flash_attn_with_kvcache = self.flash_attn_with_kvcache + flash_attn_varlen_func = self.flash_attn_varlen_func if score_mod is not None and self.fa_impl_ver != 4: raise RuntimeError("score_mod is only supported by the FA4 backend.") if k is not None: @@ -2045,7 +2059,7 @@ class FlashAttentionBackend(AttentionBackend): causal=False if use_cascade_attn else causal, window_size=window_size, softcap=layer.logit_cap, - return_softmax_lse=use_cascade_attn, + return_softmax_lse=use_cascade_attn or return_lse, num_splits=( self.decode_num_splits if not is_swa_layer @@ -2090,6 +2104,8 @@ class FlashAttentionBackend(AttentionBackend): o_expand, softmax_lse_expand.T.contiguous(), ) + elif return_lse: + o, lse_out, *_ = result else: o = result else: @@ -2168,7 +2184,12 @@ class FlashAttentionBackend(AttentionBackend): else: o = result - return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) + o = o.view(-1, layer.tp_q_head_num * layer.v_head_dim) + if return_lse: + assert lse_out is not None + # The varlen kernel emits LSE head-major [num_heads, total_q]. + return o, lse_out.transpose(0, 1).contiguous() + return o def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): """Initialize CUDA graph state for the attention backend. diff --git a/test/registered/attention/unittests/dense/test_fa4.py b/test/registered/attention/unittests/dense/test_fa4.py index 2adbce5d6..c150248b7 100644 --- a/test/registered/attention/unittests/dense/test_fa4.py +++ b/test/registered/attention/unittests/dense/test_fa4.py @@ -3,9 +3,13 @@ import unittest import torch from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.kits.attention_unittest.attention_methods.dense_attention import ( + DENSE_ATOL, + DENSE_RTOL, DenseAttentionCase, + build_dense_attention_fixture, make_dense_cases, run_dense_attention_case, ) @@ -423,6 +427,102 @@ class TestFA4DenseAttentionBackendCorrectness(CustomTestCase): hidden_size=self.HIDDEN_SIZE, ) + RETURN_LSE_CASES = ( + DenseAttentionCase( + name="return_lse_mha_extend", + backend="fa4", + forward_mode=ForwardMode.EXTEND, + num_heads=4, + num_kv_heads=4, + page_size=1, + prefix_lens=(2, 4), + extend_lens=(3, 1), + ), + DenseAttentionCase( + name="return_lse_gqa_decode", + backend="fa4", + forward_mode=ForwardMode.DECODE, + num_heads=8, + num_kv_heads=2, + page_size=1, + prefix_lens=(5, 9), + ), + ) + + def _reference_out_and_lse(self, fixture): + """Causal fp32 reference: output and per-query LSE.""" + case, module = fixture.case, fixture.reference_module + dim = module.head_dim + rep = case.num_heads // case.num_kv_heads + q, k, v = module.project_qkv(fixture.input_hidden) + q = q.view(-1, case.num_heads, dim).float() + k = k.view(-1, case.num_kv_heads, dim).float() + v = v.view(-1, case.num_kv_heads, dim).float() + + outs, lses, seen = [], [], 0 + for req, prefix in enumerate(fixture.prefix_hidden): + _, prefix_k, prefix_v = module.project_qkv(prefix) + n = case.input_lens[req] + keys = torch.cat( + [prefix_k.view(-1, case.num_kv_heads, dim).float(), k[seen : seen + n]] + ).repeat_interleave(rep, dim=1) + values = torch.cat( + [prefix_v.view(-1, case.num_kv_heads, dim).float(), v[seen : seen + n]] + ).repeat_interleave(rep, dim=1) + for offset in range(n): + end = case.prefix_lens[req] + offset + 1 + scores = ( + torch.einsum("hd,khd->hk", q[seen + offset], keys[:end]) + * module.scaling + ) + probs = torch.softmax(scores, dim=-1) + outs.append(torch.einsum("hk,khd->hd", probs, values[:end]).reshape(-1)) + lses.append(torch.logsumexp(scores, dim=-1)) + seen += n + return torch.stack(outs), torch.stack(lses) + + def test_return_lse(self): + """Calls the backend directly: return_lse is a backend-level contract, + and the RadixAttention dispatcher's custom-op schema cannot carry it. + """ + for case in self.RETURN_LSE_CASES: + with self.subTest(case=case.name): + fixture = build_dense_attention_fixture( + self, case, head_dim=self.HEAD_DIM, hidden_size=self.HIDDEN_SIZE + ) + module = fixture.actual_module + forward = ( + fixture.backend.forward_decode + if case.forward_mode.is_decode() + else fixture.backend.forward_extend + ) + with ( + torch.no_grad(), + forward_context(ForwardContext(attn_backend=fixture.backend)), + ): + fixture.backend.init_forward_metadata(fixture.forward_batch) + q, k, v = module.project_qkv(fixture.input_hidden) + result = forward( + q, k, v, module.attn, fixture.forward_batch, return_lse=True + ) + + self.assertIsInstance(result, tuple) + out, lse = result + self.assertEqual( + tuple(lse.shape), (case.num_input_tokens, case.num_heads) + ) + + expected_out, expected_lse = self._reference_out_and_lse(fixture) + torch.testing.assert_close( + lse.float(), expected_lse, atol=DENSE_ATOL, rtol=DENSE_RTOL + ) + torch.testing.assert_close( + out.float().reshape(expected_out.shape), + expected_out, + atol=DENSE_ATOL, + rtol=DENSE_RTOL, + ) + if __name__ == "__main__": unittest.main()