[FlashInfer V0.6.18] feat(dsv4): support --dsa-topk-backend flashinfer with fused top-k (#33237)

This commit is contained in:
Ziang Li
2026-09-01 01:18:10 -07:00
committed by GitHub
parent 3484f7f836
commit 5edcd0a445
8 changed files with 333 additions and 31 deletions
@@ -19,11 +19,11 @@ PAGE_SIZE = 64
DISABLE_TORCH = True
def _make_inputs(batch_size: int, seq_len: int, k: int):
def _make_inputs(batch_size: int, seq_len: int, k: int, page_size: int = PAGE_SIZE):
torch.random.manual_seed(42)
scores = torch.randn(batch_size, seq_len, dtype=torch.float32, device="cuda")
seq_lens = torch.full((batch_size,), seq_len, dtype=torch.int32, device="cuda")
num_pages = (seq_len + PAGE_SIZE - 1) // PAGE_SIZE
num_pages = (seq_len + page_size - 1) // page_size
page_table = (
torch.arange(num_pages, dtype=torch.int32, device="cuda")
.unsqueeze(0)
@@ -34,22 +34,11 @@ def _make_inputs(batch_size: int, seq_len: int, k: int):
return scores, seq_lens, page_table, out
def _make_p1_table(batch_size: int, seq_len: int):
# flashinfer / torch do a per-token (page_size=1) gather, so they need a
# (batch, seq) table (one entry per position) rather than the page-size-64 one.
src_page_table = (
torch.arange(seq_len, dtype=torch.int32, device="cuda")
.unsqueeze(0)
.expand(batch_size, -1)
.contiguous()
)
lengths = torch.full((batch_size,), seq_len, dtype=torch.int32, device="cuda")
return src_page_table, lengths
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
def _build_paged_fn(
provider: str, batch_size: int, seq_len: int, k: int, page_size: int
):
scores, seq_lens, page_table, out = _make_inputs(batch_size, seq_len, k, page_size)
N = page_size
def fn(scores, seq_lens, page_table):
if provider == "jit_v1":
@@ -61,15 +50,23 @@ def _build_paged_fn(provider: str, batch_size: int, seq_len: int, k: int):
elif provider == "flashinfer":
from flashinfer import top_k_page_table_transform
return top_k_page_table_transform(scores, page_table, seq_lens, k)
return top_k_page_table_transform(
scores,
page_table,
seq_lens,
k,
dsa_graph_safe=True,
page_size=N,
out=out,
)
elif provider == "torch":
idx = scores.topk(k, dim=-1).indices # (batch, k) int64
return torch.gather(page_table, 1, idx)
raw_indices = scores.topk(k, dim=-1).indices
physical_pages = torch.gather(page_table, 1, raw_indices // N)
out.copy_(physical_pages * N + raw_indices % N)
return out
else:
raise ValueError(f"unknown provider {provider}")
if provider in ("flashinfer", "torch"):
page_table, seq_lens = _make_p1_table(batch_size, seq_len)
if provider == "jit_v2":
metadata = plan_topk_v2(seq_lens)
return fn, (scores, seq_lens, page_table)
@@ -110,14 +107,17 @@ if not DISABLE_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.parametrize("page_size", [1, 64], [1, 64])
@marker.benchmark("provider", PRROVIDERS)
def benchmark_paged(seq_len: int, batch_size: int, k: int, provider: str):
def benchmark_paged(
seq_len: int, batch_size: int, k: int, page_size: 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_paged_fn(provider, batch_size, seq_len, k)
fn, input_args = _build_paged_fn(provider, batch_size, seq_len, k, page_size)
return marker.do_bench(fn, input_args=input_args, memory_args=input_args[:2])
@@ -1,4 +1,5 @@
import unittest
from itertools import product
from typing import List, Optional, Tuple
from unittest.mock import MagicMock, patch
@@ -25,6 +26,7 @@ from sglang.srt.layers.attention.dsa_backend import (
DeepseekSparseAttnBackend,
DSAMetadata,
)
from sglang.srt.layers.attention.dsv4.indexer import C4IndexerBackendMixin
from sglang.srt.layers.layernorm import LayerNorm
from sglang.srt.layers.linear import LinearBase
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
@@ -32,7 +34,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=18, stage="base-b", runner_config="1-gpu-large")
register_cuda_ci(est_time=30, stage="base-b", runner_config="1-gpu-large")
# Global configuration for all indexer tests
DEFAULT_CONFIG = {
@@ -450,6 +452,151 @@ class TestDSAIndexer(CustomTestCase):
index=perm,
)
def _assert_dsv4_compact_topk_result(
self,
scores: torch.Tensor,
seq_lens: torch.Tensor,
page_table: torch.Tensor,
out_page_indices: torch.Tensor,
out_raw_indices: Optional[torch.Tensor],
page_size: int,
) -> None:
topk = out_page_indices.shape[1]
for row in range(scores.shape[0]):
seq_len = int(seq_lens[row].item())
expected_count = min(seq_len, topk)
if seq_len <= topk:
expected_raw = torch.arange(
seq_len, dtype=torch.int32, device=scores.device
)
else:
expected_raw = torch.topk(
scores[row, :seq_len], topk, sorted=False
).indices.to(torch.int32)
expected_page_indices = (
page_table[row, expected_raw.long() // page_size] * page_size
+ expected_raw % page_size
)
actual_page_indices = out_page_indices[row]
actual_page_prefix = actual_page_indices[:expected_count]
self.assertTrue(
torch.equal(
torch.sort(actual_page_prefix).values,
torch.sort(expected_page_indices).values,
)
)
self.assertTrue(torch.all(actual_page_indices[expected_count:] == -1))
if out_raw_indices is not None:
actual_raw = out_raw_indices[row]
actual_raw_prefix = actual_raw[:expected_count]
self.assertTrue(
torch.equal(
torch.sort(actual_raw_prefix).values,
torch.sort(expected_raw).values,
)
)
self.assertTrue(torch.all(actual_raw[expected_count:] == -1))
translated_raw = (
page_table[row, actual_raw_prefix.long() // page_size] * page_size
+ actual_raw_prefix % page_size
)
self.assertTrue(torch.equal(actual_page_prefix, translated_raw))
def test_dsv4_flashinfer_compact_topk_cuda_graph(self):
num_rows, max_len, page_size = 4, 2048, 64
score_storage = torch.empty(
(num_rows, max_len + 16), dtype=torch.float32, device=self.device
)
scores = score_storage[:, 1 : max_len + 1]
self.assertFalse(scores.is_contiguous())
page_table = 17 + torch.arange(
num_rows * (max_len // page_size),
dtype=torch.int32,
device=self.device,
).reshape(num_rows, max_len // page_size)
for fuse_topk, topk, with_raw_output in product(
(False, True), (512, 1024), (False, True)
):
with self.subTest(
fuse_topk=fuse_topk,
topk=topk,
with_raw_output=with_raw_output,
):
scores.copy_(
torch.arange(max_len, dtype=torch.float32, device=self.device)
.unsqueeze(0)
.expand(num_rows, -1)
)
seq_lens = torch.full(
(num_rows,), max_len, dtype=torch.int32, device=self.device
)
out_page_indices = torch.empty(
(num_rows, topk), dtype=torch.int32, device=self.device
)
out_raw_indices = (
torch.empty_like(out_page_indices) if with_raw_output else None
)
# Capture with every output slot valid so shrinking lengths on
# replay leaves meaningful stale values in any unwritten suffix.
with (
envs.SGLANG_DSA_FUSE_TOPK.override(fuse_topk),
envs.SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC.override(True),
envs.SGLANG_DSA_TOPK_FLASHINFER_TIE_BREAK.override("small"),
):
backend = C4IndexerBackendMixin()
backend.flashinfer_topk_transform(
scores,
seq_lens,
page_table,
out_page_indices,
page_size,
out_raw_indices,
)
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
backend.flashinfer_topk_transform(
scores,
seq_lens,
page_table,
out_page_indices,
page_size,
out_raw_indices,
)
scores.copy_(
torch.arange(
max_len, 0, -1, dtype=torch.float32, device=self.device
)
.unsqueeze(0)
.expand(num_rows, -1)
)
page_table.add_(100)
seq_lens.copy_(
torch.tensor(
[0, topk - 1, topk + 73, max_len],
dtype=torch.int32,
device=self.device,
)
)
graph.replay()
torch.cuda.synchronize()
self._assert_dsv4_compact_topk_result(
scores,
seq_lens,
page_table,
out_page_indices,
out_raw_indices,
page_size,
)
def _run_unfused_topk_backend_validity_test(
self,
batch_size: int,