From a9654eacc12cb27abfb39f586d386b11ccd102bf Mon Sep 17 00:00:00 2001 From: Ziang Li Date: Fri, 14 Aug 2026 15:22:18 -0700 Subject: [PATCH] fix(dsa): use FlashInfer fused top-k for packed PAGED rows (#33006) --- .../layers/attention/dsa/dsa_topk_backend.py | 35 ++----- .../kernels/ops/attention/test_dsa_indexer.py | 96 +++++++++++++++---- 2 files changed, 85 insertions(+), 46 deletions(-) diff --git a/python/sglang/srt/layers/attention/dsa/dsa_topk_backend.py b/python/sglang/srt/layers/attention/dsa/dsa_topk_backend.py index c80315fd8..a8f0d727d 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_topk_backend.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_topk_backend.py @@ -158,27 +158,7 @@ class DSATopKBackend(Enum): import flashinfer if topk_transform_method == TopkTransformMethod.PAGED: - if row_starts is not None: - # 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( + row_to_batch, page_table_row_starts = _build_flashinfer_paged_args( attn_metadata=attn_metadata, row_starts=row_starts, cu_seqlens_q_topk=cu_seqlens_q_topk, @@ -195,7 +175,8 @@ class DSATopKBackend(Enum): deterministic=envs.SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC.get(), tie_break=_flashinfer_tie_break_value(), 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_indices_offset is None: @@ -373,13 +354,13 @@ def _build_flashinfer_paged_args( "PAGED topk_transform with row_starts requires cu_seqlens_q metadata." ) - local_row_starts = row_starts - if local_row_starts is not None and row_to_batch is not None: - local_row_starts = ( - local_row_starts - attn_metadata.cu_seqlens_k[:-1][row_to_batch] + page_table_row_starts = row_starts + if page_table_row_starts is not None and row_to_batch is not None: + page_table_row_starts = ( + 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: diff --git a/test/registered/kernels/ops/attention/test_dsa_indexer.py b/test/registered/kernels/ops/attention/test_dsa_indexer.py index 93d53b382..00fc9864d 100644 --- a/test/registered/kernels/ops/attention/test_dsa_indexer.py +++ b/test/registered/kernels/ops/attention/test_dsa_indexer.py @@ -570,25 +570,55 @@ class TestDSAIndexer(CustomTestCase): query_lens: Optional[List[int]] = None, ): 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) if with_row_starts: - 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) + if topk_transform_method == TopkTransformMethod.PAGED: + packed_row_size = max_score_len // batch_size + self.assertGreaterEqual(packed_row_size, topk) + cu_seqlens_k = ( + torch.arange(batch_size + 1, dtype=torch.int32, device=self.device) + * packed_row_size + ) + if query_lens is None: + row_to_batch = torch.arange( + batch_size, dtype=torch.int32, device=self.device + ) + else: + row_to_batch = torch.repeat_interleave( + torch.arange(batch_size, dtype=torch.int32, device=self.device), + torch.tensor(query_lens, dtype=torch.int32, device=self.device), + 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: row_starts = None seq_lens_expanded = torch.randint( @@ -616,9 +646,10 @@ class TestDSAIndexer(CustomTestCase): ) cu_seqlens_q[1:] = torch.cumsum(q_lens, dim=0) batch_idx_list = list(range(batch_size)) - cu_seqlens_k = torch.zeros( - batch_size + 1, dtype=torch.int32, device=self.device - ) + if not (with_row_starts and topk_transform_method == TopkTransformMethod.PAGED): + cu_seqlens_k = torch.zeros( + batch_size + 1, dtype=torch.int32, device=self.device + ) dsa_cu_seqlens_k = torch.zeros( num_rows + 1, dtype=torch.int32, device=self.device ) @@ -677,13 +708,21 @@ class TestDSAIndexer(CustomTestCase): topk_backend=DSATopKBackend.FLASHINFER, ) + import flashinfer + repeat_interleave = torch.repeat_interleave + top_k_page_table_transform = flashinfer.top_k_page_table_transform with ( envs.SGLANG_DSA_FUSE_TOPK.override(True), patch( "sglang.srt.layers.attention.dsa_backend.torch.repeat_interleave", wraps=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( logits, @@ -700,6 +739,25 @@ class TestDSAIndexer(CustomTestCase): 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: self.assertTrue(mock_repeat_interleave.call_args_list) self.assertTrue(