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, row_starts=ks,
) )
elif self.topk_transform_method == TopkTransformMethod.RAGGED: 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( return fast_topk_transform_ragged_fused(
score=logits, score=logits,
lengths=seq_lens_topk, lengths=seq_lens_topk,
@@ -402,7 +407,9 @@ class NativeSparseAttnBackend(
# Centralized dispatch: decide all strategies for this batch # Centralized dispatch: decide all strategies for this batch
self.set_nsa_prefill_impl(forward_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, # Batch indices selected when cp enabled: After splitting multiple sequences,
# a certain cp rank may not have some of these 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, # 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]) topk_indices = self._pad_topk_indices(topk_indices, q_nope.shape[0])
# NOTE(dark): here, we use page size = 1 # 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(): if envs.SGLANG_NSA_FUSE_TOPK.get():
page_table_1 = topk_indices page_table_1 = topk_indices
else: else:
@@ -2095,7 +2104,9 @@ class NativeSparseAttnBackend(
# bf16 kv cache # bf16 kv cache
self.nsa_prefill_impl = "flashmla_sparse" 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. 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. 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" and self.nsa_prefill_impl == "flashmla_sparse"
): ):
topk_transform_method = TopkTransformMethod.RAGGED topk_transform_method = TopkTransformMethod.RAGGED
if forward_mode is not None and (forward_mode.is_decode_or_idle()):
topk_transform_method = TopkTransformMethod.PAGED
else: else:
topk_transform_method = TopkTransformMethod.PAGED topk_transform_method = TopkTransformMethod.PAGED
return topk_transform_method return topk_transform_method
@@ -2119,7 +2133,9 @@ class NativeSparseAttnBackend(
) )
return NSAIndexerMetadata( return NSAIndexerMetadata(
attn_metadata=self.forward_metadata, 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, paged_mqa_schedule_metadata=self.forward_metadata.paged_mqa_schedule_metadata,
force_unfused_topk=force_unfused, force_unfused_topk=force_unfused,
) )