From b1e1fe8eee688daa299a4f79cb8d427afa6146e4 Mon Sep 17 00:00:00 2001 From: cen121212 Date: Tue, 28 Apr 2026 09:08:28 +0800 Subject: [PATCH] =?UTF-8?q?=E3=80=90NPU=E3=80=91=E3=80=90bugfix=E3=80=91ac?= =?UTF-8?q?curacy=20fix=20when=20enable=20both=20nsa=20cp=20and=20prefixca?= =?UTF-8?q?che=20(#23268)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../modules/deepseek_v2_attention_mla_npu.py | 4 +++- .../srt/layers/attention/nsa/nsa_indexer.py | 22 +++++++++++++++---- 2 files changed, 21 insertions(+), 5 deletions(-) diff --git a/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py b/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py index 199d3f74f..6726f8589 100644 --- a/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py +++ b/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py @@ -359,7 +359,9 @@ def forward_dsa_prepare_npu( if q_event is not None: torch.npu.current_stream().wait_event(q_event) else: - if fused_qkv_a_proj_out.shape[0] < 65535: + if fused_qkv_a_proj_out.shape[0] < 65535 and not nsa_use_prefill_cp( + forward_batch + ): q_lora, k_nope, k_pe = fused_split_qk_norm( fused_qkv_a_proj_out, m.q_a_layernorm, diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index 3aff6f5d7..6e92533b7 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -1463,10 +1463,24 @@ class Indexer(MultiPlatformOp): forward_batch.attn_cp_metadata.actual_seq_q_prev_tensor, forward_batch.attn_cp_metadata.actual_seq_q_next_tensor, ) - forward_batch.attn_backend.forward_metadata.actual_seq_lengths_kv = ( - forward_batch.attn_cp_metadata.kv_len_prev_tensor, - forward_batch.attn_cp_metadata.kv_len_next_tensor, - ) + if sum(forward_batch.extend_prefix_lens_cpu) > 0: + total_kv_len_prev_tensor = ( + forward_batch.attn_cp_metadata.kv_len_prev_tensor + + forward_batch.extend_prefix_lens.squeeze() + ) + total_kv_len_next_tensor = ( + forward_batch.attn_cp_metadata.kv_len_next_tensor + + forward_batch.extend_prefix_lens.squeeze() + ) + forward_batch.attn_backend.forward_metadata.actual_seq_lengths_kv = ( + total_kv_len_prev_tensor, + total_kv_len_next_tensor, + ) + else: + forward_batch.attn_backend.forward_metadata.actual_seq_lengths_kv = ( + forward_batch.attn_cp_metadata.kv_len_prev_tensor, + forward_batch.attn_cp_metadata.kv_len_next_tensor, + ) actual_seq_lengths_q = ( forward_batch.attn_backend.forward_metadata.actual_seq_lengths_q )