[fa] Make the FlashAttention backend extensible by subclasses (#33426)
Signed-off-by: Kurt Shuster <kurt@thinkingmachines.ai> Co-authored-by: Baizhou Zhang <sobereddiezhang@gmail.com>
This commit is contained in:
co-authored by
Baizhou Zhang
parent
39e147443b
commit
f2111715cd
@@ -50,11 +50,6 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
from sgl_kernel import merge_state_v2
|
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:
|
def _should_disable_scheduler_metadata_precompute() -> bool:
|
||||||
return bool(get_parallel().enable_prefill_cp or get_parallel().enable_dp_attention)
|
return bool(get_parallel().enable_prefill_cp or get_parallel().enable_dp_attention)
|
||||||
@@ -1283,7 +1278,13 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
aux_tensors=None,
|
aux_tensors=None,
|
||||||
rel_bias=None,
|
rel_bias=None,
|
||||||
rel_bias_event=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:
|
if score_mod is not None and self.fa_impl_ver != 4:
|
||||||
raise RuntimeError("score_mod is only supported by the FA4 backend.")
|
raise RuntimeError("score_mod is only supported by the FA4 backend.")
|
||||||
cp_active = is_cp_active(forward_batch)
|
cp_active = is_cp_active(forward_batch)
|
||||||
@@ -1563,7 +1564,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
causal=False if use_cascade_attn else causal,
|
causal=False if use_cascade_attn else causal,
|
||||||
window_size=window_size,
|
window_size=window_size,
|
||||||
softcap=layer.logit_cap,
|
softcap=layer.logit_cap,
|
||||||
return_softmax_lse=use_cascade_attn,
|
return_softmax_lse=use_cascade_attn or return_lse,
|
||||||
num_splits=self.num_splits,
|
num_splits=self.num_splits,
|
||||||
out=_fa_out,
|
out=_fa_out,
|
||||||
ver=self.fa_impl_ver,
|
ver=self.fa_impl_ver,
|
||||||
@@ -1623,6 +1624,8 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
o_expand,
|
o_expand,
|
||||||
softmax_lse_expand.T.contiguous(),
|
softmax_lse_expand.T.contiguous(),
|
||||||
)
|
)
|
||||||
|
elif return_lse:
|
||||||
|
o, lse_out, *_ = result
|
||||||
else:
|
else:
|
||||||
o = result
|
o = result
|
||||||
else:
|
else:
|
||||||
@@ -1823,7 +1826,12 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
else:
|
else:
|
||||||
o = result
|
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(
|
def forward_decode(
|
||||||
self,
|
self,
|
||||||
@@ -1844,7 +1852,13 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
aux_tensors=None,
|
aux_tensors=None,
|
||||||
rel_bias=None,
|
rel_bias=None,
|
||||||
rel_bias_event=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:
|
if score_mod is not None and self.fa_impl_ver != 4:
|
||||||
raise RuntimeError("score_mod is only supported by the FA4 backend.")
|
raise RuntimeError("score_mod is only supported by the FA4 backend.")
|
||||||
if k is not None:
|
if k is not None:
|
||||||
@@ -2045,7 +2059,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
causal=False if use_cascade_attn else causal,
|
causal=False if use_cascade_attn else causal,
|
||||||
window_size=window_size,
|
window_size=window_size,
|
||||||
softcap=layer.logit_cap,
|
softcap=layer.logit_cap,
|
||||||
return_softmax_lse=use_cascade_attn,
|
return_softmax_lse=use_cascade_attn or return_lse,
|
||||||
num_splits=(
|
num_splits=(
|
||||||
self.decode_num_splits
|
self.decode_num_splits
|
||||||
if not is_swa_layer
|
if not is_swa_layer
|
||||||
@@ -2090,6 +2104,8 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
o_expand,
|
o_expand,
|
||||||
softmax_lse_expand.T.contiguous(),
|
softmax_lse_expand.T.contiguous(),
|
||||||
)
|
)
|
||||||
|
elif return_lse:
|
||||||
|
o, lse_out, *_ = result
|
||||||
else:
|
else:
|
||||||
o = result
|
o = result
|
||||||
else:
|
else:
|
||||||
@@ -2168,7 +2184,12 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
else:
|
else:
|
||||||
o = result
|
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):
|
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
|
||||||
"""Initialize CUDA graph state for the attention backend.
|
"""Initialize CUDA graph state for the attention backend.
|
||||||
|
|||||||
@@ -3,9 +3,13 @@ import unittest
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
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.ci.ci_register import register_cuda_ci
|
||||||
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
|
from sglang.test.kits.attention_unittest.attention_methods.dense_attention import (
|
||||||
|
DENSE_ATOL,
|
||||||
|
DENSE_RTOL,
|
||||||
DenseAttentionCase,
|
DenseAttentionCase,
|
||||||
|
build_dense_attention_fixture,
|
||||||
make_dense_cases,
|
make_dense_cases,
|
||||||
run_dense_attention_case,
|
run_dense_attention_case,
|
||||||
)
|
)
|
||||||
@@ -423,6 +427,102 @@ class TestFA4DenseAttentionBackendCorrectness(CustomTestCase):
|
|||||||
hidden_size=self.HIDDEN_SIZE,
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user