[DSV4] Support raw-index output in TopK v2 (#33672)

Co-authored-by: weireweire <20922698+weireweire@users.noreply.github.com>
Co-authored-by: Brayden Zhong <b8zhong@uwaterloo.ca>
Co-authored-by: Po-Han Huang (NVIDIA) <53919306+nvpohanh@users.noreply.github.com>
This commit is contained in:
weireweire
2026-09-11 07:17:36 -07:00
committed by GitHub
co-authored by weireweire Brayden Zhong Po-Han Huang
parent ab9750fb35
commit 335f6aab27
5 changed files with 167 additions and 18 deletions
@@ -183,6 +183,20 @@ def _run_raw(scores, seq_lens, k):
return [[v for v in out_cpu[i] if v != -1] for i in range(batch)]
def _run_dual(scores, seq_lens, page_table, inv_cpu, k):
batch = scores.shape[0]
metadata = _plan(seq_lens)
out = torch.full((batch, k), -1, dtype=torch.int32, device=scores.device)
raw = torch.full_like(out, -1)
topk_transform_paged_v2(scores, seq_lens, page_table, out, PAGE_SIZE, metadata, raw)
torch.cuda.synchronize()
out_cpu = out.cpu().tolist()
raw_cpu = raw.cpu().tolist()
transformed_raw = [_invert(out_cpu[i], inv_cpu[i]) for i in range(batch)]
direct_raw = [[v for v in raw_cpu[i] if v != -1] for i in range(batch)]
return transformed_raw, direct_raw
@pytest.mark.parametrize("page_mode", ["identity", "perm"])
@pytest.mark.parametrize("k", [512, 1024, 2048])
@pytest.mark.parametrize("batch,seq", FIXED_CONFIGS)
@@ -270,6 +284,27 @@ 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)
@pytest.mark.parametrize(
"batch,seq", [(8, 256), (8, 8192), (4, 32768), (2, 131072), (31, 131072)]
)
@torch.inference_mode()
def test_topk_v2_dual_output(batch: int, seq: int) -> None:
"""The dual mode returns the same selection before and after page transform."""
k = 512
torch.manual_seed(batch * 100003 + seq * 7 + k + 2)
device = "cuda"
scores = torch.randn(batch, seq, dtype=torch.float32, device=device)
seq_lens = torch.full((batch,), seq, dtype=torch.int32, device=device)
num_pages = (seq + PAGE_SIZE - 1) // PAGE_SIZE
page_table, inv_cpu = _make_page_table(batch, num_pages, "perm", device)
transformed_raw, direct_raw = _run_dual(scores, seq_lens, page_table, inv_cpu, k)
for row in range(batch):
assert sorted(transformed_raw[row]) == sorted(direct_raw[row])
ref_raw = _reference(scores, seq_lens, k)
_assert_topk_close(scores.cpu(), ref_raw, direct_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
@@ -7,6 +7,7 @@ from unittest.mock import MagicMock, patch
import torch
from sglang.srt.environ import envs
from sglang.srt.layers.attention.dsa.dsa_topk_backend import DSATopKBackend
from sglang.srt.layers.attention.dsv4.indexer import (
FP8_DTYPE,
C4IndexerBackendMixin,
@@ -209,6 +210,78 @@ class TestDSV4FlashInferTopK(CustomTestCase):
)
class TestDSV4TopKDispatch(CustomTestCase):
def test_v2_raw_output_uses_sparse_prefill_buffer_with_capture(self):
page_table = torch.zeros((1, 1), dtype=torch.int32)
c4_seq_lens = torch.ones(1, dtype=torch.int32)
page_indices = torch.full((1, 512), -1, dtype=torch.int32)
raw_indices = torch.full_like(page_indices, -1)
topk_metadata = torch.zeros((2, 2), dtype=torch.int32)
indexer_metadata = object.__new__(PagedIndexerMetadata)
indexer_metadata.page_size = 256
indexer_metadata.page_table = page_table
indexer_metadata.c4_seq_lens = c4_seq_lens
indexer_metadata.topk_metadata = topk_metadata
logits = torch.empty((1, 65), dtype=torch.float32)
backend = C4IndexerBackendMixin()
backend.dsa_topk_backend = DSATopKBackend.SGL_KERNEL
backend.token_to_kv_pool = SimpleNamespace(
layer_mapping={0: SimpleNamespace(compress_layer_id=7)}
)
backend.forward_metadata = SimpleNamespace(
indexer_metadata=indexer_metadata,
core_metadata=SimpleNamespace(
positions=torch.arange(1, dtype=torch.int64),
page_table=page_table,
c4_sparse_page_indices=page_indices,
c4_sparse_raw_indices=raw_indices,
),
)
backend.hisparse_coordinator = None
backend._forward_prepare_normal = MagicMock(
return_value=(
torch.empty((1, 1, 128)),
torch.empty((1, 1, 1)),
)
)
backend._get_nonpaged_indexer_plan = MagicMock(return_value=object())
backend._forward_nonpaged_indexer = MagicMock(return_value=logits)
indexer_capturer = MagicMock()
with (
envs.SGLANG_OPT_USE_TILELANG_INDEXER.override(False),
envs.SGLANG_OPT_USE_AITER_INDEXER.override(False),
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(True),
envs.SGLANG_OPT_USE_TOPK_V2.override(True),
patch(
f"{_INDEXER}.get_global_indexer_capturer",
return_value=indexer_capturer,
),
patch(f"{_INDEXER}.topk_transform_paged") as topk_v1,
patch(f"{_INDEXER}.topk_transform_paged_v2") as topk_v2,
):
backend.forward_c4_indexer(
x=torch.empty((1, 1)),
q_lora=torch.empty((1, 1)),
c4_indexer=SimpleNamespace(use_fp4_indexer=False, layer_id=0),
forward_batch=SimpleNamespace(forward_mode=ForwardMode.EXTEND),
)
topk_v2.assert_called_once()
args = topk_v2.call_args.args
self.assertIs(args[0], logits)
torch.testing.assert_close(args[1], c4_seq_lens)
torch.testing.assert_close(args[2], page_table)
torch.testing.assert_close(args[3], page_indices)
self.assertEqual(args[4], 64)
self.assertIs(args[5], topk_metadata)
self.assertEqual(args[6].data_ptr(), raw_indices.data_ptr())
topk_v1.assert_not_called()
indexer_capturer.capture.assert_called_once_with(7, raw_indices)
class TestDSV4NonPagedIndexer(CustomTestCase):
def _is_eligible(self, **overrides):
backend = SimpleNamespace(hisparse_coordinator=None)