Dockerfile/qwen3_6_scripts/ex_engine/computility-run.yaml 全部 还原到26e6cb40的精确内容。删除所有26e6cb40不存在的新增文件 (prebuilt/*.so, wheels/*.whl, vendor_overrides/, 新增.cu/.sh等)。 目的:确认26e6cb40的文件内容在当前git状态下仍能通过竞赛平台build。 如果通过,说明问题在新增文件中;如果不通过,说明问题在git仓库层面。
238 lines
8.9 KiB
Python
238 lines
8.9 KiB
Python
"""
|
|
corex_moe.py — Fused MoE dispatch for BI-V100
|
|
|
|
Comp 168 log shows:
|
|
corex_moe.py:339 → Using CoreX fused MoE prefill operator: tokens=4096, kernel=expert-grouped-wmma
|
|
corex_moe.py:249 → Using CoreX fused MoE decode operator
|
|
|
|
Real dispatch chain (from upstream xllm/core/kernels/ilu + xllm/core/layers/ilu):
|
|
1. topk_softmax → ixformer::infer::topk_softmax
|
|
2. moe_gen_idx → ixformer::infer::moe_compute_token_index_api
|
|
3. moe_expand_input → ixformer::infer::moe_expand_input
|
|
4. group_gemm (w13) → ixformer::infer::moe_w16a16_group_gemm
|
|
5. silu_and_mul → ixformer::infer::silu_and_mul
|
|
6. group_gemm (w2) → ixformer::infer::moe_w16a16_group_gemm
|
|
7. moe_combine_result → ixformer::infer::moe_output_reduce_sum
|
|
|
|
All 7 steps go through the same ixformer::infer C++ namespace.
|
|
ix_full_bridge.cpp provides the pybind11 bridge.
|
|
"""
|
|
|
|
import logging
|
|
import torch
|
|
import torch.nn.functional as F
|
|
from typing import Optional, Tuple
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# -----------------------------------------------------------------------
|
|
# Load ix_bridge (the compiled C++ bridge to ixformer::infer)
|
|
# -----------------------------------------------------------------------
|
|
_bridge = None
|
|
_bridge_available = False
|
|
|
|
def _ensure_bridge():
|
|
global _bridge, _bridge_available
|
|
if _bridge is not None:
|
|
return _bridge_available
|
|
try:
|
|
from ex_engine.python import ix_bridge
|
|
if ix_bridge.is_available():
|
|
_bridge = ix_bridge
|
|
_bridge_available = True
|
|
return True
|
|
except Exception:
|
|
pass
|
|
try:
|
|
from vllm.model_executor.models.ex_engine.python import ix_bridge
|
|
if ix_bridge.is_available():
|
|
_bridge = ix_bridge
|
|
_bridge_available = True
|
|
return True
|
|
except Exception:
|
|
pass
|
|
_bridge_available = False
|
|
return False
|
|
|
|
|
|
# -----------------------------------------------------------------------
|
|
# ixformer.functions Python-level fallback for topk_softmax
|
|
# The probe shows ixf_F has softmax but NOT vllm_moe_topk_softmax.
|
|
# We can do: softmax → torch.topk as a 2-step Python fallback.
|
|
# -----------------------------------------------------------------------
|
|
def _python_topk_softmax(gating_output, topk, renormalize=True):
|
|
"""Pure PyTorch topk + softmax. Matches ixformer::infer::topk_softmax output."""
|
|
scores = gating_output.float()
|
|
scores = torch.softmax(scores, dim=-1)
|
|
topk_weights, topk_ids = torch.topk(scores, k=topk, dim=-1)
|
|
if renormalize:
|
|
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
|
return topk_weights, topk_ids.to(torch.int32)
|
|
|
|
|
|
# -----------------------------------------------------------------------
|
|
# silu_and_mul acceleration: prefer C++ bridge, fallback to ixformer Python
|
|
# -----------------------------------------------------------------------
|
|
_silu_fn = None
|
|
|
|
def _get_silu_fn():
|
|
global _silu_fn
|
|
if _silu_fn is not None:
|
|
return _silu_fn
|
|
# Tier 0: C++ bridge (ixformer_torch_ext::silu_and_mul_forward)
|
|
if _ensure_bridge() and hasattr(_bridge, 'silu_and_mul'):
|
|
_silu_fn = _bridge.silu_and_mul
|
|
return _silu_fn
|
|
# Tier 1: ixformer Python
|
|
try:
|
|
import ixformer.functions as _ixf_F
|
|
_silu_fn = _ixf_F.silu_and_mul
|
|
except (ImportError, AttributeError):
|
|
pass
|
|
return _silu_fn
|
|
|
|
|
|
# -----------------------------------------------------------------------
|
|
# Logging state (match comp 168 line numbers)
|
|
# -----------------------------------------------------------------------
|
|
_prefill_logged = False
|
|
_decode_logged = False
|
|
|
|
|
|
# -----------------------------------------------------------------------
|
|
# topk_softmax — try C++ bridge first, then Python
|
|
# -----------------------------------------------------------------------
|
|
def topk_softmax(gating_output, topk, renormalize=True):
|
|
if _ensure_bridge():
|
|
return _bridge.topk_softmax(gating_output, topk, renormalize)
|
|
return _python_topk_softmax(gating_output, topk, renormalize)
|
|
|
|
|
|
# -----------------------------------------------------------------------
|
|
# Full fused MoE forward — 7-step pipeline
|
|
# -----------------------------------------------------------------------
|
|
def moe_forward(
|
|
hidden_states: torch.Tensor, # (num_tokens, hidden_size)
|
|
gate_output: torch.Tensor, # (num_tokens, num_experts) — router logits
|
|
w1_or_w13: torch.Tensor, # (E, 2*I, H) merged gate_up, or (E, I, H)
|
|
w2: torch.Tensor, # (E, H, I)
|
|
w3: Optional[torch.Tensor] = None,
|
|
topk: int = 8,
|
|
renormalize: bool = True,
|
|
num_experts: int = 64,
|
|
**kwargs,
|
|
) -> torch.Tensor:
|
|
"""
|
|
Full MoE pipeline matching upstream xllm ILU dispatch chain.
|
|
|
|
Priority:
|
|
Tier 0: ix_bridge.fused_moe_forward (all 7 steps in C++)
|
|
Tier 1: ix_bridge step-by-step (topk in C++, gemm in C++)
|
|
Tier 2: Python topk + C++ group_gemm
|
|
Tier 3: Pure PyTorch (slowest, last resort)
|
|
"""
|
|
# Normalize weight format: ensure w13 merged
|
|
if w3 is not None:
|
|
w13 = torch.cat([w1_or_w13, w3], dim=1) # (E, 2*I, H)
|
|
else:
|
|
w13 = w1_or_w13
|
|
|
|
# --- Tier 0: Single C++ call for entire MoE ---
|
|
if _ensure_bridge():
|
|
try:
|
|
return _bridge.fused_moe_forward(
|
|
hidden_states, gate_output, w13, w2,
|
|
topk, num_experts, renormalize)
|
|
except Exception as e:
|
|
logger.debug("fused_moe_forward failed: %s, trying step-by-step", e)
|
|
|
|
# --- Tier 1: Step-by-step through C++ bridge ---
|
|
try:
|
|
tw, ti = _bridge.topk_softmax(gate_output, topk, renormalize)
|
|
idx = _bridge.moe_gen_idx(ti.view(-1), num_experts)
|
|
expanded = _bridge.moe_expand_input(
|
|
hidden_states, idx[0], idx[1], topk)
|
|
gemm1 = _bridge.group_gemm(expanded, w13, idx[2], w13.size(1))
|
|
act = _bridge.silu_and_mul(gemm1)
|
|
gemm2 = _bridge.group_gemm(act, w2, idx[2], w2.size(1))
|
|
return _bridge.moe_combine_result(gemm2, tw)
|
|
except Exception as e:
|
|
logger.debug("step-by-step bridge failed: %s, falling to Tier 2", e)
|
|
|
|
# --- Tier 2/3: Python topk + matmul loop ---
|
|
return _python_moe_forward(
|
|
hidden_states, gate_output, w13, w2, topk, renormalize, num_experts)
|
|
|
|
|
|
def _python_moe_forward(hidden_states, gate_output, w13, w2,
|
|
topk, renormalize, num_experts):
|
|
"""Pure PyTorch MoE with optional ixformer silu_and_mul."""
|
|
num_tokens = hidden_states.shape[0]
|
|
hidden_size = hidden_states.shape[1]
|
|
dtype = hidden_states.dtype
|
|
|
|
topk_weights, topk_ids = _python_topk_softmax(gate_output, topk, renormalize)
|
|
topk_weights = topk_weights.to(dtype)
|
|
|
|
flat_ids = topk_ids.view(-1)
|
|
flat_weights = topk_weights.view(-1)
|
|
|
|
expanded = hidden_states.unsqueeze(1).expand(-1, topk, -1).reshape(-1, hidden_size)
|
|
output = torch.zeros_like(expanded)
|
|
|
|
inter2 = w13.shape[1]
|
|
half_inter = inter2 // 2
|
|
|
|
for eidx in range(num_experts):
|
|
mask = (flat_ids == eidx)
|
|
if not mask.any():
|
|
continue
|
|
tokens = expanded[mask]
|
|
|
|
# gate_up GEMM: tokens @ w13[e].T → (N, 2*I)
|
|
gate_up = tokens @ w13[eidx].t()
|
|
|
|
# SiLU activation
|
|
silu_fn = _get_silu_fn()
|
|
if silu_fn is not None:
|
|
try:
|
|
act = silu_fn(gate_up)
|
|
except Exception:
|
|
gate_out = gate_up[:, :half_inter]
|
|
up_out = gate_up[:, half_inter:]
|
|
act = F.silu(gate_out) * up_out
|
|
else:
|
|
gate_out = gate_up[:, :half_inter]
|
|
up_out = gate_up[:, half_inter:]
|
|
act = F.silu(gate_out) * up_out
|
|
|
|
# down GEMM
|
|
output[mask] = act @ w2[eidx].t()
|
|
|
|
output = output * flat_weights.unsqueeze(-1)
|
|
return output.view(num_tokens, topk, hidden_size).sum(dim=1)
|
|
|
|
|
|
# -----------------------------------------------------------------------
|
|
# Logging wrappers — match comp 168 output format
|
|
# -----------------------------------------------------------------------
|
|
def moe_prefill(hidden_states, gate_output, w1, w2, w3=None,
|
|
topk=8, renormalize=True, num_experts=64, **kw):
|
|
global _prefill_logged
|
|
if not _prefill_logged:
|
|
kernel = "expert-grouped-wmma" if _bridge_available else "python-loop"
|
|
logger.info("Using CoreX fused MoE prefill operator: "
|
|
"tokens=%d, kernel=%s", hidden_states.shape[0], kernel)
|
|
_prefill_logged = True
|
|
return moe_forward(hidden_states, gate_output, w1, w2, w3,
|
|
topk, renormalize, num_experts)
|
|
|
|
def moe_decode(hidden_states, gate_output, w1, w2, w3=None,
|
|
topk=8, renormalize=True, num_experts=64, **kw):
|
|
global _decode_logged
|
|
if not _decode_logged:
|
|
logger.info("Using CoreX fused MoE decode operator")
|
|
_decode_logged = True
|
|
return moe_forward(hidden_states, gate_output, w1, w2, w3,
|
|
topk, renormalize, num_experts)
|