[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:
co-authored by
Claude Opus 5
parent
60ff1e33a5
commit
7fd5454335
@@ -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"]))
|
||||
|
||||
Reference in New Issue
Block a user