[NSA] Fix NSA backend assertion error when running DeepSeek-V3.2 PP with radix-cache (#15086)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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],
|
||||
|
||||
Reference in New Issue
Block a user