[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
|
import logging
|
||||||
from contextlib import nullcontext
|
from contextlib import nullcontext
|
||||||
from functools import lru_cache
|
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 torch
|
||||||
import triton
|
import triton
|
||||||
@@ -93,6 +93,7 @@ from sglang.srt.utils import (
|
|||||||
get_compiler_backend,
|
get_compiler_backend,
|
||||||
is_cuda,
|
is_cuda,
|
||||||
is_non_idle_and_non_empty,
|
is_non_idle_and_non_empty,
|
||||||
|
is_npu,
|
||||||
make_layers,
|
make_layers,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.custom_op import register_custom_op
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
_is_cuda = is_cuda()
|
_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
|
@triton.jit
|
||||||
@@ -816,6 +821,43 @@ class MiniMaxM2Attention(nn.Module):
|
|||||||
inner_state = q, k, v, forward_batch
|
inner_state = q, k, v, forward_batch
|
||||||
return None, forward_batch, inner_state
|
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):
|
def forward_core(self, intermediate_state):
|
||||||
hidden_states, forward_batch, inner_state = intermediate_state
|
hidden_states, forward_batch, inner_state = intermediate_state
|
||||||
if inner_state is None:
|
if inner_state is None:
|
||||||
@@ -830,11 +872,18 @@ class MiniMaxM2Attention(nn.Module):
|
|||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
s = self.forward_prepare(
|
if not _is_npu:
|
||||||
positions=positions,
|
s = self.forward_prepare(
|
||||||
hidden_states=hidden_states,
|
positions=positions,
|
||||||
forward_batch=forward_batch,
|
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)
|
return self.forward_core(s)
|
||||||
|
|
||||||
def op_prepare(self, state):
|
def op_prepare(self, state):
|
||||||
@@ -912,10 +961,16 @@ class MiniMaxM2DecoderLayer(nn.Module):
|
|||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
residual: Optional[torch.Tensor],
|
residual: Optional[torch.Tensor],
|
||||||
|
captured_last_layer_outputs: Optional[List[torch.Tensor]] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
# Self Attention
|
# Self Attention
|
||||||
hidden_states, residual = self.layer_communicator.prepare_attn(
|
hidden_states, residual = (
|
||||||
hidden_states, residual, forward_batch
|
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():
|
if not forward_batch.forward_mode.is_idle():
|
||||||
hidden_states = self.self_attn(
|
hidden_states = self.self_attn(
|
||||||
@@ -1099,14 +1154,15 @@ class MiniMaxM2Model(nn.Module):
|
|||||||
else get_global_expert_distribution_recorder().with_current_layer(i)
|
else get_global_expert_distribution_recorder().with_current_layer(i)
|
||||||
)
|
)
|
||||||
with ctx:
|
with ctx:
|
||||||
if i in self.layers_to_capture:
|
|
||||||
aux_hidden_states.append(hidden_states + residual)
|
|
||||||
layer = self.layers[i]
|
layer = self.layers[i]
|
||||||
hidden_states, residual = layer(
|
hidden_states, residual = layer(
|
||||||
positions=positions,
|
positions=positions,
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
residual=residual,
|
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:
|
if not self.pp_group.is_last_rank:
|
||||||
|
|||||||
Reference in New Issue
Block a user