FIX: (NSA) Compute topk_indices_offset when NSA prefill flashmla_sparse is used with FP8 KV cache (#20606)

Signed-off-by: Ho-Ren (Jack) Chuang <horenchuang@bytedance.com>
This commit is contained in:
Ho-Ren (Jack) Chuang
2026-03-26 12:50:50 -07:00
committed by GitHub
parent 3867c6431a
commit 4b5f63e1b8
@@ -260,6 +260,11 @@ class NSAIndexerMetadata(BaseIndexerMetadata):
row_starts=ks,
)
elif self.topk_transform_method == TopkTransformMethod.RAGGED:
if cu_topk_indices_offset is None:
raise RuntimeError(
"RAGGED topk_transform requires topk_indices_offset; "
"expected extend-without-speculative metadata."
)
return fast_topk_transform_ragged_fused(
score=logits,
lengths=seq_lens_topk,
@@ -402,7 +407,9 @@ class NativeSparseAttnBackend(
# Centralized dispatch: decide all strategies for this batch
self.set_nsa_prefill_impl(forward_batch)
topk_transform_method = self.get_topk_transform_method()
topk_transform_method = self.get_topk_transform_method(
forward_batch.forward_mode
)
# Batch indices selected when cp enabled: After splitting multiple sequences,
# a certain cp rank may not have some of these sequences.
# We use bs_idx_cpu to mark which sequences are finally selected by the current cp rank,
@@ -1343,7 +1350,9 @@ class NativeSparseAttnBackend(
topk_indices = self._pad_topk_indices(topk_indices, q_nope.shape[0])
# NOTE(dark): here, we use page size = 1
topk_transform_method = self.get_topk_transform_method()
topk_transform_method = self.get_topk_transform_method(
forward_batch.forward_mode
)
if envs.SGLANG_NSA_FUSE_TOPK.get():
page_table_1 = topk_indices
else:
@@ -2095,7 +2104,9 @@ class NativeSparseAttnBackend(
# bf16 kv cache
self.nsa_prefill_impl = "flashmla_sparse"
def get_topk_transform_method(self) -> TopkTransformMethod:
def get_topk_transform_method(
self, forward_mode: Optional[ForwardMode] = None
) -> TopkTransformMethod:
"""
SGLANG_NSA_FUSE_TOPK controls whether to fuse the topk transform into the topk kernel.
This method is used to select the topk transform method which can be fused or unfused.
@@ -2106,6 +2117,9 @@ class NativeSparseAttnBackend(
and self.nsa_prefill_impl == "flashmla_sparse"
):
topk_transform_method = TopkTransformMethod.RAGGED
if forward_mode is not None and (forward_mode.is_decode_or_idle()):
topk_transform_method = TopkTransformMethod.PAGED
else:
topk_transform_method = TopkTransformMethod.PAGED
return topk_transform_method
@@ -2119,7 +2133,9 @@ class NativeSparseAttnBackend(
)
return NSAIndexerMetadata(
attn_metadata=self.forward_metadata,
topk_transform_method=self.get_topk_transform_method(),
topk_transform_method=self.get_topk_transform_method(
forward_batch.forward_mode
),
paged_mqa_schedule_metadata=self.forward_metadata.paged_mqa_schedule_metadata,
force_unfused_topk=force_unfused,
)