From 3553fd0322518d424d528cae4512168b4e07d16e Mon Sep 17 00:00:00 2001 From: heziiop Date: Thu, 30 Apr 2026 08:51:22 +0800 Subject: [PATCH] [NPU] add split_qkv_tp_rmsnorm_rope ops for minimax2 & fix eagle3 hidden states capture in dp attn mode (#23190) --- python/sglang/srt/models/minimax_m2.py | 76 ++++++++++++++++++++++---- 1 file changed, 66 insertions(+), 10 deletions(-) diff --git a/python/sglang/srt/models/minimax_m2.py b/python/sglang/srt/models/minimax_m2.py index 5f7cab05e..14afc0d2f 100644 --- a/python/sglang/srt/models/minimax_m2.py +++ b/python/sglang/srt/models/minimax_m2.py @@ -18,7 +18,7 @@ import logging from contextlib import nullcontext from functools import lru_cache -from typing import Any, Dict, Iterable, Optional, Set, Tuple, Union +from typing import Any, Dict, Iterable, List, Optional, Set, Tuple, Union import torch import triton @@ -93,6 +93,7 @@ from sglang.srt.utils import ( get_compiler_backend, is_cuda, is_non_idle_and_non_empty, + is_npu, make_layers, ) from sglang.srt.utils.custom_op import register_custom_op @@ -100,6 +101,10 @@ from sglang.srt.utils.hf_transformers_utils import get_rope_config logger = logging.getLogger(__name__) _is_cuda = is_cuda() +_is_npu = is_npu() + +if _is_npu: + from sgl_kernel_npu.norm.split_qkv_tp_rmsnorm_rope import split_qkv_tp_rmsnorm_rope @triton.jit @@ -816,6 +821,43 @@ class MiniMaxM2Attention(nn.Module): inner_state = q, k, v, forward_batch return None, forward_batch, inner_state + def forward_prepare_npu( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + ): + if hidden_states.shape[0] == 0: + assert ( + not self.o_proj.reduce_results + ), "short-circuiting allreduce will lead to hangs" + return hidden_states, forward_batch, None + qkv, _ = self.qkv_proj(hidden_states) + if self.use_qk_norm: + cos_sin = self.rotary_emb.cos_sin_cache.index_select(0, positions.flatten()) + cos, sin = cos_sin.chunk(2, dim=-1) + q, k, v = split_qkv_tp_rmsnorm_rope( + input=qkv, + cos=cos, + sin=sin, + q_weight=self.q_norm.weight, + k_weight=self.k_norm.weight, + q_hidden_size=self.q_size, + kv_hidden_size=self.kv_size, + head_dim=self.head_dim, + rotary_dim=self.rotary_dim, + eps=self.q_norm.variance_epsilon, + tp_world=self.q_norm.attn_tp_size, + tp_group=get_attention_tp_group().device_group, + ) + else: + q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) + q, k = q.contiguous(), k.contiguous() + q, k = self.rotary_emb(positions, q, k) + + inner_state = q, k, v, forward_batch + return None, forward_batch, inner_state + def forward_core(self, intermediate_state): hidden_states, forward_batch, inner_state = intermediate_state if inner_state is None: @@ -830,11 +872,18 @@ class MiniMaxM2Attention(nn.Module): hidden_states: torch.Tensor, forward_batch: ForwardBatch, ) -> torch.Tensor: - s = self.forward_prepare( - positions=positions, - hidden_states=hidden_states, - forward_batch=forward_batch, - ) + if not _is_npu: + s = self.forward_prepare( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, + ) + else: + s = self.forward_prepare_npu( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, + ) return self.forward_core(s) def op_prepare(self, state): @@ -912,10 +961,16 @@ class MiniMaxM2DecoderLayer(nn.Module): hidden_states: torch.Tensor, forward_batch: ForwardBatch, residual: Optional[torch.Tensor], + captured_last_layer_outputs: Optional[List[torch.Tensor]] = None, ) -> torch.Tensor: # Self Attention - hidden_states, residual = self.layer_communicator.prepare_attn( - hidden_states, residual, forward_batch + hidden_states, residual = ( + self.layer_communicator.prepare_attn_and_capture_last_layer_outputs( + hidden_states, + residual, + forward_batch, + captured_last_layer_outputs=captured_last_layer_outputs, + ) ) if not forward_batch.forward_mode.is_idle(): hidden_states = self.self_attn( @@ -1099,14 +1154,15 @@ class MiniMaxM2Model(nn.Module): else get_global_expert_distribution_recorder().with_current_layer(i) ) with ctx: - if i in self.layers_to_capture: - aux_hidden_states.append(hidden_states + residual) layer = self.layers[i] hidden_states, residual = layer( positions=positions, forward_batch=forward_batch, hidden_states=hidden_states, residual=residual, + captured_last_layer_outputs=( + aux_hidden_states if i in self.layers_to_capture else None + ), ) if not self.pp_group.is_last_rank: