sgl-kernel: bump sgl-attn for varlen num_splits OOM fix (#29551)
This commit is contained in:
@@ -81,8 +81,8 @@ FetchContent_Populate(repo-flashinfer)
|
||||
# flash-attention
|
||||
FetchContent_Declare(
|
||||
repo-flash-attention
|
||||
URL https://${GITHUB_ARTIFACTORY}/sgl-project/sgl-attn/archive/65c54cc5a6d29fee56036484c749bc5b8e00fd66.tar.gz
|
||||
URL_HASH SHA256=53307eefba65d9ac26092433a8bc4d48182ce0bb27132e69ab0b8d6ff8ee47db
|
||||
URL https://${GITHUB_ARTIFACTORY}/sgl-project/sgl-attn/archive/f89bc2306632d1ec5f97b014dded4254f5b4a907.tar.gz
|
||||
URL_HASH SHA256=418b5681584dc3efff496a1cab5ffd58d2728d89dcfe0ea16e6985d6ef35c68c
|
||||
)
|
||||
FetchContent_Populate(repo-flash-attention)
|
||||
|
||||
|
||||
@@ -1365,5 +1365,55 @@ def test_flash_attn_varlen_output(
|
||||
).abs().max().item() + dv_atol
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not is_fa3_supported(),
|
||||
reason="flash_attn at sgl-kernel is only supported on sm90 or sm80",
|
||||
)
|
||||
@pytest.mark.parametrize("dtype", [torch.bfloat16])
|
||||
def test_flash_attn_varlen_many_short_segments_num_splits(dtype):
|
||||
# Regression for the varlen get_num_splits heuristic (sgl-attn #46). A batch of
|
||||
# many short-Q segments (NSA-style cu_seqlens_q = arange(N+1)) over long K made the
|
||||
# old "pretend batch=1" upper bound under-count total_mblocks and over-split, blowing
|
||||
# up the fp32 out_accum workspace (OOM). seqlen_k is large enough that num_n_blocks > 4
|
||||
# so the heuristic's split loop actually runs; with num_splits=0 the output must still
|
||||
# match a per-segment reference.
|
||||
from sgl_kernel.flash_attn import flash_attn_varlen_func
|
||||
|
||||
torch.manual_seed(0)
|
||||
device = "cuda"
|
||||
n_seg, seqlen_k = 64, 2048
|
||||
nheads, nheads_kv, d = 8, 1, 128
|
||||
scale = d**-0.5
|
||||
|
||||
q = torch.randn(n_seg, nheads, d, device=device, dtype=dtype)
|
||||
k = torch.randn(n_seg * seqlen_k, nheads_kv, d, device=device, dtype=dtype)
|
||||
v = torch.randn(n_seg * seqlen_k, nheads_kv, d, device=device, dtype=dtype)
|
||||
cu_seqlens_q = torch.arange(n_seg + 1, device=device, dtype=torch.int32)
|
||||
cu_seqlens_k = torch.arange(
|
||||
0, n_seg * seqlen_k + 1, seqlen_k, device=device, dtype=torch.int32
|
||||
)
|
||||
|
||||
out, lse, *rest = flash_attn_varlen_func(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
cu_seqlens_q,
|
||||
cu_seqlens_k,
|
||||
max_seqlen_q=1,
|
||||
max_seqlen_k=seqlen_k,
|
||||
causal=False,
|
||||
softmax_scale=scale,
|
||||
num_splits=0, # exercise the get_num_splits heuristic
|
||||
return_softmax_lse=True,
|
||||
)
|
||||
|
||||
qf = q.float()
|
||||
kf = k.float().view(n_seg, seqlen_k, nheads_kv, d)[:, :, 0]
|
||||
vf = v.float().view(n_seg, seqlen_k, nheads_kv, d)[:, :, 0]
|
||||
probs = (torch.einsum("nhd,nkd->nhk", qf, kf) * scale).softmax(dim=-1)
|
||||
ref = torch.einsum("nhk,nkd->nhd", probs, vf)
|
||||
torch.testing.assert_close(out.float(), ref, atol=2e-2, rtol=2e-2)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__]))
|
||||
|
||||
Reference in New Issue
Block a user