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:
30
ixformer_sdk/contrib/vllm/layers/__init__.py
Normal file
30
ixformer_sdk/contrib/vllm/layers/__init__.py
Normal file
@@ -0,0 +1,30 @@
|
||||
from .llama import forward_smoothquant
|
||||
from .mixtral import mixtral_decoder_layer_forward
|
||||
|
||||
SUPPORT_REPLACE_METHOD = {
|
||||
"llama": forward_smoothquant,
|
||||
}
|
||||
|
||||
SUPPORT_REPLACE_LAYER = {
|
||||
"llama": None,
|
||||
}
|
||||
|
||||
|
||||
def get_replace_forward(name: str):
|
||||
try:
|
||||
method = SUPPORT_REPLACE_METHOD[name]
|
||||
except:
|
||||
raise ValueError(
|
||||
f"Only support replace names: {SUPPORT_REPLACE_METHOD.keys()}, but got {name}"
|
||||
)
|
||||
return method
|
||||
|
||||
|
||||
def get_replace_layer(name: str):
|
||||
try:
|
||||
layer = SUPPORT_REPLACE_LAYER[name]
|
||||
except:
|
||||
raise ValueError(
|
||||
f"Only support replace names: {SUPPORT_REPLACE_LAYER.keys()}, but got {name}"
|
||||
)
|
||||
return layer
|
||||
100
ixformer_sdk/contrib/vllm/layers/llama.py
Normal file
100
ixformer_sdk/contrib/vllm/layers/llama.py
Normal file
@@ -0,0 +1,100 @@
|
||||
from typing import Any, Dict, Iterable, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
from transformers import LlamaConfig
|
||||
|
||||
from vllm.attention import AttentionMetadata
|
||||
from vllm.config import CacheConfig
|
||||
from vllm.distributed import tensor_model_parallel_all_reduce
|
||||
from vllm.model_executor.layers.quantization.base_config import QuantizationConfig
|
||||
from vllm.model_executor.models.llama import LlamaDecoderLayer as VllmLlamaDecoderLayer
|
||||
|
||||
import vllm._custom_ops as ops
|
||||
# from ..overlap_comm import DecoderLayerOverlapComm, get_overlap_linear_method
|
||||
|
||||
|
||||
# This method is needed for support smoothquant no overlap forward
|
||||
def forward_smoothquant(
|
||||
input_ids: Optional[torch.Tensor],
|
||||
positions: torch.Tensor,
|
||||
kv_caches: List[torch.Tensor],
|
||||
attn_metadata: AttentionMetadata,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
self = None, # will be set by partial
|
||||
) -> torch.Tensor:
|
||||
dtype = self.dtype
|
||||
|
||||
def forward_smoothquant_mlp(self,x,scales):
|
||||
# gate_up_proj
|
||||
# Int8 Matrix multiply.
|
||||
bias = self.gate_up_proj.bias if not self.gate_up_proj.skip_bias_add else None
|
||||
gate_up = ops.w8a8(x, self.gate_up_proj.weight, scales, self.gate_up_proj.weight_scales, dtype)
|
||||
if bias:
|
||||
gate_up += bias
|
||||
|
||||
# act_fun
|
||||
x, scales = ops.silu_and_mul_smoothquant(gate_up, self.down_proj.smooth_scales)
|
||||
|
||||
# down_proj
|
||||
output_parallel = ops.w8a8(x, self.down_proj.weight, scales, self.down_proj.weight_scales, dtype)
|
||||
if self.down_proj.reduce_results and self.down_proj.tp_size > 1:
|
||||
output = tensor_model_parallel_all_reduce(output_parallel)
|
||||
else:
|
||||
output = output_parallel
|
||||
|
||||
if not self.down_proj.skip_bias_add:
|
||||
output = output + self.down_proj.bias if self.down_proj.bias is not None else output
|
||||
|
||||
return output
|
||||
|
||||
def forward_smoothquant_attn(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
scales: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
# qkv proj
|
||||
bias = self.qkv_proj.bias if not self.qkv_proj.skip_bias_add else None
|
||||
|
||||
qkv = ops.w8a8(hidden_states, self.qkv_proj.weight, scales, self.qkv_proj.weight_scales, dtype)
|
||||
if bias:
|
||||
qkv += bias
|
||||
|
||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||
q, k = self.rotary_emb(positions, q, k)
|
||||
attn_output = self.attn(q, k, v, kv_cache, attn_metadata)
|
||||
output, _ = self.o_proj(attn_output) # TODO
|
||||
return output
|
||||
|
||||
if inputs_embeds is not None:
|
||||
hidden_states = inputs_embeds
|
||||
else:
|
||||
hidden_states = self.get_input_embeddings(input_ids)
|
||||
residual = None
|
||||
for i in range(len(self.layers)):
|
||||
layer = self.layers[i]
|
||||
if residual is None:
|
||||
residual = hidden_states
|
||||
hidden_states, scales = ops.rms_norm_smoothquant(hidden_states,layer.input_layernorm.weight,layer.input_layernorm.variance_epsilon, layer.self_attn.qkv_proj.smooth_scales)
|
||||
else:
|
||||
hidden_states, residual, scales = ops.fused_add_rms_norm_smoothquant(hidden_states, residual, layer.input_layernorm.weight, layer.input_layernorm.variance_epsilon, layer.self_attn.qkv_proj.smooth_scales)
|
||||
|
||||
hidden_states = forward_smoothquant_attn(
|
||||
layer.self_attn,
|
||||
positions=positions,
|
||||
hidden_states=hidden_states,
|
||||
kv_cache=kv_caches[i],
|
||||
attn_metadata=attn_metadata,
|
||||
scales=scales,
|
||||
)
|
||||
|
||||
# Fully Connected
|
||||
hidden_states, residual, scales = ops.fused_add_rms_norm_smoothquant(hidden_states, residual, layer.post_attention_layernorm.weight, layer.post_attention_layernorm.variance_epsilon, layer.mlp.gate_up_proj.smooth_scales)
|
||||
|
||||
hidden_states = forward_smoothquant_mlp(layer.mlp, hidden_states, scales)
|
||||
|
||||
hidden_states, _ = self.norm(hidden_states, residual)
|
||||
return hidden_states
|
||||
|
||||
331
ixformer_sdk/contrib/vllm/layers/mixtral.py
Normal file
331
ixformer_sdk/contrib/vllm/layers/mixtral.py
Normal file
@@ -0,0 +1,331 @@
|
||||
import functools
|
||||
from typing import Dict, Optional, Tuple
|
||||
|
||||
import ixformer.inference.functions as ixf
|
||||
import torch
|
||||
|
||||
|
||||
def mixtral_decoder_layer_forward(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata,
|
||||
residual: Optional[torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
if self.use_int_w8a8:
|
||||
return w8a8_forward(
|
||||
self, positions, hidden_states, kv_cache, attn_metadata, residual
|
||||
)
|
||||
else:
|
||||
return original_forward(
|
||||
self, positions, hidden_states, kv_cache, attn_metadata, residual
|
||||
)
|
||||
|
||||
|
||||
def original_forward(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata,
|
||||
residual: Optional[torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
# Self Attention
|
||||
if residual is None:
|
||||
residual = hidden_states
|
||||
hidden_states = self.input_layernorm(hidden_states)
|
||||
else:
|
||||
hidden_states, residual = self.input_layernorm(hidden_states, residual)
|
||||
hidden_states = self.self_attn(
|
||||
positions=positions,
|
||||
hidden_states=hidden_states,
|
||||
kv_cache=kv_cache,
|
||||
attn_metadata=attn_metadata,
|
||||
)
|
||||
|
||||
# Fully Connected
|
||||
hidden_states, residual = self.post_attention_layernorm(hidden_states, residual)
|
||||
hidden_states = self.block_sparse_moe(hidden_states)
|
||||
return hidden_states, residual
|
||||
|
||||
|
||||
def dynamic_scaled_int8_quant(x):
|
||||
m, k = x.shape
|
||||
i8_x = x.new_empty([m, k], dtype=torch.int8, device="cuda")
|
||||
i8_scales = torch.empty([m], dtype=torch.float32, device="cuda")
|
||||
ixf.dynamic_scaled_int8_quant(i8_x, x, i8_scales)
|
||||
return i8_x, i8_scales
|
||||
|
||||
|
||||
def dynamic_w8a8(x, i8_weight, weight_scale):
|
||||
i8_x, i8_scale = dynamic_scaled_int8_quant(x)
|
||||
m, k = x.shape
|
||||
k, n = i8_weight.shape
|
||||
output = x.new_empty([m, n], dtype=x.dtype, device="cuda")
|
||||
ixf.w8a8(
|
||||
i8_x,
|
||||
i8_weight.transpose(0, 1),
|
||||
i8_scale,
|
||||
weight_scale,
|
||||
output=output,
|
||||
out_dtype=x.dtype,
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
def fused_rms_norm_quant_linear(
|
||||
self,
|
||||
hidden_states,
|
||||
ln_weight,
|
||||
eps,
|
||||
linear_weight,
|
||||
linear_weight_scale,
|
||||
residual=None,
|
||||
):
|
||||
# lower rouge
|
||||
# if residual is None:
|
||||
# residual = hidden_states
|
||||
# i8_hidden_states, _, i8_scales = ixf.residual_rms_norm_dynamic_int8(
|
||||
# input=hidden_states,
|
||||
# weight=ln_weight,
|
||||
# residual=None,
|
||||
# eps=eps,
|
||||
# )
|
||||
# else:
|
||||
# i8_hidden_states, residual, i8_scales = ixf.residual_rms_norm_dynamic_int8(
|
||||
# input=hidden_states,
|
||||
# weight=ln_weight,
|
||||
# residual=residual,
|
||||
# eps=eps,
|
||||
# )
|
||||
|
||||
if residual is None:
|
||||
residual = hidden_states
|
||||
hidden_states = self.input_layernorm(hidden_states)
|
||||
else:
|
||||
hidden_states, residual = self.input_layernorm(hidden_states, residual)
|
||||
i8_hidden_states, i8_scales = dynamic_scaled_int8_quant(hidden_states)
|
||||
|
||||
qkv = hidden_states.new_empty(hidden_states.shape[0], linear_weight.shape[1])
|
||||
ixf.w8a8(
|
||||
i8_hidden_states,
|
||||
linear_weight.transpose(0, 1),
|
||||
i8_scales,
|
||||
linear_weight_scale,
|
||||
output=qkv,
|
||||
out_dtype=hidden_states.dtype,
|
||||
)
|
||||
return qkv, residual
|
||||
|
||||
|
||||
def attention(qkv, positions, kv_cache, attn_metadata, self_attn):
|
||||
q, k, v = qkv.split(
|
||||
[self_attn.q_size, self_attn.kv_size, self_attn.kv_size], dim=-1
|
||||
)
|
||||
q, k = self_attn.rotary_emb(positions, q, k)
|
||||
attn_output = self_attn.attn(q, k, v, kv_cache, attn_metadata)
|
||||
return attn_output
|
||||
|
||||
|
||||
def fused_rms_norm_attention(
|
||||
self,
|
||||
hidden_states,
|
||||
ln_weight,
|
||||
eps,
|
||||
positions,
|
||||
kv_cache,
|
||||
attn_metadata,
|
||||
self_attn,
|
||||
residual=None,
|
||||
):
|
||||
hidden_states, residual = fused_rms_norm_quant_linear(
|
||||
self,
|
||||
hidden_states,
|
||||
ln_weight,
|
||||
eps,
|
||||
self_attn.qkv_proj.weight,
|
||||
self_attn.qkv_proj.weight_scale,
|
||||
residual,
|
||||
)
|
||||
hidden_states = attention(
|
||||
hidden_states, positions, kv_cache, attn_metadata, self_attn
|
||||
)
|
||||
|
||||
hidden_states = dynamic_w8a8(
|
||||
hidden_states, self_attn.o_proj.weight, self_attn.o_proj.weight_scale
|
||||
)
|
||||
# hidden_states,_ = self_attn.o_proj(hidden_states) # quant+linear+allreduce
|
||||
return hidden_states, residual
|
||||
|
||||
|
||||
def w8a8_forward(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata,
|
||||
residual: Optional[torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
|
||||
# qkv,_ = self.self_attn.qkv_proj(hidden_states)
|
||||
hidden_states, residual = fused_rms_norm_attention(
|
||||
self,
|
||||
hidden_states,
|
||||
self.input_layernorm.weight,
|
||||
self.input_layernorm.variance_epsilon,
|
||||
positions,
|
||||
kv_cache,
|
||||
attn_metadata,
|
||||
self.self_attn,
|
||||
residual,
|
||||
)
|
||||
|
||||
# allreduce
|
||||
tp_size = self.block_sparse_moe.experts.tp_size
|
||||
if tp_size > 1:
|
||||
from vllm.distributed import tensor_model_parallel_all_reduce
|
||||
|
||||
hidden_states = tensor_model_parallel_all_reduce(hidden_states)
|
||||
|
||||
# rms norm
|
||||
hidden_states, residual = self.post_attention_layernorm(hidden_states, residual)
|
||||
# moe
|
||||
hidden_states = fused_moe(
|
||||
hidden_states,
|
||||
self.block_sparse_moe.gate.weight,
|
||||
top_k=self.block_sparse_moe.experts.top_k,
|
||||
w1=self.block_sparse_moe.experts.w13_weight,
|
||||
w2=self.block_sparse_moe.experts.w2_weight,
|
||||
w1_scale=self.block_sparse_moe.experts.w13_weight_scale,
|
||||
w2_scale=self.block_sparse_moe.experts.w2_weight_scale,
|
||||
)
|
||||
|
||||
# allreduce
|
||||
if tp_size > 1:
|
||||
from vllm.distributed import tensor_model_parallel_all_reduce
|
||||
|
||||
hidden_states = tensor_model_parallel_all_reduce(hidden_states)
|
||||
|
||||
return hidden_states, residual
|
||||
|
||||
|
||||
def fused_experts(hidden_states, router_logits, top_k, w1, w2, w1_scale, w2_scale):
|
||||
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (num_tokens, k) dtype
|
||||
router_logits: (num_tokens, num_experts) torch.float32
|
||||
top_k int
|
||||
w1: (num_experts, 2n, k) torch.int8
|
||||
w2: (num_experts, k, n) torch.int8
|
||||
w1_scale: (num_experts, 2n) torch.float32
|
||||
w2_scale: (num_experts, k) torch.float32
|
||||
Returns
|
||||
final_hidden_states: (num_tokens, k) dtype
|
||||
"""
|
||||
|
||||
# topk_weight: (num_tokens, top_k) torch.float32
|
||||
# topk_ids: (num_tokens, top_k) torch.int32
|
||||
topk_weight, topk_ids = ixf.moe_topk_softmax(
|
||||
gating_output=router_logits,
|
||||
topk=top_k,
|
||||
renormalize=True,
|
||||
)
|
||||
|
||||
dtype = hidden_states.dtype
|
||||
num_tokens, num_experts = router_logits.shape
|
||||
expand_tokens = num_tokens * top_k
|
||||
|
||||
(
|
||||
src_to_dst,
|
||||
sorted_token_ids,
|
||||
expert_sizes_gpu,
|
||||
expert_sizes_cpu,
|
||||
) = ixf.moe_compute_token_index(
|
||||
topk_ids=topk_ids,
|
||||
num_experts=num_experts,
|
||||
)
|
||||
expert_sizes_cpu = expert_sizes_gpu.cpu()
|
||||
|
||||
# expand + reorder + quant
|
||||
# i8_hidden_states: (expand_tokens, k) torch.int8
|
||||
i8_hidden_states, a_scale = ixf.moe_expand_input_dynamic_scaled_int8(
|
||||
hidden_states=hidden_states,
|
||||
dst_to_src=sorted_token_ids,
|
||||
dst_tokens=expand_tokens,
|
||||
topk=top_k,
|
||||
src_to_dst=src_to_dst,
|
||||
topk_ids=None, # use smooth quant
|
||||
smooth_scales=None, # use smooth quant
|
||||
)
|
||||
|
||||
# w8a8 group gemm 1
|
||||
# pt_output_1: (expand_tokens, 2n) dtype
|
||||
pt_output_1 = ixf.moe_w8a8_group_gemm(
|
||||
input=i8_hidden_states,
|
||||
weight=w1,
|
||||
i_scales=a_scale,
|
||||
w_scales=w1_scale,
|
||||
output_dtype=dtype,
|
||||
tokens_per_experts=expert_sizes_cpu,
|
||||
dst_to_src=None,
|
||||
format="TN",
|
||||
)
|
||||
|
||||
# act + quant
|
||||
# pt_output_2: (expand_tokens, n) torch.int8
|
||||
pt_output_2, a2_scale = ixf.activation_dynamic_scaled_int8(
|
||||
input=pt_output_1,
|
||||
bias=None, # add gemm bias
|
||||
smooth_scales=None, # use smooth quant
|
||||
dst_to_src=sorted_token_ids,
|
||||
topk_ids=None, # add gemm bias or use smooth quant
|
||||
act_type="swiglu",
|
||||
)
|
||||
|
||||
# w8a8 group gemm 2 + reorder
|
||||
# pt_output_3: (expand_tokens, k) dtype
|
||||
pt_output_3 = ixf.moe_w8a8_group_gemm(
|
||||
input=pt_output_2,
|
||||
weight=w2,
|
||||
i_scales=a2_scale,
|
||||
w_scales=w2_scale,
|
||||
output_dtype=dtype,
|
||||
tokens_per_experts=expert_sizes_cpu,
|
||||
dst_to_src=sorted_token_ids,
|
||||
format="TN",
|
||||
)
|
||||
|
||||
# mul + reduce_sum
|
||||
# final_hidden_states: (num_tokens, k)
|
||||
final_hidden_states = ixf.moe_output_reduce_sum(
|
||||
input=pt_output_3.view(num_tokens, top_k, -1),
|
||||
topk_weight=topk_weight,
|
||||
)
|
||||
|
||||
return final_hidden_states
|
||||
|
||||
|
||||
def fused_moe(hidden_states, gate_weight, top_k, w1, w2, w1_scale, w2_scale):
|
||||
orig_shape = hidden_states.shape
|
||||
hidden_size = hidden_states.shape[-1]
|
||||
|
||||
hidden_states = hidden_states.view(-1, hidden_size)
|
||||
|
||||
# router_logits: (num_tokens, n_experts)
|
||||
# gate_weight: fp16
|
||||
router_logits = ixf.linear(hidden_states, gate_weight)
|
||||
router_logits = router_logits.to(torch.float32)
|
||||
|
||||
final_hidden_states = fused_experts(
|
||||
hidden_states,
|
||||
router_logits,
|
||||
top_k,
|
||||
w1,
|
||||
w2,
|
||||
w1_scale,
|
||||
w2_scale,
|
||||
)
|
||||
|
||||
return final_hidden_states.view(orig_shape)
|
||||
Reference in New Issue
Block a user