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 b16fa2b01..d5938c36e 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py @@ -2104,7 +2104,7 @@ class AscendAttnBackend(AttentionBackend): v, ) - if sinks is not None: + if sinks is not None or self.is_hybrid_swa: # Use SWA block tables if hybrid SWA is enabled for this layer if self._is_swa_layer(layer): block_tables = self.forward_metadata.block_tables_swa @@ -2158,7 +2158,9 @@ class AscendAttnBackend(AttentionBackend): if layer.sliding_window_size != -1 else FULL_ATTENTION_WINDOW ), - next_tokens=0, + next_tokens=( + 0 if layer.sliding_window_size == -1 else FULL_ATTENTION_WINDOW + ), atten_mask=self.fia_mask.to(torch.int8), sparse_mode=sparse_mode, softmax_scale=layer.scaling, @@ -2213,59 +2215,6 @@ class AscendAttnBackend(AttentionBackend): else: actual_seq_len_kv = seq_lens_cpu_int.cpu().int().tolist() - if (layer.qk_head_dim != layer.v_head_dim) and ( - self.is_hybrid_swa and layer.sliding_window_size == -1 - ): - query_v2 = q.reshape( - -1, layer.tp_q_head_num, layer.qk_head_dim - ).contiguous() - actual_seq_qlen = ( - torch.tensor([1] * len(actual_seq_len_kv), dtype=torch.int32) - .cumsum(dim=0) - .tolist() - ) - common_kwargs = dict( - num_query_heads=layer.tp_q_head_num, - num_key_value_heads=layer.tp_k_head_num, - input_layout="TND", - pre_tokens=FULL_ATTENTION_WINDOW, - next_tokens=0, - atten_mask=self.fia_mask.to(torch.int8), - sparse_mode=3, - softmax_scale=layer.scaling, - block_table=self.forward_metadata.block_tables, - block_size=self.page_size, - actual_seq_qlen=actual_seq_qlen, - actual_seq_kvlen=actual_seq_len_kv, - ) - workspace = ( - torch_npu._npu_fused_infer_attention_score_v2_get_max_workspace( - query_v2, - k_cache.contiguous(), - v_cache.contiguous(), - **common_kwargs, - ) - ) - attn_output = torch.empty( - ( - query_v2.shape[0], - layer.tp_q_head_num, - layer.v_head_dim, - ), - dtype=q.dtype, - device=q.device, - ) - softmax_lse = torch.empty(1, dtype=q.dtype, device=q.device) - torch_npu.npu_fused_infer_attention_score_v2.out( - query_v2, - k_cache.contiguous(), - v_cache.contiguous(), - **common_kwargs, - workspace=workspace, - out=[attn_output, softmax_lse], - ) - return attn_output.view(-1, layer.tp_q_head_num * layer.v_head_dim) - num_tokens = query.shape[0] workspace = torch_npu._npu_fused_infer_attention_score_get_max_workspace( query, diff --git a/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py b/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py index 2869c4bd1..0edb2afde 100644 --- a/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py +++ b/python/sglang/srt/hardware_backend/npu/graph_runner/npu_graph_runner.py @@ -108,7 +108,8 @@ class NPUGraphRunner(DecodeCudaGraphRunner): self._init_arch_map() self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False") self.if_use_v2 = any( - arch in ("MiMoV2ForCausalLM", "MiMoV2FlashForCausalLM") + arch + in ("MiMoV2ForCausalLM", "MiMoV2FlashForCausalLM", "Step3p5ForCausalLM") for arch in (model_runner.model_config.hf_config.architectures or []) ) diff --git a/python/sglang/srt/layers/quantization/unquant.py b/python/sglang/srt/layers/quantization/unquant.py index f36e9b6d3..a30872231 100644 --- a/python/sglang/srt/layers/quantization/unquant.py +++ b/python/sglang/srt/layers/quantization/unquant.py @@ -69,20 +69,6 @@ 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.""" @@ -729,8 +715,15 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp): hidden_states = swiglu_oai(layer, hidden_states) elif self.moe_runner_config.activation == "silu": 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 + from sgl_kernel_npu.activation.swiglu_quant import swiglu_quant + + hidden_states, _ = swiglu_quant( + hidden_states, + group_list=expert_tokens, + group_list_type=1, + need_quant=False, + do_limit=True, + limit=self.moe_runner_config.gemm1_clamp_limit, ) else: hidden_states = torch.ops.npu.npu_swiglu(hidden_states)