Files
project_6/ixformer_sdk/inference/overlap/w8a8_allreduce.py
project6-dev 87a19d2d00 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
2026-08-11 02:32:06 +00:00

155 lines
5.0 KiB
Python

import math
from typing import Optional
import ixformer.distributed as ixfd
import ixformer.functions as F
import torch
import torch.distributed as dist
from ixformer.distributed.overlap_comm import SplitOverlapComm
from ixformer.core import config as ixff_config
__all__ = ["w8a8_allreduce"]
class W8A8AllReduceOverlap(SplitOverlapComm):
def compute(
self,
input: torch.Tensor,
weight: torch.Tensor,
input_scale: torch.Tensor,
weight_scale: torch.Tensor,
bias: Optional[torch.Tensor] = None,
output: Optional[torch.Tensor] = None,
format: str = "TN",
out_dtype: torch.dtype = None,
comm_group=None,
split_ratio=0.5,
):
# compute the chunk size of input
input_chunk_sizes = [int(math.ceil(input.shape[0] * split_ratio))]
input_chunk_sizes.append(input.shape[0] - input_chunk_sizes[0])
# split input and input_scale
input_chunks = list(torch.split_with_sizes(input, input_chunk_sizes, dim=0))
input_scale_chunks = torch.split(
input_scale,
input_chunk_sizes,
)
# create output and split it
if output is None:
if out_dtype is None:
raise RuntimeError(
"w8a8 gemm need out_dtype argument when output is none."
)
output = torch.empty(
(input.shape[:-1] + (weight.shape[0],)),
dtype=out_dtype,
device=input.device,
)
out_chunks = torch.split(output, input_chunk_sizes)
# overlap gemm and allreduce
for chunk_idx in range(len(input_chunks)):
# submit gemm kernel into compute stream
with self.compute_stream_context(chunk_idx):
F.w8a8(
input=input_chunks[chunk_idx],
weight=weight,
i_scales=input_scale_chunks[chunk_idx],
w_scales=weight_scale,
bias=bias,
output=out_chunks[chunk_idx],
format=format,
persistent=chunk_idx != 0,
)
# recode compute stream and wait gemm
self.start_comm(chunk_idx)
# submit allreduce kernel into communication stream by set use_comm_stream to true
ixfd.all_reduce(
out_chunks[chunk_idx],
async_op=True,
group=self.comm_group,
use_comm_stream=True,
)
return output
_w8a8_allreduce_overlap = None
def w8a8_allreduce(
enable_overlap: bool,
input: torch.Tensor,
weight: torch.Tensor,
input_scale: torch.Tensor,
weight_scale: torch.Tensor,
bias: Optional[torch.Tensor] = None,
output: Optional[torch.Tensor] = None,
format: str = "TN",
out_dtype: torch.dtype = None,
comm_group=None,
split_ratio=0.5,
) -> torch.Tensor:
"""
Gemm(w8a8) + AllReduce
Args:
enable_overlap: whether enable gemm and allreduce overlap
input: shape: [M, K], dtype: int8, linear input
weight: shape: [N, K], dtype: int8, linear weight
input_scale: shape: [M], dtype: float32, quantized scale of input
weight_scale: shape: [N], dtype: float32, quantized scale of weight
bias: shape: [N], dtype: float16 or bfloat16, linear bias
output: shape: [M, N], dtype: float16 or bfloat16, allreduce output
format: options include TN, NN and NT
out_dtype: use the argument to decide to the dtype of output when output is None
comm_group: communication group
split_ratio: split the ratio of input.shape[0] when using overlap, range: (0, 1),
it will affect area of the overlap for gemm and allreduce.
Returns: output
"""
if (
enable_overlap
and ixff_config.IXFORMER_ENABLE_OVERLAP_COMM
and dist.is_initialized()
and dist.get_world_size(comm_group) > 1
and input.shape[0] > 1
):
global _w8a8_allreduce_overlap
if _w8a8_allreduce_overlap is None:
_w8a8_allreduce_overlap = W8A8AllReduceOverlap.dispatcher(
num_chunks=2, comm_group=comm_group
).forward
return _w8a8_allreduce_overlap(
input=input,
weight=weight,
input_scale=input_scale,
weight_scale=weight_scale,
bias=bias,
output=output,
format=format,
out_dtype=out_dtype,
split_ratio=split_ratio,
)
out = F.w8a8(
input=input,
weight=weight,
i_scales=input_scale,
w_scales=weight_scale,
bias=bias,
output=output,
format=format,
out_dtype=out_dtype,
)
if dist.get_world_size() > 1:
ixfd.all_reduce(out, op=ixfd.ReduceOp.SUM, async_op=True, group=comm_group)
return out