fix(dsa): use FlashInfer fused top-k for packed PAGED rows (#33006)
This commit is contained in:
@@ -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(
|
||||||
|
|||||||
Reference in New Issue
Block a user