Files
project_6/ixformer_sdk/inference/functions/overlap_comm.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

85 lines
2.7 KiB
Python

import itertools
from functools import partial
from typing import Callable, Dict, Iterable, Tuple
import torch
import torch.distributed as dist
import ixformer.distributed as ixfd
from ixformer.core.dispatcher import Dispatcher
from ixformer.core.operator_autotuning import (
OperatorPreBaseRangeAutotuning,
sync_ranks_metric,
)
from ixformer.distributed import overlap_comm
from ixformer.inference.overlap.linear_mlp_overlap_comm import linear_mlp_overlap
from ixformer.distributed.overlap_comm import GemmMethod
__all__ = ["linear_allreduce_overlap", "linear_mlp_overlap"]
class LinearAllReducePreAutotuning(OperatorPreBaseRangeAutotuning, Dispatcher):
def __init__(self, comm_group, *args, **kwargs):
dist_barrier = True
if "dist_barrier" in kwargs:
dist_barrier = kwargs.pop("dist_barrier")
super().__init__(dist_barrier=dist_barrier, *args, **kwargs)
self._comm_group = comm_group
self._world_size = ixfd.get_group_world_size(comm_group)
@classmethod
def dispatcher_key(cls, comm_group, *args, **kwargs):
return (comm_group,)
def operators(self):
chunks = [2, 4]
gemm_algos = [GemmMethod.kCUINFER, GemmMethod.kCUBLAS, GemmMethod.kLIMITED_GEMM]
candidate_ops = [overlap_comm.GemmAllReduceSplitOverlapComm.native_forward]
for num_chunks, algo in itertools.product(chunks, gemm_algos):
candidate_ops.append(
partial(
overlap_comm.linear_allreduce_overlap,
num_chunks=num_chunks,
gemm_method=algo,
)
)
return candidate_ops
@property
def _gemm_shapes(self):
basic_k = [4096, 6114, 8192]
tp_k = [k // self._world_size for k in basic_k]
basic_k = tp_k
basic_m = (512, 1024, 2048, 4096, 8192)
shapes = set(itertools.product(basic_m, basic_k))
return shapes
def get_operator_key(self, input, *args, **kwargs):
ndim = input.ndim
shape = input.shape
if ndim == 1:
return (1, shape[0])
elif ndim == 2:
return shape
else:
return (sum(shape[:-1]), shape[-1])
def generate_operator_inputs(self) -> Iterable[Tuple[Tuple, Dict]]:
for m, kn in self._gemm_shapes:
input = torch.randn(m, kn, device="cuda", dtype=torch.half)
weight = torch.randn(kn, kn, device="cuda", dtype=torch.half)
yield (input, weight), {}
def perf_operator_time(self, op: Callable, *args, **kwargs) -> float:
op_time = super().perf_operator_time(op, *args, **kwargs)
return sync_ranks_metric(op_time, group=self._comm_group)
linear_allreduce_overlap = overlap_comm.linear_allreduce_overlap