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:
project6-dev
2026-08-11 02:31:56 +00:00
parent a8b16da5da
commit 87a19d2d00
250 changed files with 76690 additions and 0 deletions

View File

@@ -0,0 +1,496 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import ixformer.functions as ixf_F
def time_embed(t_emb, weight1, bias1, weight2, bias2):
# unet time_emd
# linear + silu + linear
emb = ixf_F.act_bias_mm(
t_emb, weight1, act_type="silu", bias=bias1, scale=1, trans_format="TN"
)
emb = ixf_F.act_bias_mm(
emb, weight2, act_type="none", bias=bias2, scale=1, trans_format="TN"
)
return emb
def ixf_layer_norm(input, normalized_shape, weight=None, bias=None, eps=1e-05):
return ixf_F.layernorm(input, weight, bias, normalized_shape)
def ixf_pt_scaled_dot_product_attention(
query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False
):
if (
not query.is_contiguous()
and query.transpose(1, 2).is_contiguous()
and key.transpose(1, 2).is_contiguous()
and value.transpose(1, 2).is_contiguous()
and attn_mask is None
):
batch_size, head_num, seq_len_q, head_dim = query.shape
_, _, seq_len_k, _ = key.shape
query = query.transpose(1, 2).view(batch_size * seq_len_q, head_num, head_dim)
key = key.transpose(1, 2).view(batch_size * seq_len_k, head_num, head_dim)
value = value.transpose(1, 2).view(batch_size * seq_len_k, head_num, head_dim)
cu_seqlens_q = torch.arange(
0,
seq_len_q * (batch_size + 1),
seq_len_q,
dtype=torch.int32,
device=query.device,
)
if seq_len_q == seq_len_k:
cu_seqlens_k = cu_seqlens_q
else:
cu_seqlens_k = torch.arange(
0,
seq_len_k * (batch_size + 1),
seq_len_k,
dtype=torch.int32,
device=query.device,
)
res = ixf_F.flash_attn_varlen_func(
query,
key,
value,
cu_seqlens_q.int(),
cu_seqlens_k.int(),
seq_len_q,
seq_len_k,
)
res = res.view(batch_size, seq_len_q, head_num, head_dim).transpose(1, 2)
return res
if not query.is_contiguous():
query = query.contiguous()
if not key.is_contiguous():
key = key.contiguous()
if not value.is_contiguous():
value = value.contiguous()
return ixf_F.scaled_dot_product_attention(
query, key, value, attn_mask=attn_mask, is_causal=is_causal
)
class UnetIxformerFunction:
def __init__(self) -> None:
self.ixf_linear = ixf_F.linear
self.pt_linear = F.linear
self.pt_layer_norm = F.layer_norm
self.pt_scaled_dot_product_attention = F.scaled_dot_product_attention
def __enter__(self):
F.linear = self.ixf_linear
F.layer_norm = ixf_layer_norm
F.scaled_dot_product_attention = ixf_pt_scaled_dot_product_attention
return self
def __exit__(self, exc_type, exc_val, exc_tb):
F.linear = self.pt_linear
F.layer_norm = self.pt_layer_norm
F.scaled_dot_product_attention = self.pt_scaled_dot_product_attention
if exc_tb is not None:
print(f"{exc_type} {exc_val}")
return False
return True
def ForwardWrapper(fun):
def wrap(*args, **kwargs):
with UnetIxformerFunction() as w:
return fun(*args, **kwargs)
return wrap
class IxformerComfyWrapper(nn.Module):
def __init__(self):
super().__init__()
self.is_ixf_wrapper = True
class Conv2dNhwcWrapper(IxformerComfyWrapper):
def __init__(self, module):
super().__init__()
module.weight.data = module.weight.permute(0, 2, 3, 1).contiguous()
module.bias.data = module.bias.float()
self.weight = module.weight.data
self.bias = module.bias.data
self.stride = module.stride
self.padding = module.padding
self.dilation = module.dilation
self.groups = module.groups
def forward(self, x):
h2 = ixf_F.conv2d(
x,
self.weight,
self.bias,
self.stride,
self.padding,
self.dilation,
self.groups,
)
return h2
class ResBlockNhwcWrapper(IxformerComfyWrapper):
def __init__(self, module) -> None:
super().__init__()
assert not module.updown
assert not module.use_scale_shift_norm
assert not module.skip_t_emb
assert not module.exchange_temb_dims
if isinstance(module.skip_connection, nn.Identity):
self.skip_connection = module.skip_connection
elif get_class_name(module.skip_connection) == "Conv2d":
self.skip_connection = Conv2dNhwcWrapper(module.skip_connection)
else:
raise NotImplementedError(
f"ResBlockNhwcWrapper support Conv2d or nn.Identity, but got {module.skip_connection}"
)
self.in_layers = module.in_layers
self.out_layers = module.out_layers
self.emb_layers = module.emb_layers
self.in_layers_conv = Conv2dNhwcWrapper(module.in_layers[2])
self.out_layers_conv = Conv2dNhwcWrapper(module.out_layers[3])
def forward(self, x, emb):
# x: nhwc
x1 = x
# print(x1.shape)
# fused group_norm silu
h = ixf_F.group_norm(
x1, # nchw->nhwc
self.in_layers[0].num_groups,
self.in_layers[0].weight,
self.in_layers[0].bias,
format=False,
act_type=1,
)
h = self.in_layers_conv(h)
emb_out = self.emb_layers(emb)
while len(emb_out.shape) < len(h.shape):
emb_out = emb_out[..., None]
h = h + emb_out.permute(0, 2, 3, 1)
# print(h.shape)
h = ixf_F.group_norm(
h,
self.out_layers[0].num_groups,
self.out_layers[0].weight,
self.out_layers[0].bias,
format=False,
act_type=1,
)
h = self.out_layers[2](h)
h = self.out_layers_conv(h)
# TODO: support other skip_connection
return self.skip_connection(x) + h
class DownsampleNhwcWrapper(IxformerComfyWrapper):
def __init__(self, module) -> None:
# TODO: support avg_pool_nd
super().__init__()
assert module.use_conv
self.channels = module.channels
self.op = Conv2dNhwcWrapper(module.op)
def forward(self, x):
assert x.shape[-1] == self.channels
return self.op(x)
def ffn_forward(self, x):
if get_class_name(self.net[0]) == "GEGLU":
net = self.net[1:]
geglu_net = self.net[0]
x = geglu_net.proj(x)
x = ixf_F.gelu_and_mul(x)
return net(x)
else:
return self.net(x)
# ComfyUI/comfy/ldm/modules/attention.py `class BasicTransformerBlock(nn.Module)`
def transformer_block_forward(self, x, context=None, transformer_options={}):
extra_options = {}
block = transformer_options.get("block", None)
block_index = transformer_options.get("block_index", 0)
transformer_patches = {}
transformer_patches_replace = {}
for k in transformer_options:
if k == "patches":
transformer_patches = transformer_options[k]
elif k == "patches_replace":
transformer_patches_replace = transformer_options[k]
else:
extra_options[k] = transformer_options[k]
extra_options["n_heads"] = self.n_heads
extra_options["dim_head"] = self.d_head
if self.ff_in:
x_skip = x
x = self.ff_in(self.norm_in(x))
if self.is_res:
x += x_skip
n = self.norm1(x)
if self.disable_self_attn:
context_attn1 = context
else:
context_attn1 = None
value_attn1 = None
if "attn1_patch" in transformer_patches:
patch = transformer_patches["attn1_patch"]
if context_attn1 is None:
context_attn1 = n
value_attn1 = context_attn1
for p in patch:
n, context_attn1, value_attn1 = p(
n, context_attn1, value_attn1, extra_options
)
if block is not None:
transformer_block = (block[0], block[1], block_index)
else:
transformer_block = None
attn1_replace_patch = transformer_patches_replace.get("attn1", {})
block_attn1 = transformer_block
if block_attn1 not in attn1_replace_patch:
block_attn1 = block
if block_attn1 in attn1_replace_patch:
if context_attn1 is None:
context_attn1 = n
value_attn1 = n
n = self.attn1.to_q(n)
context_attn1 = self.attn1.to_k(context_attn1)
value_attn1 = self.attn1.to_v(value_attn1)
n = attn1_replace_patch[block_attn1](
n, context_attn1, value_attn1, extra_options
)
n = self.attn1.to_out(n)
else:
n = self.attn1(n, context=context_attn1, value=value_attn1)
if "attn1_output_patch" in transformer_patches:
patch = transformer_patches["attn1_output_patch"]
for p in patch:
n = p(n, extra_options)
x += n
if "middle_patch" in transformer_patches:
patch = transformer_patches["middle_patch"]
for p in patch:
x = p(x, extra_options)
if self.attn2 is not None:
n = self.norm2(x)
if self.switch_temporal_ca_to_sa:
context_attn2 = n
else:
context_attn2 = context
value_attn2 = None
if "attn2_patch" in transformer_patches:
patch = transformer_patches["attn2_patch"]
value_attn2 = context_attn2
for p in patch:
n, context_attn2, value_attn2 = p(
n, context_attn2, value_attn2, extra_options
)
attn2_replace_patch = transformer_patches_replace.get("attn2", {})
block_attn2 = transformer_block
if block_attn2 not in attn2_replace_patch:
block_attn2 = block
if block_attn2 in attn2_replace_patch:
if value_attn2 is None:
value_attn2 = context_attn2
n = self.attn2.to_q(n)
context_attn2 = self.attn2.to_k(context_attn2)
value_attn2 = self.attn2.to_v(value_attn2)
n = attn2_replace_patch[block_attn2](
n, context_attn2, value_attn2, extra_options
)
n = self.attn2.to_out(n)
else:
n = self.attn2(n, context=context_attn2, value=value_attn2)
if "attn2_output_patch" in transformer_patches:
patch = transformer_patches["attn2_output_patch"]
for p in patch:
n = p(n, extra_options)
# x += n
# if self.is_res:
# x_skip = x
# x = self.ff(self.norm3(x))
x, x_skip = ixf_F.residual_layer_norm(
n,
self.norm3.normalized_shape,
self.norm3.weight,
self.norm3.bias,
x,
eps=self.norm3.eps,
is_post_ln=False,
)
x = ffn_forward(self.ff, x)
# x = ffn_forward(self.ff, self.norm3(x))
if self.is_res:
x += x_skip
return x
class SpatialTransformerNhwcWrapper(IxformerComfyWrapper):
def __init__(self, module):
super().__init__()
self.use_linear = module.use_linear
self.transformer_blocks = module.transformer_blocks
self.norm = module.norm
if not self.use_linear:
self.proj_in = Conv2dNhwcWrapper(module.proj_in)
self.proj_out = Conv2dNhwcWrapper(module.proj_out)
else:
self.proj_in = module.proj_in
self.proj_out = module.proj_out
@ForwardWrapper
def forward(self, x, context=None, transformer_options={}):
# note: if no context is given, cross-attention defaults to self-attention
if not isinstance(context, list):
context = [context] * len(self.transformer_blocks)
b, h, w, c = x.shape
x_in = x
# group_norm
x = ixf_F.group_norm(
x,
self.norm.num_groups,
self.norm.weight,
self.norm.bias,
format=False,
)
# conv2d
if not self.use_linear:
x = self.proj_in(x)
# n,(hw),c
x = x.view(x.shape[0], -1, x.shape[-1])
if self.use_linear:
x = self.proj_in(x)
for i, block in enumerate(self.transformer_blocks):
transformer_options["block_index"] = i
# x = block(x, context=context[i], transformer_options=transformer_options)
x = transformer_block_forward(
block, x, context=context[i], transformer_options=transformer_options
)
if self.use_linear:
x = self.proj_out(x)
x = x.view(b, h, w, c)
if not self.use_linear:
x = self.proj_out(x)
return x + x_in
class UpsampleNhwcWrapper(IxformerComfyWrapper):
def __init__(self, module) -> None:
# TODO: support mhwc interpolate
super().__init__()
self.dims = module.dims
self.use_conv = module.use_conv
self.channels = module.channels
if self.use_conv:
self.conv = Conv2dNhwcWrapper(module.conv)
def forward(self, x, output_shape=None):
# print("================== Upsample is running ==================")
assert x.shape[-1] == self.channels
assert len(x.shape) == 4
# nhwc -> nchw
if output_shape is not None:
assert len(output_shape) == 4
output_shape = [
output_shape[0],
output_shape[3],
output_shape[1],
output_shape[2],
]
x = x.permute(0, 3, 1, 2).contiguous()
if self.dims == 3:
shape = [x.shape[2], x.shape[3] * 2, x.shape[4] * 2]
if output_shape is not None:
shape[1] = output_shape[3]
shape[2] = output_shape[4]
else:
shape = [x.shape[2] * 2, x.shape[3] * 2]
if output_shape is not None:
shape[0] = output_shape[2]
shape[1] = output_shape[3]
# TODO: interpolate 支持 nhwc 去掉前后转置
x = F.interpolate(x, size=shape, mode="nearest")
# nchw -> nhwc
x = x.permute(0, 2, 3, 1).contiguous()
if self.use_conv:
x = self.conv(x)
return x
unet_wrappers = {
"Conv2d": Conv2dNhwcWrapper,
"ResBlock": ResBlockNhwcWrapper,
"Downsample": DownsampleNhwcWrapper,
"SpatialTransformer": SpatialTransformerNhwcWrapper,
"Upsample": UpsampleNhwcWrapper,
}
def get_class_name(module):
return module.__class__.__name__
def module_wrapper(module):
# 将原始的 module 封装为 nhwc 模式
module_name = get_class_name(module)
assert (
module_name == "TimestepEmbedSequential"
), f"ixformer unet_model_wrapper only support 'TimestepEmbedSequential' now, but got {module_name}"
num_sequential = len(module)
for idx_seq in range(num_sequential):
sub_module = module[idx_seq]
sub_module_name = get_class_name(sub_module)
# 判断模块是否已经封装
if not getattr(sub_module, "is_ixf_wrapper", False):
if sub_module_name in unet_wrappers:
module[idx_seq].forward = unet_wrappers[sub_module_name](
sub_module
).forward
module[idx_seq].is_ixf_wrapper = True
else:
raise NotImplementedError(f"{sub_module_name} not support")
return module