fix(dsa): use FlashInfer fused top-k for packed PAGED rows (#33006)

This commit is contained in:
Ziang Li
2026-08-14 15:22:18 -07:00
committed by GitHub
parent bfb224ff01
commit a9654eacc1
2 changed files with 85 additions and 46 deletions
@@ -158,27 +158,7 @@ class DSATopKBackend(Enum):
import flashinfer import flashinfer
if topk_transform_method == TopkTransformMethod.PAGED: if topk_transform_method == TopkTransformMethod.PAGED:
if row_starts is not None: row_to_batch, page_table_row_starts = _build_flashinfer_paged_args(
# Packed PAGED extend uses batch-global score offsets with
# request-local page tables. FlashInfer applies row_starts
# to both, so reuse the SGL transform.
from sgl_kernel import fast_topk_transform_fused
page_table_size_1 = (
attn_metadata.page_table_1[batch_idx_list]
if batch_idx_list is not None
else attn_metadata.page_table_1
)
return fast_topk_transform_fused(
score=logits,
lengths=lengths,
page_table_size_1=page_table_size_1,
cu_seqlens_q=cu_seqlens_q_topk,
topk=topk,
row_starts=row_starts,
)
row_to_batch, local_row_starts = _build_flashinfer_paged_args(
attn_metadata=attn_metadata, attn_metadata=attn_metadata,
row_starts=row_starts, row_starts=row_starts,
cu_seqlens_q_topk=cu_seqlens_q_topk, cu_seqlens_q_topk=cu_seqlens_q_topk,
@@ -195,7 +175,8 @@ class DSATopKBackend(Enum):
deterministic=envs.SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC.get(), deterministic=envs.SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC.get(),
tie_break=_flashinfer_tie_break_value(), tie_break=_flashinfer_tie_break_value(),
dsa_graph_safe=True, dsa_graph_safe=True,
row_starts=local_row_starts, row_starts=row_starts,
page_table_row_starts=page_table_row_starts,
) )
if topk_transform_method == TopkTransformMethod.RAGGED: if topk_transform_method == TopkTransformMethod.RAGGED:
if topk_indices_offset is None: if topk_indices_offset is None:
@@ -373,13 +354,13 @@ def _build_flashinfer_paged_args(
"PAGED topk_transform with row_starts requires cu_seqlens_q metadata." "PAGED topk_transform with row_starts requires cu_seqlens_q metadata."
) )
local_row_starts = row_starts page_table_row_starts = row_starts
if local_row_starts is not None and row_to_batch is not None: if page_table_row_starts is not None and row_to_batch is not None:
local_row_starts = ( page_table_row_starts = (
local_row_starts - attn_metadata.cu_seqlens_k[:-1][row_to_batch] page_table_row_starts - attn_metadata.cu_seqlens_k[:-1][row_to_batch]
) )
return row_to_batch, local_row_starts return row_to_batch, page_table_row_starts
def _flashinfer_tie_break_value() -> int: def _flashinfer_tie_break_value() -> int:
@@ -570,25 +570,55 @@ class TestDSAIndexer(CustomTestCase):
query_lens: Optional[List[int]] = None, query_lens: Optional[List[int]] = None,
): ):
num_rows = sum(query_lens) if query_lens is not None else batch_size num_rows = sum(query_lens) if query_lens is not None else batch_size
# Shifted PAGED rows use global score offsets and request-local page tables.
# Give each packed row more than topk entries to exercise actual selection.
if with_row_starts and topk_transform_method == TopkTransformMethod.PAGED:
max_score_len = max(max_score_len, batch_size * (topk + 1))
logits = self._make_tie_free_logits(num_rows, max_score_len) logits = self._make_tie_free_logits(num_rows, max_score_len)
if with_row_starts: if with_row_starts:
row_starts = torch.randint( if topk_transform_method == TopkTransformMethod.PAGED:
0, packed_row_size = max_score_len // batch_size
max_score_len - 1, self.assertGreaterEqual(packed_row_size, topk)
(num_rows,), cu_seqlens_k = (
dtype=torch.int32, torch.arange(batch_size + 1, dtype=torch.int32, device=self.device)
device=self.device, * packed_row_size
) )
max_lengths = max_score_len - row_starts if query_lens is None:
random_lengths = torch.randint( row_to_batch = torch.arange(
1, batch_size, dtype=torch.int32, device=self.device
max_score_len, )
(num_rows,), else:
dtype=torch.int32, row_to_batch = torch.repeat_interleave(
device=self.device, torch.arange(batch_size, dtype=torch.int32, device=self.device),
) torch.tensor(query_lens, dtype=torch.int32, device=self.device),
seq_lens_expanded = torch.minimum(max_lengths, random_lengths) output_size=num_rows,
)
row_starts = cu_seqlens_k[:-1][row_to_batch]
seq_lens_expanded = torch.randint(
topk + 1,
packed_row_size + 1,
(num_rows,),
dtype=torch.int32,
device=self.device,
)
else:
row_starts = torch.randint(
0,
max_score_len - 1,
(num_rows,),
dtype=torch.int32,
device=self.device,
)
max_lengths = max_score_len - row_starts
random_lengths = torch.randint(
1,
max_score_len,
(num_rows,),
dtype=torch.int32,
device=self.device,
)
seq_lens_expanded = torch.minimum(max_lengths, random_lengths)
else: else:
row_starts = None row_starts = None
seq_lens_expanded = torch.randint( seq_lens_expanded = torch.randint(
@@ -616,9 +646,10 @@ class TestDSAIndexer(CustomTestCase):
) )
cu_seqlens_q[1:] = torch.cumsum(q_lens, dim=0) cu_seqlens_q[1:] = torch.cumsum(q_lens, dim=0)
batch_idx_list = list(range(batch_size)) batch_idx_list = list(range(batch_size))
cu_seqlens_k = torch.zeros( if not (with_row_starts and topk_transform_method == TopkTransformMethod.PAGED):
batch_size + 1, dtype=torch.int32, device=self.device cu_seqlens_k = torch.zeros(
) batch_size + 1, dtype=torch.int32, device=self.device
)
dsa_cu_seqlens_k = torch.zeros( dsa_cu_seqlens_k = torch.zeros(
num_rows + 1, dtype=torch.int32, device=self.device num_rows + 1, dtype=torch.int32, device=self.device
) )
@@ -677,13 +708,21 @@ class TestDSAIndexer(CustomTestCase):
topk_backend=DSATopKBackend.FLASHINFER, topk_backend=DSATopKBackend.FLASHINFER,
) )
import flashinfer
repeat_interleave = torch.repeat_interleave repeat_interleave = torch.repeat_interleave
top_k_page_table_transform = flashinfer.top_k_page_table_transform
with ( with (
envs.SGLANG_DSA_FUSE_TOPK.override(True), envs.SGLANG_DSA_FUSE_TOPK.override(True),
patch( patch(
"sglang.srt.layers.attention.dsa_backend.torch.repeat_interleave", "sglang.srt.layers.attention.dsa_backend.torch.repeat_interleave",
wraps=repeat_interleave, wraps=repeat_interleave,
) as mock_repeat_interleave, ) as mock_repeat_interleave,
patch.object(
flashinfer,
"top_k_page_table_transform",
wraps=top_k_page_table_transform,
) as mock_top_k_page_table_transform,
): ):
out_sgl = metadata_sgl.topk_transform( out_sgl = metadata_sgl.topk_transform(
logits, logits,
@@ -700,6 +739,25 @@ class TestDSAIndexer(CustomTestCase):
batch_idx_list=batch_idx_list, batch_idx_list=batch_idx_list,
) )
if topk_transform_method == TopkTransformMethod.PAGED:
mock_top_k_page_table_transform.assert_called_once()
call_kwargs = mock_top_k_page_table_transform.call_args.kwargs
if row_starts is None:
self.assertIsNone(call_kwargs["row_starts"])
self.assertIsNone(call_kwargs["page_table_row_starts"])
else:
self.assertTrue(torch.equal(call_kwargs["row_starts"], row_starts))
self.assertIsNotNone(call_kwargs["row_to_batch"])
expected_page_table_row_starts = (
row_starts - cu_seqlens_k[:-1][call_kwargs["row_to_batch"]]
)
self.assertTrue(
torch.equal(
call_kwargs["page_table_row_starts"],
expected_page_table_row_starts,
)
)
if query_lens is not None: if query_lens is not None:
self.assertTrue(mock_repeat_interleave.call_args_list) self.assertTrue(mock_repeat_interleave.call_args_list)
self.assertTrue( self.assertTrue(