[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,
@@ -1,12 +1,16 @@
import sys
import unittest
from itertools import product
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import torch
from sglang.srt.environ import envs
from sglang.srt.layers.attention.dsv4.indexer import FP8_DTYPE, C4IndexerBackendMixin
from sglang.srt.layers.attention.dsv4.indexer import (
FP8_DTYPE,
C4IndexerBackendMixin,
)
from sglang.srt.layers.attention.dsv4.metadata import (
NonPagedIndexerPlan,
PagedIndexerMetadata,
@@ -43,6 +47,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
page_size=256,
page_table=torch.zeros((1, 1), dtype=torch.int32),
c4_seq_lens=torch.tensor([65], dtype=torch.int32),
use_topk_v2=False,
force_deep_gemm_metadata=True,
)
@@ -64,10 +69,119 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
page_size=256,
page_table=torch.zeros((1, 1), dtype=torch.int32),
c4_seq_lens=torch.tensor([65], dtype=torch.int32),
use_topk_v2=False,
)
self.assertIsNone(metadata.deep_gemm_metadata)
def test_topk_v2_ineligible_backend_skips_plan(self):
with (
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(True),
envs.SGLANG_OPT_USE_AITER_INDEXER.override(False),
envs.SGLANG_OPT_USE_TOPK_V2.override(True),
patch("sglang.kernels.ops.attention.dsv4.plan_topk_v2") as plan_topk_v2,
):
metadata = PagedIndexerMetadata(
page_size=256,
page_table=torch.zeros((1, 1), dtype=torch.int32),
c4_seq_lens=torch.tensor([65], dtype=torch.int32),
use_topk_v2=False,
)
plan_topk_v2.assert_not_called()
self.assertEqual(metadata.topk_metadata.numel(), 0)
class TestDSV4FlashInferTopK(CustomTestCase):
def test_compact_page_transform_respects_fuse_topk(self):
score_storage = torch.arange(160, dtype=torch.float32).reshape(2, 80)
scores = score_storage[:, 1:65]
self.assertFalse(scores.is_contiguous())
seq_lens = torch.tensor([63, 64], dtype=torch.int32)
page_tables = torch.tensor(
[[7, 17, 8, 18], [11, 21, 12, 22]], dtype=torch.int32
)[:, ::2]
self.assertFalse(page_tables.is_contiguous())
out_page_indices = torch.empty((2, 8), dtype=torch.int32)
for fuse_topk, with_raw_output in product((False, True), repeat=2):
with self.subTest(fuse_topk=fuse_topk, with_raw_output=with_raw_output):
out_raw_indices = (
torch.empty_like(out_page_indices) if with_raw_output else None
)
def fake_top_k(input: torch.Tensor, k: int, **kwargs):
return torch.topk(
input,
k,
dim=-1,
largest=True,
sorted=kwargs["sorted"],
)
top_k = MagicMock(side_effect=fake_top_k)
top_k_page_table_transform = MagicMock()
flashinfer = SimpleNamespace(
top_k=top_k,
top_k_page_table_transform=top_k_page_table_transform,
)
with (
patch.dict(sys.modules, {"flashinfer": flashinfer}),
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_tables,
out_page_indices,
page_size=64,
out_raw_indices=out_raw_indices,
)
if not fuse_topk:
top_k_page_table_transform.assert_not_called()
top_k.assert_called_once()
call = top_k.call_args
self.assertTrue(call.args[0].is_contiguous())
self.assertEqual(call.args[0].shape, scores.shape)
self.assertEqual(call.args[1], out_page_indices.shape[1])
self.assertEqual(
call.kwargs,
{
"sorted": False,
"deterministic": True,
"tie_break": 1,
"dsa_graph_safe": True,
},
)
continue
top_k.assert_not_called()
top_k_page_table_transform.assert_called_once()
call = top_k_page_table_transform.call_args
self.assertIs(call.args[0], scores)
self.assertIsNot(call.args[1], page_tables)
self.assertTrue(call.args[1].is_contiguous())
self.assertTrue(torch.equal(call.args[1], page_tables))
self.assertIs(call.args[2], seq_lens)
self.assertEqual(call.args[3], out_page_indices.shape[1])
self.assertEqual(
call.kwargs,
{
"deterministic": True,
"tie_break": 1,
"dsa_graph_safe": True,
"page_size": 64,
"out": out_page_indices,
"out_raw_indices": out_raw_indices,
},
)
class TestDSV4NonPagedIndexer(CustomTestCase):
def _is_eligible(self, **overrides):