[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
+5 -1
View File
@@ -1698,7 +1698,11 @@ def _set_envs_and_config(server_args: ServerArgs):
# Check flashinfer version # Check flashinfer version
if not get_bool_env_var("SGLANG_SKIP_SGL_KERNEL_VERSION_CHECK"): if not get_bool_env_var("SGLANG_SKIP_SGL_KERNEL_VERSION_CHECK"):
if "flashinfer" in attention_backends_of(resolved_view(cfg)): if (
"flashinfer" in attention_backends_of(resolved_view(cfg))
or cfg.dsa_topk_backend == "flashinfer"
or cfg.speculative_dsa_topk_backend == "flashinfer"
):
assert_pkg_version( assert_pkg_version(
"flashinfer_python", "flashinfer_python",
"0.6.18", "0.6.18",
@@ -704,6 +704,7 @@ class DeepseekV4AttnBackend(
page_size=self.page_size, page_size=self.page_size,
page_table=core_attn_metadata.page_table, page_table=core_attn_metadata.page_table,
c4_seq_lens=core_attn_metadata.c4_topk_lengths_raw, c4_seq_lens=core_attn_metadata.c4_topk_lengths_raw,
use_topk_v2=self.dsa_topk_backend.should_use_topk_v2() and not _is_xpu,
# The SM120 FP4 kernel schedules split_kv=128, while the generic # The SM120 FP4 kernel schedules split_kv=128, while the generic
# JIT metadata planner encodes split_kv=256. # JIT metadata planner encodes split_kv=256.
force_deep_gemm_metadata=( force_deep_gemm_metadata=(
@@ -492,6 +492,7 @@ class DeepseekV4HipRadixBackend(
page_size=self.page_size, page_size=self.page_size,
page_table=core_attn_metadata.page_table, page_table=core_attn_metadata.page_table,
c4_seq_lens=core_attn_metadata.c4_topk_lengths_raw, c4_seq_lens=core_attn_metadata.c4_topk_lengths_raw,
use_topk_v2=False,
) )
def init_forward_metadata_decode( def init_forward_metadata_decode(
@@ -404,11 +404,44 @@ def topk_transform_512_flashinfer_unfused(
) )
def topk_transform_512_flashinfer_fused(
scores: torch.Tensor,
seq_lens: torch.Tensor,
page_tables: torch.Tensor,
out_page_indices: torch.Tensor,
page_size: int,
out_raw_indices: Optional[torch.Tensor] = None,
) -> None:
import flashinfer
from sglang.srt.layers.attention.dsa.dsa_topk_backend import (
_flashinfer_tie_break_value,
)
flashinfer.top_k_page_table_transform(
scores,
page_tables.contiguous(),
seq_lens.contiguous(),
out_page_indices.shape[1],
deterministic=envs.SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC.get(),
tie_break=_flashinfer_tie_break_value(),
dsa_graph_safe=True,
page_size=page_size,
out=out_page_indices,
out_raw_indices=out_raw_indices,
)
class C4IndexerBackendMixin: class C4IndexerBackendMixin:
def __init__(self): def __init__(self):
super().__init__() super().__init__()
self.debug_use_external_c4_sparse_indices: bool = False self.debug_use_external_c4_sparse_indices: bool = False
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend.SGL_KERNEL self.dsa_topk_backend: DSATopKBackend = DSATopKBackend.SGL_KERNEL
self.flashinfer_topk_transform: Callable[..., None] = (
topk_transform_512_flashinfer_fused
if envs.SGLANG_DSA_FUSE_TOPK.get()
else topk_transform_512_flashinfer_unfused
)
def _forward_prepare_multi_stream( def _forward_prepare_multi_stream(
self, self,
@@ -829,7 +862,7 @@ class C4IndexerBackendMixin:
raw_indices, raw_indices,
) )
elif self.dsa_topk_backend.is_flashinfer(): elif self.dsa_topk_backend.is_flashinfer():
topk_transform_512_flashinfer_unfused( self.flashinfer_topk_transform(
logits, logits,
c4_seq_lens, c4_seq_lens,
page_table, page_table,
@@ -837,7 +870,7 @@ class C4IndexerBackendMixin:
indexer_metadata.c4_page_size, indexer_metadata.c4_page_size,
raw_indices, raw_indices,
) )
elif envs.SGLANG_OPT_USE_TOPK_V2.get() and raw_indices is None: elif self.dsa_topk_backend.should_use_topk_v2() and raw_indices is None:
topk_transform_512_v2( topk_transform_512_v2(
logits, logits,
c4_seq_lens, c4_seq_lens,
@@ -112,6 +112,7 @@ class PagedIndexerMetadata:
page_size: int page_size: int
page_table: torch.Tensor page_table: torch.Tensor
c4_seq_lens: torch.Tensor c4_seq_lens: torch.Tensor
use_topk_v2: bool
force_deep_gemm_metadata: bool = False force_deep_gemm_metadata: bool = False
use_prefill_cuda_graph: bool = False use_prefill_cuda_graph: bool = False
deep_gemm_metadata: Any = field(init=False, repr=False) deep_gemm_metadata: Any = field(init=False, repr=False)
@@ -152,7 +153,7 @@ class PagedIndexerMetadata:
assert isinstance(self.deep_gemm_metadata, torch.Tensor) assert isinstance(self.deep_gemm_metadata, torch.Tensor)
if envs.SGLANG_OPT_USE_TOPK_V2.get() and not is_xpu(): if self.use_topk_v2:
from sglang.kernels.ops.attention.dsv4 import plan_topk_v2 from sglang.kernels.ops.attention.dsv4 import plan_topk_v2
self.topk_metadata = plan_topk_v2(self.c4_seq_lens) self.topk_metadata = plan_topk_v2(self.c4_seq_lens)
@@ -188,6 +189,7 @@ class PagedIndexerMetadata:
"page_size", "page_size",
"force_deep_gemm_metadata", "force_deep_gemm_metadata",
"use_prefill_cuda_graph", "use_prefill_cuda_graph",
"use_topk_v2",
], ],
copy_fields=copy_fields, copy_fields=copy_fields,
assign_fields=assign_fields, assign_fields=assign_fields,
@@ -19,11 +19,11 @@ PAGE_SIZE = 64
DISABLE_TORCH = True 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) torch.random.manual_seed(42)
scores = torch.randn(batch_size, seq_len, dtype=torch.float32, device="cuda") 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") 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 = ( page_table = (
torch.arange(num_pages, dtype=torch.int32, device="cuda") torch.arange(num_pages, dtype=torch.int32, device="cuda")
.unsqueeze(0) .unsqueeze(0)
@@ -34,22 +34,11 @@ def _make_inputs(batch_size: int, seq_len: int, k: int):
return scores, seq_lens, page_table, out return scores, seq_lens, page_table, out
def _make_p1_table(batch_size: int, seq_len: int): def _build_paged_fn(
# flashinfer / torch do a per-token (page_size=1) gather, so they need a provider: str, batch_size: int, seq_len: int, k: int, page_size: int
# (batch, seq) table (one entry per position) rather than the page-size-64 one. ):
src_page_table = ( scores, seq_lens, page_table, out = _make_inputs(batch_size, seq_len, k, page_size)
torch.arange(seq_len, dtype=torch.int32, device="cuda") N = page_size
.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 fn(scores, seq_lens, page_table): def fn(scores, seq_lens, page_table):
if provider == "jit_v1": 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": elif provider == "flashinfer":
from flashinfer import top_k_page_table_transform 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": elif provider == "torch":
idx = scores.topk(k, dim=-1).indices # (batch, k) int64 raw_indices = scores.topk(k, dim=-1).indices
return torch.gather(page_table, 1, idx) physical_pages = torch.gather(page_table, 1, raw_indices // N)
out.copy_(physical_pages * N + raw_indices % N)
return out
else: else:
raise ValueError(f"unknown provider {provider}") 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": if provider == "jit_v2":
metadata = plan_topk_v2(seq_lens) metadata = plan_topk_v2(seq_lens)
return fn, (scores, seq_lens, page_table) return fn, (scores, seq_lens, page_table)
@@ -110,14 +107,17 @@ if not DISABLE_TORCH:
@marker.parametrize("k", [512, 1024, 2048], [512]) @marker.parametrize("k", [512, 1024, 2048], [512])
@marker.parametrize("seq_len", [2**x for x in range(10, 19)], [4096, 65536]) @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("batch_size", [2**x for x in range(13)], [1, 128, 1024])
@marker.parametrize("page_size", [1, 64], [1, 64])
@marker.benchmark("provider", PRROVIDERS) @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: if k > seq_len:
marker.skip("k cannot be larger than seq_len") marker.skip("k cannot be larger than seq_len")
if k == 2048 and provider == "jit_v1": if k == 2048 and provider == "jit_v1":
marker.skip("jit_v1 does not support k=2048") 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]) return marker.do_bench(fn, input_args=input_args, memory_args=input_args[:2])
@@ -1,4 +1,5 @@
import unittest import unittest
from itertools import product
from typing import List, Optional, Tuple from typing import List, Optional, Tuple
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
@@ -25,6 +26,7 @@ from sglang.srt.layers.attention.dsa_backend import (
DeepseekSparseAttnBackend, DeepseekSparseAttnBackend,
DSAMetadata, DSAMetadata,
) )
from sglang.srt.layers.attention.dsv4.indexer import C4IndexerBackendMixin
from sglang.srt.layers.layernorm import LayerNorm from sglang.srt.layers.layernorm import LayerNorm
from sglang.srt.layers.linear import LinearBase from sglang.srt.layers.linear import LinearBase
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool 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.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.test.test_utils import CustomTestCase 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 # Global configuration for all indexer tests
DEFAULT_CONFIG = { DEFAULT_CONFIG = {
@@ -450,6 +452,151 @@ class TestDSAIndexer(CustomTestCase):
index=perm, 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( def _run_unfused_topk_backend_validity_test(
self, self,
batch_size: int, batch_size: int,
@@ -1,12 +1,16 @@
import sys import sys
import unittest import unittest
from itertools import product
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
import torch import torch
from sglang.srt.environ import envs 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 ( from sglang.srt.layers.attention.dsv4.metadata import (
NonPagedIndexerPlan, NonPagedIndexerPlan,
PagedIndexerMetadata, PagedIndexerMetadata,
@@ -43,6 +47,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
page_size=256, page_size=256,
page_table=torch.zeros((1, 1), dtype=torch.int32), page_table=torch.zeros((1, 1), dtype=torch.int32),
c4_seq_lens=torch.tensor([65], dtype=torch.int32), c4_seq_lens=torch.tensor([65], dtype=torch.int32),
use_topk_v2=False,
force_deep_gemm_metadata=True, force_deep_gemm_metadata=True,
) )
@@ -64,10 +69,119 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
page_size=256, page_size=256,
page_table=torch.zeros((1, 1), dtype=torch.int32), page_table=torch.zeros((1, 1), dtype=torch.int32),
c4_seq_lens=torch.tensor([65], dtype=torch.int32), c4_seq_lens=torch.tensor([65], dtype=torch.int32),
use_topk_v2=False,
) )
self.assertIsNone(metadata.deep_gemm_metadata) 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): class TestDSV4NonPagedIndexer(CustomTestCase):
def _is_eligible(self, **overrides): def _is_eligible(self, **overrides):