[NPU] add split_qkv_tp_rmsnorm_rope ops for minimax2 & fix eagle3 hidden states capture in dp attn mode (#23190)
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user