sgl-kernel: bump sgl-attn for varlen num_splits OOM fix (#29551)

This commit is contained in:
Xinyuan Tong
2026-07-01 21:13:53 -07:00
committed by GitHub
parent e6f6a353bf
commit 9ba4b8f8ba
2 changed files with 52 additions and 2 deletions
+2 -2
View File
@@ -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)
+50
View File
@@ -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__]))