[NSA] Fix NSA backend assertion error when running DeepSeek-V3.2 PP with radix-cache (#15086)

This commit is contained in:
YAMY
2025-12-14 17:13:18 -08:00
committed by GitHub
parent 4449c17011
commit c96903074c
2 changed files with 19 additions and 5 deletions

View File

@@ -189,10 +189,9 @@ def dequantize_k_cache_paged(
), f"dim_quant: {dim_quant} != 656 detected in dequantize_k_cache_paged"
quant_k_cache = quant_k_cache.view((-1, dim_quant))
total_num_tokens, _ = quant_k_cache.shape
# num_tokens can exceed kv_cache_size due to prefix sharing (multiple seqs share same KV slots)
# Index bounds validated in nsa_backend.init_forward_metadata
num_tokens = page_table_1_flattened.shape[0]
assert num_tokens <= total_num_tokens
assert quant_k_cache.dtype == torch.float8_e4m3fn
dim_nope = 512
dim_rope = 64

View File

@@ -463,14 +463,16 @@ class NativeSparseAttnBackend(AttentionBackend):
]
)
# Check if MHA with FP8 needs page_table_1_flattened for dequantization
# Check if MHA FP8 dequantization is needed
mha_dequantize_needed = (
self.use_mha
and forward_batch.token_to_kv_pool.dtype == torch.float8_e4m3fn
)
forward_batch.using_mha_one_shot_fp8_dequant = mha_dequantize_needed
if (
# page_table_1_flattened is only used when prefix sharing is enabled:
has_prefix_sharing = any(forward_batch.extend_prefix_lens_cpu)
if has_prefix_sharing and (
topk_transform_method == TopkTransformMethod.RAGGED
or mha_dequantize_needed
):
@@ -486,6 +488,19 @@ class NativeSparseAttnBackend(AttentionBackend):
page_table_1_flattened.shape[0] == forward_batch.seq_lens_sum
), f"{page_table_1_flattened.shape[0] = } must be the same as {forward_batch.seq_lens_sum = }"
# Validate indices when logical tokens exceed physical capacity
# This is likely to be triggered by PP with high kv reuse & parallelism
kv_cache_capacity = (
forward_batch.token_to_kv_pool.size
+ forward_batch.token_to_kv_pool.page_size
)
if forward_batch.seq_lens_sum > kv_cache_capacity:
max_idx = page_table_1_flattened.max().item()
assert max_idx < kv_cache_capacity, (
f"Invalid page table index: max={max_idx}, "
f"kv_cache_capacity={kv_cache_capacity}"
)
if topk_transform_method == TopkTransformMethod.RAGGED:
topk_indices_offset = torch.repeat_interleave(
cu_seqlens_k[:-1],