feat(CRITICAL): 从 GitHub 扫描搬运 ixformer SDK + xllm 完整 GDN/MoE 代码
来源:
1. Chranos/ixformer (GitHub) → ixformer_sdk/ (230 files, 70K lines)
- inference/functions/vllm.py: vllm_moe_topk_softmax 完整实现 (2033 lines)
- inference/functions/moe.py: MoE ops 完整实现 (1380 lines)
- contrib/vllm_flash_attn/: FA2 Python 接口 (1018 lines)
- contrib/tgi/fused_moe.py: TGI fused MoE (429 lines)
- csrc/include/ixformer/: C++ kernel headers + cmake
2. Deep-Spark/xllm (GitHub) → upstream_ref/xllm_latest/ (+15 files)
- npu_torch/qwen3_5_decoder_layer_impl.cpp/.h
- npu_torch/qwen3_5_gated_delta_net.cpp/.h
- npu_torch/qwen3_next_*.cpp/.h (6 files)
- npu_torch/attention.cpp/.h + fused_moe.cpp/.h + CMakeLists.txt
- models/llm/qwen3_5.h + qwen3_5_mtp.h + qwen3_next.h
- models/vlm/qwen3_5.h
调用链完整性:
ixformer_sdk/inference/functions/vllm.py
→ ops.infer.moe_topk_softmax() (C++ 层)
→ 这就是 base 镜像 libixformer.so 里的实现
upstream_ref/xllm_latest/core/layers/ilu/fused_moe.cpp
→ ixformer::infer::topk_softmax() (直接 C++ 调用)
→ ixformer::infer::group_gemm() → 完整 7-step MoE pipeline
This commit is contained in:
162
ixformer_sdk/train/speedformer/layers/baichuan/attention.py
Normal file
162
ixformer_sdk/train/speedformer/layers/baichuan/attention.py
Normal file
@@ -0,0 +1,162 @@
|
||||
import math
|
||||
import warnings
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from flash_attn import flash_attn_func, flash_attn_varlen_func
|
||||
from flash_attn.bert_padding import index_first_axis, pad_input, unpad_input
|
||||
from ixformer.train.speedformer.models.baichuan.configuration_baichuan import BaichuanConfig
|
||||
from ixformer.train.speedformer.models.baichuan.modeling_baichuan import Attention
|
||||
|
||||
from ixformer.train.functions.fused_rope import fused_apply_rotary_pos_emb
|
||||
from ixformer.train.speedformer.layers.rotary_pos_embedding import RotaryEmbedding
|
||||
|
||||
from ixformer.train.speedformer.layers.lazy import LazyInitContext
|
||||
|
||||
|
||||
class FlashAttention(Attention):
|
||||
# 这个类主要的改进包含:1. apply_rotary_pos_emb;2. flash-attn 代替 native attention
|
||||
def __init__(self, config: BaichuanConfig):
|
||||
super().__init__(config)
|
||||
self.rotary_emb = RotaryEmbedding(self.head_dim)
|
||||
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
past_key_value: Optional[Tuple[torch.Tensor]] = None,
|
||||
output_attentions: bool = False,
|
||||
use_cache: bool = False,
|
||||
) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
|
||||
bsz, q_len, _ = hidden_states.size()
|
||||
|
||||
proj = self.W_pack(hidden_states)
|
||||
proj = proj.unflatten(-1, (3, self.hidden_size)).unsqueeze(0).transpose(0, -2).squeeze(-2)
|
||||
|
||||
# fused_apply_rotary_pos_emb need qk to be in "sbhd", v stay in "bshd"
|
||||
query_states = proj[0].view(bsz, q_len, self.num_heads, self.head_dim).transpose(0, 1).contiguous()
|
||||
key_states = proj[1].view(bsz, q_len, self.num_heads, self.head_dim).transpose(0, 1).contiguous()
|
||||
value_states = proj[2].view(bsz, q_len, self.num_heads, self.head_dim)
|
||||
|
||||
kv_seq_len = key_states.shape[0]
|
||||
if past_key_value is not None:
|
||||
kv_seq_len += past_key_value[0].shape[0]
|
||||
|
||||
# fused_apply_rotary_pos_emb need emb in float32
|
||||
emb = self.rotary_emb(kv_seq_len).to(dtype=torch.float32)
|
||||
query_states = fused_apply_rotary_pos_emb(query_states, emb)
|
||||
key_states = fused_apply_rotary_pos_emb(key_states, emb)
|
||||
|
||||
if past_key_value is not None:
|
||||
# reuse k, v, self_attention
|
||||
key_states = torch.cat([past_key_value[0], key_states], dim=0)
|
||||
value_states = torch.cat([past_key_value[1], value_states], dim=0)
|
||||
|
||||
past_key_value = (key_states, value_states) if use_cache else None
|
||||
|
||||
# after fused_apply_rotary_pos_emb, qk change to "bshd" for flashattn or "bhsd" for sdpa
|
||||
if attention_mask is None: # flash-attn
|
||||
query_states = query_states.transpose(0, 1).contiguous()
|
||||
key_states = key_states.transpose(0, 1).contiguous()
|
||||
else: # sdpa
|
||||
query_states = query_states.permute(1, 2, 0, 3).contiguous()
|
||||
key_states = key_states.permute(1, 2, 0, 3).contiguous()
|
||||
value_states = value_states.transpose(1, 2).contiguous()
|
||||
|
||||
'''
|
||||
if attention_mask is not None:
|
||||
batch_size = query_states.shape[0] # bsz, q_len, self.num_heads, self.head_dim
|
||||
query_states, key_states, value_states, indices_q, cu_seq_lens, max_seq_lens = self._upad_input(
|
||||
query_states, key_states, value_states, attention_mask, q_len
|
||||
)
|
||||
|
||||
cu_seqlens_q, cu_seqlens_k = cu_seq_lens
|
||||
max_seqlen_in_batch_q, max_seqlen_in_batch_k = max_seq_lens
|
||||
attn_output_unpad = flash_attn_varlen_func(
|
||||
query_states,
|
||||
key_states,
|
||||
value_states,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
cu_seqlens_k=cu_seqlens_k,
|
||||
max_seqlen_q=max_seqlen_in_batch_q,
|
||||
max_seqlen_k=max_seqlen_in_batch_k,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
causal=True,
|
||||
)
|
||||
|
||||
attn_output = pad_input(attn_output_unpad, indices_q, batch_size, q_len)
|
||||
else:
|
||||
attn_output = flash_attn_func(
|
||||
query_states, key_states, value_states, 0.0, softmax_scale=None, causal=True
|
||||
)
|
||||
'''
|
||||
attn_output = self._flash_attention_forward(
|
||||
query_states, key_states, value_states, q_len, attention_mask, dropout=0.0
|
||||
)
|
||||
|
||||
attn_output = attn_output.reshape(bsz, q_len, self.hidden_size)
|
||||
attn_output = self.o_proj(attn_output)
|
||||
|
||||
if not output_attentions:
|
||||
attn_weights = None
|
||||
|
||||
return attn_output, attn_weights, past_key_value
|
||||
|
||||
def _flash_attention_forward(
|
||||
self,
|
||||
query_states: torch.Tensor,
|
||||
key_states: torch.Tensor,
|
||||
value_states: torch.Tensor,
|
||||
query_length: int,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
dropout=0.0,
|
||||
softmax_scale=None
|
||||
):
|
||||
if attention_mask is not None:
|
||||
attn_output = torch.nn.functional.scaled_dot_product_attention(
|
||||
query_states,
|
||||
key_states,
|
||||
value_states,
|
||||
attn_mask=attention_mask,
|
||||
dropout_p=0.0,
|
||||
# The q_len > 1 is necessary to match with AttentionMaskConverter.to_causal_4d that does not create a causal mask in case q_len == 1.
|
||||
is_causal=query_length > 1,
|
||||
)
|
||||
attn_output = attn_output.transpose(1, 2).contiguous()
|
||||
|
||||
else:
|
||||
attn_output = flash_attn_func(
|
||||
query_states, key_states, value_states, dropout, softmax_scale=softmax_scale, causal=self.is_causal
|
||||
)
|
||||
|
||||
return attn_output
|
||||
|
||||
class BaichuanAttention(FlashAttention):
|
||||
def __init__(self) -> None:
|
||||
raise NotImplementedError(
|
||||
"BaichuanAttention is not implemented as a physical class. "
|
||||
"It is meant to be used only with the from_native_module interface to Convert a native BaichuanAttention module to LlamaAttention module provided above."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def from_native_module(module: nn.Module, *args, **kwargs) -> nn.Module:
|
||||
|
||||
LazyInitContext.materialize(module)
|
||||
|
||||
# try to get normalized_shape, eps, elementwise_affine from the module
|
||||
config = getattr(module, "config")
|
||||
|
||||
attention = FlashAttention(
|
||||
config=config,
|
||||
)
|
||||
|
||||
attention.W_pack.weight = module.W_pack.weight
|
||||
attention.o_proj.weight = module.o_proj.weight
|
||||
|
||||
return attention
|
||||
141
ixformer_sdk/train/speedformer/layers/baichuan/baichuan_model.py
Normal file
141
ixformer_sdk/train/speedformer/layers/baichuan/baichuan_model.py
Normal file
@@ -0,0 +1,141 @@
|
||||
import math
|
||||
import warnings
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from ixformer.train.speedformer.models.baichuan.configuration_baichuan import BaichuanConfig
|
||||
from ixformer.train.speedformer.models.baichuan.modeling_baichuan import BaichuanModel
|
||||
from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast
|
||||
from transformers.utils import logging, ContextManagers
|
||||
|
||||
|
||||
from ixformer.train.speedformer.layers.lazy import LazyInitContext
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
|
||||
class IXFBaichuanModel(BaichuanModel):
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.LongTensor = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
past_key_values: Optional[List[torch.FloatTensor]] = None,
|
||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||
use_cache: Optional[bool] = None,
|
||||
output_attentions: Optional[bool] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
return_dict: Optional[bool] = None,
|
||||
) -> Union[Tuple, BaseModelOutputWithPast]:
|
||||
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
|
||||
output_hidden_states = (
|
||||
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
|
||||
)
|
||||
use_cache = use_cache if use_cache is not None else self.config.use_cache
|
||||
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
|
||||
# retrieve input_ids and inputs_embeds
|
||||
if input_ids is not None and inputs_embeds is not None:
|
||||
raise ValueError(
|
||||
"You cannot specify both decoder_input_ids and decoder_inputs_embeds at the same time")
|
||||
elif input_ids is not None:
|
||||
batch_size, seq_length = input_ids.shape
|
||||
elif inputs_embeds is not None:
|
||||
batch_size, seq_length, _ = inputs_embeds.shape
|
||||
else:
|
||||
raise ValueError(
|
||||
"You have to specify either decoder_input_ids or decoder_inputs_embeds")
|
||||
|
||||
seq_length_with_past = seq_length
|
||||
past_key_values_length = 0
|
||||
|
||||
if past_key_values is not None:
|
||||
past_key_values_length = past_key_values[0][0].shape[2]
|
||||
seq_length_with_past = seq_length_with_past + past_key_values_length
|
||||
|
||||
if position_ids is None:
|
||||
device = input_ids.device if input_ids is not None else inputs_embeds.device
|
||||
position_ids = torch.arange(
|
||||
past_key_values_length, seq_length + past_key_values_length, dtype=torch.long, device=device
|
||||
)
|
||||
position_ids = position_ids.unsqueeze(0).view(-1, seq_length)
|
||||
else:
|
||||
position_ids = position_ids.view(-1, seq_length).long()
|
||||
|
||||
if inputs_embeds is None:
|
||||
inputs_embeds = self.embed_tokens(input_ids)
|
||||
|
||||
hidden_states = inputs_embeds
|
||||
|
||||
if self.gradient_checkpointing and self.training:
|
||||
if use_cache:
|
||||
logger.warning_once(
|
||||
"`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."
|
||||
)
|
||||
use_cache = False
|
||||
|
||||
# decoder layers
|
||||
all_hidden_states = () if output_hidden_states else None
|
||||
all_self_attns = () if output_attentions else None
|
||||
next_decoder_cache = () if use_cache else None
|
||||
|
||||
for idx, decoder_layer in enumerate(self.layers):
|
||||
if output_hidden_states:
|
||||
all_hidden_states += (hidden_states,)
|
||||
|
||||
past_key_value = past_key_values[idx] if past_key_values is not None else None
|
||||
|
||||
if self.gradient_checkpointing and self.training:
|
||||
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
# None for past_key_value
|
||||
return module(*inputs, output_attentions, None)
|
||||
|
||||
return custom_forward
|
||||
|
||||
layer_outputs = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(decoder_layer),
|
||||
hidden_states,
|
||||
attention_mask,
|
||||
position_ids,
|
||||
None,
|
||||
)
|
||||
else:
|
||||
layer_outputs = decoder_layer(
|
||||
hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
position_ids=position_ids,
|
||||
past_key_value=past_key_value,
|
||||
output_attentions=output_attentions,
|
||||
use_cache=use_cache,
|
||||
)
|
||||
|
||||
hidden_states = layer_outputs[0]
|
||||
|
||||
if use_cache:
|
||||
next_decoder_cache += (
|
||||
layer_outputs[2 if output_attentions else 1],)
|
||||
|
||||
if output_attentions:
|
||||
all_self_attns += (layer_outputs[1],)
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
|
||||
# add hidden states from the last decoder layer
|
||||
if output_hidden_states:
|
||||
all_hidden_states += (hidden_states,)
|
||||
|
||||
next_cache = next_decoder_cache if use_cache else None
|
||||
if not return_dict:
|
||||
return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None)
|
||||
return BaseModelOutputWithPast(
|
||||
last_hidden_state=hidden_states,
|
||||
past_key_values=next_cache,
|
||||
hidden_states=all_hidden_states,
|
||||
attentions=all_self_attns,
|
||||
)
|
||||
53
ixformer_sdk/train/speedformer/layers/baichuan/mlp.py
Normal file
53
ixformer_sdk/train/speedformer/layers/baichuan/mlp.py
Normal file
@@ -0,0 +1,53 @@
|
||||
import math
|
||||
import warnings
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
import ixformer.train.functions as F
|
||||
from ixformer.train.speedformer.models.baichuan.configuration_baichuan import BaichuanConfig
|
||||
from ixformer.train.speedformer.models.baichuan.modeling_baichuan import MLP
|
||||
from transformers.utils import logging
|
||||
|
||||
from ixformer.train.speedformer.layers.lazy import LazyInitContext
|
||||
|
||||
|
||||
class BaseMLP(MLP):
|
||||
"""
|
||||
这个层主要的优化点是:将linear1(act(cat(linear2(x), linear3(x))))的结构变成 linear1(act(linear23(x)))
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, intermediate_size, hidden_act):
|
||||
super().__init__(hidden_size, intermediate_size, hidden_act)
|
||||
self.gate_up = nn.Linear(
|
||||
hidden_size, intermediate_size * 2, bias=False)
|
||||
del self.gate_proj, self.up_proj
|
||||
del self.act_fn
|
||||
|
||||
def forward(self, x):
|
||||
res = self.gate_up(x)
|
||||
down_proj = self.down_proj(F.swiglu(res))
|
||||
return down_proj
|
||||
|
||||
|
||||
class IXFBaichuanMLP(BaseMLP):
|
||||
def __init__(self) -> None:
|
||||
raise NotImplementedError(
|
||||
"IXFLlamaMLP is not implemented as a physical class. "
|
||||
"It is meant to be used only with the from_native_module interface to Convert a native LlamaAttention module to IXFLlamaMLP module provided above."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def from_native_module(module: nn.Module, *args, **kwargs) -> nn.Module:
|
||||
hidden_size, intermediate_size = module.gate_proj.in_features, module.gate_proj.out_features
|
||||
hidden_act = "silu"
|
||||
|
||||
mlp = BaseMLP(hidden_size=hidden_size,
|
||||
intermediate_size=intermediate_size, hidden_act=hidden_act)
|
||||
|
||||
mlp.gate_up.weight.data = torch.concat(
|
||||
(module.gate_proj.weight.data, module.up_proj.weight.data), dim=0)
|
||||
mlp.down_proj.weight.data = module.down_proj.weight.data
|
||||
|
||||
return mlp
|
||||
Reference in New Issue
Block a user