Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions python/sglang/srt/layers/attention/dsv4/indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -585,6 +585,13 @@ def to_cpu_int_list(values) -> Optional[List[int]]:
ke = c4_seq_lens[:query_rows].reshape(-1).to(torch.int32).contiguous()
gather_seq_lens = ke[-1:]
ks = torch.zeros_like(ke)
# SGL Top-K synthesizes sequential indices for trivial rows without
# reading logits, so DeepGEMM can receive an empty range for them.
if (
self.dsa_topk_backend.is_sgl_kernel()
and not envs.SGLANG_TOPK_TRANSFORM_512_TORCH.get()
):
ke = torch.where(ke - ks > c4_indexer.index_topk, ke, ks)
c4_page_size = indexer_metadata.c4_page_size
max_seqlen_k = (final_c4_len + c4_page_size - 1) // c4_page_size * c4_page_size
plan = NonPagedIndexerPlan(
Expand Down
18 changes: 13 additions & 5 deletions test/registered/unit/layers/test_dsv4_nonpaged_indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -128,7 +128,8 @@ def test_eligibility_is_fail_closed(self):

def test_single_request_plan_contract(self):
backend = SimpleNamespace(_can_use_nonpaged_indexer=lambda **_: True)
c4_indexer = SimpleNamespace(use_fp4_indexer=False)
backend.dsa_topk_backend = SimpleNamespace(is_sgl_kernel=lambda: True)
c4_indexer = SimpleNamespace(use_fp4_indexer=False, index_topk=64)
query_rows = 4
batch = SimpleNamespace(
seq_lens=torch.tensor([262], dtype=torch.int32),
Expand Down Expand Up @@ -156,14 +157,19 @@ def build_plan():
threshold = envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER_MIN_QUERY_TOKENS
with threshold.override(threshold.default):
self.assertIsNone(build_plan())
with threshold.override(query_rows):
with (
threshold.override(query_rows),
envs.SGLANG_TOPK_TRANSFORM_512_TORCH.override(False),
):
plan = build_plan()
self.assertEqual(
(plan.seq_len_sum, plan.max_seqlen_k, plan.query_rows),
(65, 128, query_rows),
)
torch.testing.assert_close(plan.page_table, page_table[:1])
torch.testing.assert_close(plan.ke, c4_seq_lens)
torch.testing.assert_close(
plan.ke, torch.tensor([0, 0, 0, 65], dtype=torch.int32)
)
torch.testing.assert_close(plan.gather_seq_lens, c4_seq_lens[-1:])

metadata.nonpaged_plan = None
Expand All @@ -173,7 +179,8 @@ def build_plan():

def test_extreme_plan_metadata_is_bounded_and_fail_closed(self):
backend = SimpleNamespace(_can_use_nonpaged_indexer=lambda **_: True)
c4_indexer = SimpleNamespace(use_fp4_indexer=False)
backend.dsa_topk_backend = SimpleNamespace(is_sgl_kernel=lambda: True)
c4_indexer = SimpleNamespace(use_fp4_indexer=False, index_topk=512)
query_rows = 4
batch = SimpleNamespace(
seq_lens=torch.tensor([500_000], dtype=torch.int32),
Expand Down Expand Up @@ -219,7 +226,8 @@ def build_plan():
def test_query_threshold_boundary(self):
can_use_nonpaged_indexer = MagicMock(return_value=True)
backend = SimpleNamespace(_can_use_nonpaged_indexer=can_use_nonpaged_indexer)
c4_indexer = SimpleNamespace(use_fp4_indexer=False)
backend.dsa_topk_backend = SimpleNamespace(is_sgl_kernel=lambda: True)
c4_indexer = SimpleNamespace(use_fp4_indexer=False, index_topk=512)
metadata = SimpleNamespace(nonpaged_plan=None, c4_page_size=64)

def build_plan(query_rows):
Expand Down
Loading