[DSA] Route the ragged prefill top-k to the v2 kernel (#35175)

Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
DarkSharpness
2026-08-21 16:59:41 -07:00
committed by GitHub
co-authored by Claude Opus 5
parent 60ff1e33a5
commit 7fd5454335
7 changed files with 420 additions and 15 deletions
@@ -5,6 +5,7 @@ from sglang.kernels.ops.attention.dsv4.topk import (
plan_topk_v2,
topk_transform_512,
topk_transform_512_v2,
topk_transform_ragged_v2,
)
from sglang.test.ci.ci_register import register_cuda_ci
@@ -14,6 +15,8 @@ register_cuda_ci(
# Compressed page size used by the DSA indexer (real value is 256 // 4 = 64).
PAGE_SIZE = 64
# NOTE: currently torch baseline is disabled, since it's too slow
DISABLE_TORCH = True
def _make_inputs(batch_size: int, seq_len: int, k: int):
@@ -44,7 +47,7 @@ def _make_p1_table(batch_size: int, seq_len: int):
return src_page_table, lengths
def _build_fn(provider: str, batch_size: int, seq_len: int, k: int):
def _build_paged_fn(provider: str, batch_size: int, seq_len: int, k: int):
scores, seq_lens, page_table, out = _make_inputs(batch_size, seq_len, k)
N = PAGE_SIZE
@@ -72,19 +75,67 @@ def _build_fn(provider: str, batch_size: int, seq_len: int, k: int):
return fn, (scores, seq_lens, page_table)
def _build_ragged_fn(provider: str, batch_size: int, seq_len: int, k: int):
scores, seq_lens, _, out = _make_inputs(batch_size, seq_len, k)
offsets = torch.arange(batch_size, dtype=torch.int32, device="cuda") * seq_len
def fn(scores, seq_lens, offsets):
if provider == "jit_v1":
from sgl_kernel import fast_topk_transform_ragged_fused
return fast_topk_transform_ragged_fused(scores, seq_lens, offsets, k)
elif provider == "jit_v2":
topk_transform_ragged_v2(
scores, seq_lens, out_offsets=offsets, out_indices=out
)
return out
elif provider == "flashinfer":
from flashinfer import top_k_ragged_transform
return top_k_ragged_transform(scores, offsets, seq_lens, k)
elif provider == "torch":
idx = scores.topk(k, dim=-1).indices.to(torch.int32) # (batch, k)
return idx + offsets.unsqueeze(1)
else:
raise ValueError(f"unknown provider {provider}")
return fn, (scores, seq_lens, offsets)
PRROVIDERS = ["jit_v1", "jit_v2", "flashinfer"]
if not DISABLE_TORCH:
PRROVIDERS.append("torch")
@marker.parametrize("k", [512, 1024, 2048], [512])
@marker.parametrize("seq_len", [2**x for x in range(10, 19)], [4096, 65536])
@marker.parametrize("batch_size", [2**x for x in range(13)], [1, 128, 1024])
@marker.benchmark("provider", ["jit_v1", "jit_v2", "flashinfer", "torch"])
def benchmark(seq_len: int, batch_size: int, k: int, provider: str):
@marker.benchmark("provider", PRROVIDERS)
def benchmark_paged(seq_len: int, batch_size: int, k: int, provider: str):
if k > seq_len:
marker.skip("k cannot be larger than seq_len")
if k == 2048 and provider == "jit_v1":
marker.skip("jit_v1 does not support k=2048")
fn, input_args = _build_fn(provider, batch_size, seq_len, k)
fn, input_args = _build_paged_fn(provider, batch_size, seq_len, k)
return marker.do_bench(fn, input_args=input_args, memory_args=input_args[:2])
@marker.parametrize("k", [512, 1024, 2048], [2048])
@marker.parametrize("seq_len", [2**x for x in range(10, 19)], [4096, 65536])
# NOTE: prefill workload should be heavier than decode; not common for short extend
@marker.parametrize("batch_size", [2**x for x in range(7, 14)], [128, 1024])
@marker.benchmark("provider", PRROVIDERS)
def benchmark_ragged(seq_len: int, batch_size: int, k: int, provider: str):
if k > seq_len:
marker.skip("k cannot be larger than seq_len")
if k != 2048 and provider == "jit_v1":
marker.skip("jit_v1 (here, sgl-AOT) only support k=2048")
fn, input_args = _build_ragged_fn(provider, batch_size, seq_len, k)
return marker.do_bench(fn, input_args=input_args, memory_args=input_args[:2])
if __name__ == "__main__":
benchmark.run()
benchmark_paged.run()
benchmark_ragged.run()
@@ -29,7 +29,11 @@ import sys
import pytest
import torch
from sglang.kernels.ops.attention.dsv4.topk import plan_topk_v2, topk_transform_512_v2
from sglang.kernels.ops.attention.dsv4.topk import (
plan_topk_v2,
topk_transform_512_v2,
topk_transform_ragged_v2,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=90, stage="base-b-kernel-unit", runner_config="1-gpu-large")
@@ -265,5 +269,112 @@ def test_topk_v2_output_indices(batch: int, seq: int, k: int) -> None:
_assert_topk_close(scores.cpu(), ref_raw, our_raw, batch, seq_lens.cpu(), k)
# --- ragged entry point ------------------------------------------------------
# Rows select inside `[row_start, row_start + seq_len)` of their score row and
# emit `position + offset`. The window start is an arbitrary token offset, so
# every `row_start % 4` residue must be covered: the kernel reads from a
# 16-byte-aligned base and masks the <=3 columns that pulls in ahead of the
# window. Everything outside the window is filled with OUTSIDE_SCORE, which
# beats every in-window score, so any leak shows up as a wrong selection.
OUTSIDE_SCORE = 1e3
# (name, per-row (row_start, length)) spanning every template and residue.
RAGGED_CONFIGS = [
# one length per template band, all four residues plus aligned starts
("trivial", [(s, 1500) for s in (0, 1, 2, 3, 4, 7, 4096, 4099)]),
("register2", [(s, 6000) for s in (0, 1, 2, 3, 4, 7, 4096, 4099)]),
("register4", [(s, 12000) for s in (0, 1, 2, 3, 4, 7, 4096, 4099)]),
("streaming", [(s, 40000) for s in (0, 1, 2, 3, 4, 7, 4096, 4099)]),
# mixed bands in one launch, laid out back to back like a real prefill batch
(
"mixed",
[
(0, 1000),
(1000, 3000),
(4000, 9000),
(13000, 20000),
(33000, 1),
(33001, 2047),
],
),
# boundaries: seq == k, seq == k + 1, and the register/streaming edges
(
"boundaries",
[(1, 2048), (2049, 2049), (4098, 8192), (12290, 8193), (20483, 16385)],
),
("long_ctx", [(0, 131072), (131072, 65537), (196609, 100000)]),
]
def _make_ragged(rows, offset_shift, device):
width = ((max(s + n for s, n in rows)) + 3) & ~3
scores = torch.full(
(len(rows), width), OUTSIDE_SCORE, dtype=torch.float32, device=device
)
for i, (start, length) in enumerate(rows):
scores[i, start : start + length] = torch.randn(length, device=device)
starts = torch.tensor([s for s, _ in rows], dtype=torch.int32, device=device)
lengths = torch.tensor([n for _, n in rows], dtype=torch.int32, device=device)
return scores, starts, lengths, starts + offset_shift
def _run_ragged(scores, lengths, starts, offsets, k):
"""Selected positions per row, rebased back to window-relative."""
out = torch.empty((scores.shape[0], k), dtype=torch.int32, device=scores.device)
topk_transform_ragged_v2(
scores, lengths, out_offsets=offsets, out_indices=out, row_starts=starts
)
torch.cuda.synchronize()
off = offsets.cpu().tolist()
return [
[v - off[i] for v in row if v != -1] for i, row in enumerate(out.cpu().tolist())
]
@pytest.mark.parametrize("k", [512, 1024, 2048])
@pytest.mark.parametrize("offset_shift", [0, 4321])
@pytest.mark.parametrize("name,rows", RAGGED_CONFIGS)
@torch.inference_mode()
def test_topk_v2_ragged_window(name: str, rows, k: int, offset_shift: int) -> None:
torch.manual_seed(len(rows) * 7919 + k + offset_shift)
device = "cuda"
scores, starts, lengths, offsets = _make_ragged(rows, offset_shift, device)
before = scores.clone()
our_raw = _run_ragged(scores, lengths, starts, offsets, k)
# reference on the window slice, padded to a common width for the helper
max_len = max(n for _, n in rows)
windows = torch.zeros(len(rows), max_len, dtype=torch.float32)
for i, (start, length) in enumerate(rows):
windows[i, :length] = before[i, start : start + length].cpu()
ref_raw = _reference(windows, lengths.cpu(), k)
_assert_topk_close(windows, ref_raw, our_raw, len(rows), lengths.cpu(), k)
# the only legal in-place write is the <=3 masked columns ahead of a window
# that the kernel actually reads (trivial rows read nothing)
changed = (scores != before).cpu()
for i, (start, length) in enumerate(rows):
allowed = torch.zeros(scores.shape[1], dtype=torch.bool)
if length > k:
allowed[start - start % 4 : start] = True
stray = (changed[i] & ~allowed).nonzero().flatten().tolist()
assert not stray, f"row {i} ({name}) wrote outside its masked head: {stray[:8]}"
@pytest.mark.parametrize("k", [512, 2048])
@torch.inference_mode()
def test_topk_v2_ragged_no_row_starts(k: int) -> None:
"""`row_starts=None` means every window starts at column 0."""
torch.manual_seed(4242 + k)
device = "cuda"
rows = [(0, 900), (0, 5000), (0, 20000), (0, 70000)]
scores, starts, lengths, offsets = _make_ragged(rows, 0, device)
explicit = _run_ragged(scores.clone(), lengths, starts, offsets, k)
implicit = _run_ragged(scores.clone(), lengths, None, offsets, k)
for i in range(len(rows)):
assert sorted(explicit[i]) == sorted(implicit[i]), f"row {i} differs"
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v"]))