From b1173c8c14690505c9a77f2591ccaaefa91400d4 Mon Sep 17 00:00:00 2001 From: McZyWu Date: Fri, 29 May 2026 11:29:30 +0800 Subject: [PATCH] [NPU] Enhance accuracy for model Step3_5 from 0 to 88% (#24582) --- .../npu/attention/ascend_backend.py | 162 +++++++++++++++--- .../sglang/srt/layers/quantization/unquant.py | 21 ++- 2 files changed, 155 insertions(+), 28 deletions(-) diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py index 038119688..77bde4232 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py @@ -45,6 +45,7 @@ def _reshape_kv_for_fia_nz( logger = logging.getLogger(__name__) +SWA_INT_MAX = 2147483647 @dataclass @@ -207,6 +208,16 @@ class AscendAttnMaskBuilder: ) return attn_mask + def get_swa_mask(self, seq_lens: torch.Tensor, s2: int, left_context=512): + if seq_lens.dim() == 1: + seq_lens = seq_lens.unsqueeze(1) + b = seq_lens.size(0) + device = seq_lens.device + indices = torch.arange(s2, device=device).unsqueeze(0).expand(b, -1) + start_indices = torch.clamp(seq_lens - left_context, min=0) + mask = (indices < start_indices) | (indices >= seq_lens) + return mask.unsqueeze(1).to(self.device, non_blocking=True) + def _cp_allgather_and_save_kv_npu( forward_batch, layer, k, v, cp_size, token_to_kv_pool @@ -1066,25 +1077,77 @@ class AscendAttnBackend(AttentionBackend): k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id) v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id) - if sinks is not None: + if sinks is not None or ( + self.is_hybrid_swa and layer.sliding_window_size != -1 + ): # Use SWA block tables if hybrid SWA is enabled for this layer if self.is_hybrid_swa and layer.sliding_window_size != -1: block_tables = self.forward_metadata.block_tables_swa else: block_tables = self.forward_metadata.block_tables - attn_out = attention_sinks_prefill_triton( - q, - k_cache, - v_cache, - sinks, - self.forward_metadata.extend_seq_lens, - block_tables, - self.forward_metadata.seq_lens, - layer.scaling, - layer.sliding_window_size, - layer.tp_q_head_num, - layer.tp_k_head_num, - ) + if self.use_fia: + num_token_padding = q.shape[0] + if num_token_padding > forward_batch.num_token_non_padded_cpu: + q, k, v = [ + data[: forward_batch.num_token_non_padded_cpu] + for data in [q, k, v] + ] + q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim) + block_size = self.page_size + attn_out, _ = torch_npu.npu_fused_infer_attention_score_v2( + query=q, + key=k_cache.view( + -1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim + ), + value=v_cache.view( + -1, self.page_size, layer.tp_v_head_num * layer.v_head_dim + ), + pre_tokens=( + layer.sliding_window_size + if layer.sliding_window_size != -1 + else SWA_INT_MAX + ), + next_tokens=( + 0 if layer.sliding_window_size != -1 else SWA_INT_MAX + ), + atten_mask=self.fia_mask, + block_table=block_tables, + input_layout="TND", + block_size=block_size, + num_query_heads=layer.tp_q_head_num, + num_key_value_heads=layer.tp_k_head_num, + actual_seq_qlen=self.forward_metadata.seq_lens_list_cumsum, + actual_seq_kvlen=self.forward_metadata.seq_lens_cpu_int, + softmax_scale=layer.scaling, + sparse_mode=4 if layer.sliding_window_size != -1 else 3, + learnable_sink=sinks, + ) + if num_token_padding != forward_batch.num_token_non_padded_cpu: + attn_out = torch.cat( + [ + attn_out, + attn_out.new_zeros( + num_token_padding - attn_out.shape[0], + *attn_out.shape[1:], + ), + ], + dim=0, + ) + attn_out = attn_out.view(-1, layer.tp_q_head_num * layer.v_head_dim) + else: + attn_out = attention_sinks_prefill_triton( + q, + k_cache, + v_cache, + sinks, + self.forward_metadata.extend_seq_lens, + block_tables, + self.forward_metadata.seq_lens, + layer.scaling, + layer.sliding_window_size, + layer.tp_q_head_num, + layer.tp_k_head_num, + ) return attn_out if is_cp_mode: @@ -1986,24 +2049,69 @@ class AscendAttnBackend(AttentionBackend): k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id) v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id) - if sinks is not None: + if sinks is not None or ( + self.is_hybrid_swa and layer.sliding_window_size != -1 + ): # Use SWA block tables if hybrid SWA is enabled for this layer if self.is_hybrid_swa and layer.sliding_window_size != -1: block_tables = self.forward_metadata.block_tables_swa else: block_tables = self.forward_metadata.block_tables - attn_out = attention_sinks_triton( - q, - k_cache, - v_cache, - sinks, - block_tables, - self.forward_metadata.seq_lens, - layer.scaling, - layer.sliding_window_size, - layer.tp_q_head_num, - layer.tp_k_head_num, - ) + if self.use_fia: + if self.forward_metadata.seq_lens_cpu_int is None: + actual_seq_len_kv = self.forward_metadata.seq_lens_cpu_list + else: + actual_seq_len_kv = ( + self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist() + ) + block_size = self.page_size + max_model_len = block_tables.shape[-1] * block_size + swa_mask = self.ascend_attn_mask_builder.get_swa_mask( + self.forward_metadata.seq_lens, + max_model_len, + layer.sliding_window_size, + ) + attn_out, _ = torch_npu.npu_fused_infer_attention_score_v2( + q.view( + forward_batch.batch_size, + -1, + layer.tp_q_head_num, + layer.qk_head_dim, + ), + k_cache.view( + -1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim + ), + v_cache.view( + -1, self.page_size, layer.tp_v_head_num * layer.v_head_dim + ), + num_query_heads=layer.tp_q_head_num, + num_key_value_heads=layer.tp_k_head_num, + input_layout="BSND", + block_size=block_size, + atten_mask=( + swa_mask if layer.sliding_window_size != -1 else None + ), + sparse_mode=4 if layer.sliding_window_size != -1 else 0, + softmax_scale=layer.scaling, + block_table=block_tables, + actual_seq_qlen=[1] * len(self.forward_metadata.seq_lens), + actual_seq_kvlen=actual_seq_len_kv, + learnable_sink=sinks, + ) + attn_out = attn_out.view(-1, layer.tp_q_head_num * layer.v_head_dim) + else: + attn_out = attention_sinks_triton( + q, + k_cache, + v_cache, + sinks, + block_tables, + self.forward_metadata.seq_lens, + layer.scaling, + layer.sliding_window_size, + layer.tp_q_head_num, + layer.tp_k_head_num, + ) return attn_out if self.use_fia: diff --git a/python/sglang/srt/layers/quantization/unquant.py b/python/sglang/srt/layers/quantization/unquant.py index 99d0a2468..1a4ab64ae 100644 --- a/python/sglang/srt/layers/quantization/unquant.py +++ b/python/sglang/srt/layers/quantization/unquant.py @@ -69,6 +69,20 @@ except ImportError: flashinfer_cutlass_fused_moe = None +def swiglustep_and_mul(x: torch.Tensor, limit: float = 7.0) -> torch.Tensor: + """Out-variant of swiglustep activation. + + Writes into `out`: + silu(x[:d]).clamp(max=limit) * x[d:].clamp(-limit, limit) + """ + gate, up = x.chunk(2, dim=-1) + gate = F.silu(gate) + gate = gate.clamp(max=limit) + up = up.clamp(min=-limit, max=limit) + out = gate * up + return out + + class UnquantizedEmbeddingMethod(QuantizeMethodBase): """Unquantized method for embeddings.""" @@ -693,7 +707,12 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp): hidden_states = swiglu_oai(layer, hidden_states) elif self.moe_runner_config.activation == "silu": - hidden_states = torch.ops.npu.npu_swiglu(hidden_states) + if self.moe_runner_config.gemm1_clamp_limit is not None: + hidden_states = swiglustep_and_mul( + hidden_states, self.moe_runner_config.gemm1_clamp_limit + ) + else: + hidden_states = torch.ops.npu.npu_swiglu(hidden_states) else: from sglang.srt.layers.activation import GeluAndMul