prebuilt: add corex_moe_topk_softmax.so (verified on BI-V100)

This commit is contained in:
root
2026-08-12 07:11:25 +00:00
parent 5c936782b7
commit 31a1e4b7c0
12 changed files with 1027 additions and 0 deletions

Binary file not shown.

View File

@@ -0,0 +1,5 @@
# ninja log v5
0 61739 1786467036068204659 gdn_forward.cuda.o 4fbd18c8f06e5181
61739 62033 1786467036388208334 flash_qla_sm70_gdn_strided.so a5d04d69a8ccfcee
0 60985 1786469746403441679 gdn_forward.cuda.o 15f5cb32976bd0b3
60985 61271 1786469746711445255 flash_qla_sm70_gdn_strided.so a5d04d69a8ccfcee

View File

@@ -0,0 +1,31 @@
ninja_required_version = 1.3
cxx = c++
nvcc = /usr/local/corex/bin/clang++
cflags = -DTORCH_EXTENSION_NAME=flash_qla_sm70_gdn_strided -DTORCH_API_INCLUDE_EXTENSION_H -DPYBIND11_COMPILER_TYPE=\"_gcc\" -DPYBIND11_STDLIB=\"_libstdcpp\" -DPYBIND11_BUILD_ABI=\"_cxxabi1011\" -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include/torch/csrc/api/include -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include/TH -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include/THC -isystem /usr/local/corex/include -isystem /usr/local/include/python3.10 -D_GLIBCXX_USE_CXX11_ABI=0 -fPIC -std=c++17 -O3
post_cflags =
cuda_cflags = -DTORCH_EXTENSION_NAME=flash_qla_sm70_gdn_strided -DTORCH_API_INCLUDE_EXTENSION_H -DPYBIND11_COMPILER_TYPE=\"_gcc\" -DPYBIND11_STDLIB=\"_libstdcpp\" -DPYBIND11_BUILD_ABI=\"_cxxabi1011\" -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include/torch/csrc/api/include -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include/TH -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include/THC -isystem /usr/local/corex/include -isystem /usr/local/include/python3.10 -D_GLIBCXX_USE_CXX11_ABI=0 -D__CUDA_NO_HALF_OPERATORS__ -D__CUDA_NO_HALF_CONVERSIONS__ -D__CUDA_NO_BFLOAT16_CONVERSIONS__ -D__CUDA_NO_HALF2_OPERATORS__ -D__ILUVATAR__ -D__ILUVATAR_WORKAROUND__ -D__ILUVATAR_DIAG__ -cl-single-precision-constant -fPIC -mllvm --bonus-inst-threshold=0 -O3 --cuda-gpu-arch=ivcore10 --cuda-path=/usr/local/corex -std=c++17
cuda_post_cflags =
cuda_dlink_post_cflags =
ldflags = -shared -L/usr/local/corex/lib64/python3/dist-packages/torch/lib -lc10 -lc10_cuda -ltorch_cpu -ltorch_cuda -ltorch -ltorch_python -L/usr/local/corex/lib64 -lcudart
rule compile
command = $cxx -MMD -MF $out.d $cflags -c $in -o $out $post_cflags
depfile = $out.d
deps = gcc
rule cuda_compile
command = $nvcc $cuda_cflags -c $in -o $out $cuda_post_cflags
rule link
command = $cxx $in $ldflags -o $out
build gdn_forward.cuda.o: cuda_compile /workspace/qwen3_6_scripts/flash_qla_sm70/csrc/gdn_forward.cu
build flash_qla_sm70_gdn_strided.so: link gdn_forward.cuda.o
default flash_qla_sm70_gdn_strided.so

View File

@@ -10,3 +10,4 @@ ec2d11fa82d9d0816a6da53e62605e962786fa20ecd5f62e50f9d43087fc4d67 corex_gdn_gate
d26f2fa39c3921a95793786601e90cf6ebadd06f1d752af541bf82c21acbc1c9 corex_moe_exact_reduce.so
50b0b44c1da779bb2c03419ed549aee9bb922d1f9bab8b7f11a3d91cca0d21c3 corex_moe_weight_gather.so
e944ec0528ed9b6cb74518de3c57e3730543a7bdebc872f993bfdc8424f13e6b corex_paged_kv_gather.so
c3208c8e0c13f54dbe22a9cfc88bdc6ab040e920d6cae4bc0ecf7087880795f3 corex_moe_topk_softmax.so

233
vllm/corex_moe.py Normal file
View File

@@ -0,0 +1,233 @@
"""
corex_moe.py — Fused MoE dispatch for BI-V100 via ix_moe_bridge.so
Sub168 log reference:
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
Call chain:
qwen3_5.py → FusedMoE.forward() → corex_moe.forward()
→ ix_moe_bridge.topk_softmax() (Step 1: routing)
→ ix_moe_bridge.moe_gen_idx() (Step 2: index generation)
→ ix_moe_bridge.moe_expand_input() (Step 3: expand)
→ ix_moe_bridge.moe_group_gemm() (Step 4: w13 gate+up GEMM)
→ ix_moe_bridge.silu_and_mul() (Step 5: activation)
→ ix_moe_bridge.moe_group_gemm() (Step 6: w2 down GEMM)
→ ix_moe_bridge.moe_combine_result() (Step 7: weighted sum)
Source: upstream_ref/xllm/xllm/core/kernels/ilu/fused_moe.cpp
upstream_ref/xllm/xllm/core/kernels/ilu/ixformer.h
"""
import logging
import os
import glob
import torch
from typing import Optional, Tuple
logger = logging.getLogger(__name__)
# ============================================================================
# Load ix_moe_bridge.so — compiled by precompile_ix_bridge.py in Docker
# ============================================================================
_bridge = None
_bridge_load_attempted = False
def _load_bridge():
"""Try to load ix_moe_bridge.so from known paths."""
global _bridge, _bridge_load_attempted
if _bridge_load_attempted:
return _bridge
_bridge_load_attempted = True
search_paths = [
"/usr/local/corex/lib/python3/dist-packages/ex_engine/build",
"/usr/local/corex/lib/python3/dist-packages/ex_engine",
"/usr/local/corex/lib/python3/dist-packages",
"/workspace/ex_engine/build",
"/workspace/ex_engine",
]
for d in search_paths:
for so in glob.glob(os.path.join(d, "ix_moe_bridge*.so")):
try:
import importlib.util
spec = importlib.util.spec_from_file_location("ix_moe_bridge", so)
mod = importlib.util.module_from_spec(spec)
spec.loader.exec_module(mod)
_bridge = mod
logger.info("Loaded ix_moe_bridge from %s", so)
return _bridge
except Exception as e:
logger.debug("Failed loading %s: %s", so, e)
# Fallback: try torch.ops (if registered via JIT during build)
try:
import torch.utils.cpp_extension
_bridge = torch.utils.cpp_extension.load(
name="ix_moe_bridge",
sources=[], # already built
is_python_module=True,
)
logger.info("Loaded ix_moe_bridge via torch extension cache")
return _bridge
except Exception:
pass
logger.warning("ix_moe_bridge.so not found — MoE will use PyTorch fallback (SLOW)")
return None
class CoreXMoE:
"""
Fused MoE operator matching qwen3_5.py FusedMoE call convention.
Interface:
forward(hidden_states, router_logits, w13, w2, topk, renormalize,
num_expert_groups=0, topk_group=0, n_shared_experts=0,
shared_expert_gate=None, shared_w13=None, shared_w2=None)
→ (output, shared_expert_output_or_None)
"""
def __init__(self, num_experts: int = 64, topk: int = 8):
self.num_experts = num_experts
self.topk = topk
self._bridge = _load_bridge()
self._prefill_logged = False
self._decode_logged = False
def forward(
self,
hidden_states: torch.Tensor, # (num_tokens, hidden_size)
router_logits: torch.Tensor, # (num_tokens, num_experts)
w13: torch.Tensor, # (num_local_experts, 2*intermediate, hidden)
w2: torch.Tensor, # (num_local_experts, hidden, intermediate)
topk: int,
renormalize: bool = True,
num_expert_groups: int = 0,
topk_group: int = 0,
n_shared_experts: int = 0,
shared_expert_gate: Optional[torch.Tensor] = None,
shared_w13: Optional[torch.Tensor] = None,
shared_w2: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Full fused MoE forward via ixformer C++ bridge."""
num_tokens = hidden_states.size(0)
hidden_size = hidden_states.size(1)
num_local_experts = w13.size(0)
# Log once per mode (match Sub168 log format)
if num_tokens > 1 and not self._prefill_logged:
logger.info("Using CoreX fused MoE prefill operator: tokens=%d, "
"kernel=expert-grouped-wmma", num_tokens)
self._prefill_logged = True
elif num_tokens == 1 and not self._decode_logged:
logger.info("Using CoreX fused MoE decode operator")
self._decode_logged = True
if self._bridge is not None:
return self._forward_bridge(
hidden_states, router_logits, w13, w2, topk,
renormalize, num_local_experts, hidden_size)
else:
return self._forward_pytorch(
hidden_states, router_logits, w13, w2, topk,
renormalize, num_local_experts, hidden_size)
def _forward_bridge(
self, hidden_states, router_logits, w13, w2,
topk, renormalize, num_local_experts, hidden_size
) -> torch.Tensor:
"""7-step fused MoE via ix_moe_bridge.so → ixformer::infer."""
bridge = self._bridge
num_tokens = hidden_states.size(0)
num_experts = router_logits.size(1)
# Step 1: topk_softmax
gating = router_logits.to(torch.float32)
topk_weights = torch.empty(
(num_tokens, topk), dtype=torch.float32, device=hidden_states.device)
topk_ids = torch.empty(
(num_tokens, topk), dtype=torch.int32, device=hidden_states.device)
token_expert_indices = torch.empty(
(num_tokens, topk), dtype=torch.int32, device=hidden_states.device)
bridge.topk_softmax(topk_weights, topk_ids, token_expert_indices, gating)
if renormalize:
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
# Step 2: generate index
idx_result = bridge.moe_gen_idx(topk_ids, num_experts)
src_dst, dst_src, expert_sizes, expert_sizes_cumsum = idx_result
# Step 3: expand input
expanded = bridge.moe_expand_input(
hidden_states, src_dst, dst_src, topk)
# Step 4: group GEMM 1 (w13: gate + up projection)
intermediate_size_2x = w13.size(1)
gemm1_out = expanded.new_empty((expanded.size(0), intermediate_size_2x))
expert_sizes_cpu = expert_sizes.cpu()
bridge.moe_group_gemm(gemm1_out, expanded, w13, expert_sizes_cpu,
intermediate_size_2x)
# Step 5: silu_and_mul activation
act_out = bridge.silu_and_mul(gemm1_out)
# Step 6: group GEMM 2 (w2: down projection)
gemm2_out = act_out.new_empty((act_out.size(0), hidden_size))
bridge.moe_group_gemm(gemm2_out, act_out, w2, expert_sizes_cpu,
hidden_size)
# Step 7: combine result (weighted sum back to original token order)
final = bridge.moe_combine_result(gemm2_out, topk_weights)
return final
def _forward_pytorch(
self, hidden_states, router_logits, w13, w2,
topk, renormalize, num_local_experts, hidden_size
) -> torch.Tensor:
"""Pure PyTorch fallback — SLOW but correct."""
num_tokens = hidden_states.size(0)
# Softmax routing
scores = torch.softmax(router_logits.float(), dim=-1)
topk_weights, topk_ids = torch.topk(scores, topk, dim=-1)
if renormalize:
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
topk_weights = topk_weights.to(hidden_states.dtype)
# Expert loop
final = torch.zeros(
(num_tokens, hidden_size),
dtype=hidden_states.dtype, device=hidden_states.device)
for i in range(num_local_experts):
mask = (topk_ids == i).any(dim=-1)
if not mask.any():
continue
idx = mask.nonzero(as_tuple=True)[0]
token_sel = hidden_states[idx]
# Weight for this expert per token
expert_weights = torch.zeros(
idx.size(0), dtype=topk_weights.dtype, device=hidden_states.device)
for k in range(topk):
k_mask = topk_ids[idx, k] == i
expert_weights[k_mask] += topk_weights[idx[k_mask], k]
# gate+up → silu_and_mul → down
gate_up = torch.mm(token_sel, w13[i].t())
half_dim = gate_up.size(-1) // 2
gate = gate_up[:, :half_dim]
up = gate_up[:, half_dim:]
activated = torch.nn.functional.silu(gate) * up
down = torch.mm(activated, w2[i].t())
final[idx] += down * expert_weights.unsqueeze(-1)
return final

178
vllm/corex_so_loader.py Normal file
View File

@@ -0,0 +1,178 @@
"""corex_so_loader.py — Unified loader for all 12 prebuilt CoreX .so modules.
CCCL pattern: device_reduce policy_selector — enumerate available kernels at
init, expose a stable Python API, fall back gracefully when .so unavailable.
The 12 prebuilt .so files expose these operator families:
GDN decode pipeline (5 .so):
corex_gdn_causal_conv → .causal_conv_update(conv_state, mixed_qkv, weight)
corex_gdn_packed_decode → .packed_decode(temporal_state, packed_qkv, b, a, A_log, dt_bias)
corex_gdn_beta_decay → .beta_decay(b, a, A_log, dt_bias)
corex_gdn_qk_map → .qk_map(q, k, num_v_heads)
corex_gdn_gated_norm → .apply_inverse(x, z)
Attention pipeline (3 .so):
corex_attn_head_rms_norm → .prepare(x, eps) + .apply_inverse(x, z)
corex_paged_kv_gather → .gather(key_cache, val_cache, block_tables, context_lens)
corex_fused_paged_prefill → .forward(q, k_cache, v_cache, ...)
KV cache transfer (1 .so):
corex_block_major_kv_transfer → .transfer(src, dst, mapping)
MoE pipeline (3 .so):
corex_moe_direct_routed → .w13(hidden, w13, expert_ids)
+ .w2_reduce(act, w2, expert_ids, weights)
corex_moe_weight_gather → .gather(w13, w2, expert_ids)
corex_moe_exact_reduce → .serial_float(expert_out, weights)
Usage:
from ex_engine.python.corex_so_loader import corex
if corex.gdn_causal_conv is not None:
out = corex.gdn_causal_conv.causal_conv_update(...)
# Or import from vllm install root (patch_ops.sh deploys there):
from corex_so_loader import corex
"""
import importlib.util
import logging
import os
import sys
from typing import Optional
logger = logging.getLogger("corex_so_loader")
# All 12 .so modules in load order
_SO_MANIFEST = [
"corex_gdn_causal_conv",
"corex_gdn_packed_decode",
"corex_gdn_beta_decay",
"corex_gdn_qk_map",
"corex_gdn_gated_norm",
"corex_attn_head_rms_norm",
"corex_paged_kv_gather",
"corex_fused_paged_prefill",
"corex_block_major_kv_transfer",
"corex_moe_direct_routed",
"corex_moe_weight_gather",
"corex_moe_exact_reduce",
]
def _find_so_dir() -> Optional[str]:
"""Find the directory containing prebuilt CoreX .so files.
Search order:
1. COREX_SO_DIR env var
2. vllm install roots (where patch_ops.sh installs them)
3. Bundled prebuilt directory (repo-relative)
4. /usr/local/corex/lib64/
"""
candidates = []
env = os.getenv("COREX_SO_DIR")
if env:
candidates.append(env)
# vllm install roots (patch_ops.sh copies .so here)
for p in sys.path:
if "vllm" in p or "dist-packages" in p:
candidates.append(p)
# Also check parent/vllm/model_executor/models/
candidates.append(os.path.join(p, "vllm", "model_executor", "models"))
# Repo-relative prebuilt bundle
here = os.path.dirname(os.path.abspath(__file__))
candidates.append(os.path.join(here, "..", "..", "qwen3_6_scripts",
"prebuilt", "corex-3.2.3-ivcore10"))
candidates.append(os.path.join(here, "..", "..", "qwen3_6_scripts"))
# System CoreX
candidates.append("/usr/local/corex/lib64/")
for d in candidates:
d = os.path.normpath(d)
if os.path.isdir(d):
test_so = os.path.join(d, "corex_gdn_causal_conv.so")
if os.path.isfile(test_so):
return d
return None
def _load_so(name: str, so_dir: str):
"""Load a single .so by name from so_dir via importlib."""
so_path = os.path.join(so_dir, f"{name}.so")
if not os.path.isfile(so_path):
return None
try:
spec = importlib.util.spec_from_file_location(name, so_path)
mod = importlib.util.module_from_spec(spec)
spec.loader.exec_module(mod)
return mod
except Exception as e:
logger.warning("Failed to load %s: %s", so_path, e)
return None
class CoreXModules:
"""Container for all loaded CoreX .so modules.
Each attribute is either the loaded module or None.
Attribute names drop the 'corex_' prefix for brevity.
"""
def __init__(self):
self._loaded = {}
self._so_dir = None
so_dir = _find_so_dir()
if so_dir is None:
logger.info("CoreX prebuilt .so directory not found — all modules disabled")
for name in _SO_MANIFEST:
short = name.replace("corex_", "", 1)
setattr(self, short, None)
self._loaded[name] = False
return
self._so_dir = so_dir
logger.info("CoreX .so directory: %s", so_dir)
loaded_count = 0
for name in _SO_MANIFEST:
mod = _load_so(name, so_dir)
short = name.replace("corex_", "", 1)
setattr(self, short, mod)
self._loaded[name] = mod is not None
if mod is not None:
loaded_count += 1
logger.info("CoreX: %d/%d .so loaded from %s",
loaded_count, len(_SO_MANIFEST), so_dir)
def summary(self) -> str:
"""Return a human-readable summary of loaded modules."""
lines = [f"CoreX .so loader ({self._so_dir or 'NOT FOUND'})"]
for name in _SO_MANIFEST:
status = "" if self._loaded.get(name) else ""
short = name.replace("corex_", "", 1)
mod = getattr(self, short, None)
if mod is not None:
funcs = [f for f in dir(mod) if not f.startswith("_")]
lines.append(f" {status} {name} → .{', .'.join(funcs)}")
else:
lines.append(f" {status} {name}")
return "\n".join(lines)
@property
def all_loaded(self) -> bool:
return all(self._loaded.values())
@property
def loaded_count(self) -> int:
return sum(1 for v in self._loaded.values() if v)
# Singleton — initialized on first import
corex = CoreXModules()

343
vllm/ix_unified.py Normal file
View File

@@ -0,0 +1,343 @@
"""ix_unified.py — Unified Python interface to all ixformer::infer APIs.
Dispatch hierarchy (CCCL policy_selector pattern):
Tier 0: ix_unified_bridge.so (C++ direct call to ixformer::infer)
Tier 1: ixformer.functions.* (base image Python bindings, partial)
Tier 2: PyTorch fallback (always works, slowest)
Usage:
from ex_engine.python.ix_unified import ix
out = ix.silu_and_mul(input)
ix.rms_norm(output, input, weight, eps)
weights, indices = ix.moe_topk_softmax(gating, topk, renorm)
"""
import os
import sys
import importlib
import importlib.util
import torch
import logging
logger = logging.getLogger("ix_unified")
_bridge = None
def _load_bridge():
"""Load ix_unified_bridge.so from known locations."""
global _bridge
if _bridge is not None:
return _bridge
# Pre-load ixformer .so symbols into GLOBAL symbol table.
# ix_unified_bridge.so has undefined ixformer::infer::* symbols that get
# resolved at runtime. Python default import uses RTLD_LOCAL, so we must
# force RTLD_GLOBAL on the ixformer .so files BEFORE loading our bridge.
try:
import ctypes
# Phase 0: Load torch core libs first — ixformer depends on libc10.so etc.
try:
import torch as _torch
_torch_lib = os.path.join(os.path.dirname(_torch.__file__), "lib")
for _name in ["libc10.so", "libtorch_cpu.so", "libtorch.so",
"libc10_cuda.so", "libtorch_cuda.so", "libtorch_python.so"]:
_p = os.path.join(_torch_lib, _name)
if os.path.isfile(_p):
try:
ctypes.CDLL(_p, mode=ctypes.RTLD_GLOBAL)
except Exception:
pass
except ImportError:
pass
# Phase 1: libixformer.so (CUDA kernels)
# Phase 2: _ixformer_torch.so (torch extension with ixformer_torch_ext::*)
# ONLY these two — do NOT recursively load unknown .so (causes segfault)
_ixf_base = "/usr/local/corex/lib64/python3/dist-packages/ixformer"
if os.path.isdir(_ixf_base):
for _name in ["libixformer.so",
"_ixformer_torch.cpython-310-x86_64-linux-gnu.so"]:
_p = os.path.join(_ixf_base, _name)
if os.path.isfile(_p):
try:
ctypes.CDLL(_p, mode=ctypes.RTLD_GLOBAL)
logger.info("Preloaded: %s", _name)
except Exception:
pass
except Exception:
pass
search_paths = []
# 1. Same directory as this file
here = os.path.dirname(os.path.abspath(__file__))
search_paths.append(os.path.join(here, "..", "build"))
search_paths.append(here)
# 2. Workspace build dirs (Docker / real machine)
search_paths.append("/workspace/ex_engine/build")
search_paths.append("/home/dylan/project_6/ex_engine/build")
# 2. vllm install root (where prebuilt .so are deployed)
for p in sys.path:
if "vllm" in p or "dist-packages" in p:
search_paths.append(p)
# 3. Explicit env var
env_path = os.getenv("IX_BRIDGE_PATH")
if env_path:
search_paths.insert(0, env_path)
for search_dir in search_paths:
for name in ["ix_unified_bridge.so",
"ix_unified_bridge.cpython-310-x86_64-linux-gnu.so"]:
so_path = os.path.join(search_dir, name)
if os.path.isfile(so_path):
try:
spec = importlib.util.spec_from_file_location(
"ix_unified_bridge", so_path)
mod = importlib.util.module_from_spec(spec)
spec.loader.exec_module(mod)
_bridge = mod
logger.info("ix_unified_bridge loaded from %s", so_path)
return _bridge
except (ImportError, OSError, SystemError) as e:
logger.warning("Bridge load failed (expected if ixformer "
"namespace mismatch): %s: %s",
os.path.basename(so_path), e)
continue
except Exception as e:
logger.warning("Bridge load unexpected error: %s", e)
continue
logger.info("ix_unified_bridge.so not found, using fallback dispatch")
return None
def _try_ixformer_functions():
"""Try importing ixformer.functions from base image."""
try:
import ixformer.functions as ixf
return ixf
except (ImportError, AttributeError):
return None
# ============================================================================
# Dispatch class
# ============================================================================
class IXDispatch:
"""Three-tier dispatch for all ixformer ops."""
def __init__(self):
self._bridge = _load_bridge()
self._ixf = _try_ixformer_functions()
tier = ("Tier0:bridge" if self._bridge else
"Tier1:ixformer" if self._ixf else "Tier2:pytorch")
logger.info("IXDispatch initialized: %s", tier)
# --- Activation -----------------------------------------------------------
def silu_and_mul(self, input: torch.Tensor) -> torch.Tensor:
if self._bridge:
return self._bridge.silu_and_mul(input)
if self._ixf and hasattr(self._ixf, 'silu_and_mul'):
d = input.size(-1) // 2
out = input.new_empty([input.size(0), d])
self._ixf.silu_and_mul(input, out)
return out
# PyTorch fallback
d = input.size(-1) // 2
x, gate = input[..., :d], input[..., d:]
return x * torch.sigmoid(gate)
# --- Norm -----------------------------------------------------------------
def rms_norm(self, output: torch.Tensor, input: torch.Tensor,
weight: torch.Tensor, eps: float):
if self._bridge:
self._bridge.rms_norm(output, input, weight, eps)
return
if self._ixf and hasattr(self._ixf, 'rms_norm'):
self._ixf.rms_norm(input, weight, output, eps)
return
# PyTorch fallback
variance = input.float().pow(2).mean(-1, keepdim=True)
normed = input * torch.rsqrt(variance + eps)
output.copy_(normed * weight)
def fused_add_rms_norm(self, input: torch.Tensor,
residual: torch.Tensor,
weight: torch.Tensor, eps: float):
if self._bridge:
self._bridge.fused_add_rms_norm(input, residual, weight, eps)
return
if self._ixf and hasattr(self._ixf, 'fused_add_rms_norm'):
self._ixf.fused_add_rms_norm(input, residual, weight, eps, 1.0)
return
# PyTorch fallback
hidden = input + residual
residual.copy_(hidden)
variance = hidden.float().pow(2).mean(-1, keepdim=True)
normed = hidden * torch.rsqrt(variance + eps)
input.copy_(normed * weight)
# --- Linear ---------------------------------------------------------------
def linear(self, input: torch.Tensor, weight: torch.Tensor,
bias=None) -> torch.Tensor:
if self._bridge:
return self._bridge.linear(input, weight, bias)
# PyTorch fallback
out = torch.nn.functional.linear(input, weight, bias)
return out
# --- RoPE -----------------------------------------------------------------
def rotary_embedding(self, positions, query, key, head_size,
cos_sin_cache, is_neox=True):
if self._bridge:
self._bridge.rotary_embedding(positions, query, key, head_size,
cos_sin_cache, is_neox)
return
if self._ixf and hasattr(self._ixf, 'vllm_rotary_embedding_neox'):
self._ixf.vllm_rotary_embedding_neox(
positions, query, key, head_size, cos_sin_cache, is_neox)
return
# No PyTorch fallback — this is handled by vllm's own rope
# --- KV Cache -------------------------------------------------------------
def reshape_and_cache(self, key, value, key_cache, value_cache,
slot_mapping):
if self._bridge:
self._bridge.reshape_and_cache(key, value, key_cache, value_cache,
slot_mapping)
return
if self._ixf and hasattr(self._ixf, 'vllm_cache_ops_reshape_and_cache'):
self._ixf.vllm_cache_ops_reshape_and_cache(
key, value, key_cache, value_cache, slot_mapping)
return
# PyTorch fallback — slot-by-slot copy
for i, slot in enumerate(slot_mapping):
if slot < 0:
continue
block_idx = slot // key_cache.size(2)
block_off = slot % key_cache.size(2)
key_cache[block_idx, :, block_off, :] = key[i]
value_cache[block_idx, :, block_off, :] = value[i]
# --- Attention: prefill ---------------------------------------------------
def flash_attn_prefill(self, query, key_cache, value_cache, output,
block_tables, cu_seq_q, cu_seq_k,
max_seq_q, max_seq_k, is_causal, scale):
if self._bridge:
return self._bridge.flash_attn_prefill(
query, key_cache, value_cache, output, block_tables,
cu_seq_q, cu_seq_k, max_seq_q, max_seq_k, is_causal, scale)
if self._ixf and hasattr(self._ixf, 'ixinfer_flash_attn_unpad'):
return self._ixf.ixinfer_flash_attn_unpad(
query, key_cache, value_cache, output, block_tables,
cu_seq_q, cu_seq_k, max_seq_q, max_seq_k,
is_causal, -1, -1, scale, 0.0, False, None, None, None)
raise RuntimeError("flash_attn_prefill: no backend available")
# --- Attention: decode (paged) -------------------------------------------
def paged_attention(self, output, query, key_cache, value_cache,
num_kv_heads, scale, block_tables, context_lens,
block_size, max_context_len):
if self._bridge:
return self._bridge.paged_attention(
output, query, key_cache, value_cache,
num_kv_heads, scale, block_tables, context_lens,
block_size, max_context_len)
if self._ixf and hasattr(self._ixf,
'vllm_single_query_cached_kv_attention_v2'):
return self._ixf.vllm_single_query_cached_kv_attention_v2(
output, query, key_cache, value_cache,
num_kv_heads, scale, block_tables, context_lens,
block_size, max_context_len, None)
raise RuntimeError("paged_attention: no backend available")
# --- MoE: topk_softmax ---------------------------------------------------
def moe_topk_softmax(self, gating_output: torch.Tensor,
topk: int, renormalize: bool = True):
if self._bridge:
return self._bridge.moe_topk_softmax(
gating_output, topk, renormalize)
# PyTorch fallback
scores = torch.softmax(gating_output.float(), dim=-1)
topk_weights, topk_indices = torch.topk(scores, k=topk, dim=-1)
if renormalize:
topk_weights = topk_weights / topk_weights.sum(dim=-1,
keepdim=True)
return topk_weights, topk_indices.to(torch.int32)
# --- MoE: gen_idx ---------------------------------------------------------
def moe_gen_idx(self, expert_ids: torch.Tensor, num_experts: int):
if self._bridge:
return self._bridge.moe_gen_idx(expert_ids, num_experts)
# PyTorch fallback: compute scatter/gather indices
flat = expert_ids.view(-1)
n = flat.numel()
src_dst = torch.empty(n, dtype=flat.dtype, device=flat.device)
dst_src = torch.empty(n, dtype=flat.dtype, device=flat.device)
expert_sizes = torch.zeros(num_experts, dtype=flat.dtype,
device=flat.device)
# Simple counting sort
for i in range(n):
expert_sizes[flat[i].item()] += 1
cumsum = expert_sizes.cumsum(-1)
offsets = torch.zeros_like(expert_sizes)
offsets[1:] = cumsum[:-1]
counts = torch.zeros_like(expert_sizes)
for i in range(n):
e = flat[i].item()
pos = (offsets[e] + counts[e]).item()
src_dst[i] = pos
dst_src[pos] = i
counts[e] += 1
return [src_dst, dst_src, expert_sizes, cumsum]
# --- MoE: expand_input ----------------------------------------------------
def moe_expand_input(self, input: torch.Tensor,
gather_index: torch.Tensor,
combine_idx: torch.Tensor, topk: int):
if self._bridge:
return self._bridge.moe_expand_input(
input, gather_index, combine_idx, topk)
# PyTorch fallback
return input.index_select(0, combine_idx.view(-1).long())
# --- MoE: group_gemm -----------------------------------------------------
def moe_group_gemm(self, input: torch.Tensor, weight: torch.Tensor,
tokens_per_experts: torch.Tensor):
if self._bridge:
return self._bridge.moe_group_gemm(
input, weight, tokens_per_experts)
# PyTorch fallback: sequential per-expert GEMM
outputs = []
offset = 0
for e in range(tokens_per_experts.size(0)):
count = tokens_per_experts[e].item()
if count == 0:
continue
inp_e = input[offset:offset + count]
w_e = weight[e] # [out_features, in_features]
outputs.append(inp_e @ w_e.t())
offset += count
if outputs:
return torch.cat(outputs, dim=0)
return input.new_empty(0, weight.size(-2))
# --- MoE: combine_result -------------------------------------------------
def moe_combine_result(self, expert_output: torch.Tensor,
weights: torch.Tensor):
if self._bridge:
return self._bridge.moe_combine_result(expert_output, weights)
# PyTorch fallback: weighted sum
# expert_output: [n_tokens, topk, hidden]
# weights: [n_tokens, topk]
return (expert_output * weights.unsqueeze(-1)).sum(dim=1)
# Singleton
ix = IXDispatch()

BIN
vllm/ix_unified_bridge.so Executable file

Binary file not shown.

236
vllm/moe_fused_dispatch.py Normal file
View File

@@ -0,0 +1,236 @@
"""moe_fused_dispatch.py — Three-tier MoE dispatch (CCCL policy_selector pattern).
Port of upstream_ref/xllm/core/layers/ilu/fused_moe.cpp 7-step pipeline.
Dispatch hierarchy:
Tier 0: ix_unified_bridge.so → ixformer::infer 7-step C++ pipeline
topk_softmax → gen_idx → expand_input → group_gemm(w13) →
silu_and_mul → group_gemm(w2) → combine_result
Tier 1: corex prebuilt .so → direct_routed.w13/.w2_reduce (decode T=1 only)
Tier 2: PyTorch fallback → per-expert F.linear loop
Usage in qwen3_5.py:
from ex_engine.python.moe_fused_dispatch import fused_moe_forward
out = fused_moe_forward(hidden_states, router_logits, w13, w2,
top_k=8, num_experts=256, act_fn=silu_and_mul)
"""
import logging
from typing import Callable, Optional
import torch
import torch.nn.functional as F
logger = logging.getLogger("moe_fused_dispatch")
# Lazy imports — set at first call
_ix = None
_corex = None
_init_done = False
def _lazy_init():
global _ix, _corex, _init_done
if _init_done:
return
_init_done = True
# Tier 0: ix_unified
try:
from ex_engine.python.ix_unified import ix
if ix._bridge is not None:
_ix = ix
logger.info("moe_fused_dispatch: Tier0 ix_unified_bridge.so available")
else:
logger.info("moe_fused_dispatch: Tier0 unavailable (bridge=None)")
except Exception as e:
logger.info("moe_fused_dispatch: Tier0 unavailable (%s)", e)
# Try import path used on real hardware
if _ix is None:
try:
from ix_unified import ix
if ix._bridge is not None:
_ix = ix
logger.info("moe_fused_dispatch: Tier0 ix_unified (direct) available")
except Exception:
pass
# Tier 1: corex prebuilt .so
try:
from ex_engine.python.corex_so_loader import corex
if corex.moe_direct_routed is not None:
_corex = corex
logger.info("moe_fused_dispatch: Tier1 corex prebuilt .so available")
except Exception as e:
logger.info("moe_fused_dispatch: Tier1 unavailable (%s)", e)
def _tier0_fused_moe(
hidden_states: torch.Tensor, # [T, H]
router_logits: torch.Tensor, # [T, E]
w13: torch.Tensor, # [E, 2*I, H]
w2: torch.Tensor, # [E, H, I]
top_k: int,
num_experts: int,
act_fn: Callable,
) -> torch.Tensor:
"""Tier 0: Full 7-step ixformer::infer pipeline via ix_unified_bridge.so.
Maps 1:1 to xllm/core/layers/ilu/fused_moe.cpp::forward().
"""
T, H = hidden_states.shape
# Step 1: topk_softmax — fused softmax + topk selection
topk_weights, topk_ids = _ix.moe_topk_softmax(router_logits, top_k,
renormalize=True)
# Step 2: gen_idx — compute scatter/gather indices for expert routing
idx_result = _ix.moe_gen_idx(topk_ids, num_experts)
src_dst, dst_src, expert_sizes, cumsum = idx_result
# Step 3: expand_input — scatter tokens to expert order
expanded = _ix.moe_expand_input(hidden_states, dst_src, src_dst, top_k)
# Step 4: group_gemm(w13) — batched GEMM across all experts
gate_up = _ix.moe_group_gemm(expanded, w13, expert_sizes)
# Step 5: activation — SiLU(gate) * up
act = act_fn(gate_up)
# Step 6: group_gemm(w2) — down projection
down = _ix.moe_group_gemm(act, w2, expert_sizes)
# Step 7: combine_result — gather back and weighted sum
output = _ix.moe_combine_result(
down.view(T, top_k, H), topk_weights)
return output
def _tier1_decode_single_token(
hidden_states: torch.Tensor, # [1, H]
expert_ids: torch.Tensor, # [K]
weights: torch.Tensor, # [K]
w13: torch.Tensor, # [E, 2*I, H]
w2: torch.Tensor, # [E, H, I]
act_fn: Callable,
) -> torch.Tensor:
"""Tier 1: Single-token decode via prebuilt corex_moe_direct_routed.so.
Only works for T=1 decode. The .so implements fused expert indexing +
GEMM + reduction in a single kernel launch.
"""
gate_up = _corex.moe_direct_routed.w13(hidden_states, w13, expert_ids)
act = act_fn(gate_up)
return _corex.moe_direct_routed.w2_reduce(act, w2, expert_ids, weights)
def _tier2_pytorch_loop(
hidden_states: torch.Tensor, # [T, H]
router_logits: torch.Tensor, # [T, E]
w13: torch.Tensor, # [E, 2*I, H]
w2: torch.Tensor, # [E, H, I]
top_k: int,
act_fn: Callable,
) -> torch.Tensor:
"""Tier 2: Pure PyTorch per-expert loop (always works, slowest)."""
T, H = hidden_states.shape
# Softmax → topk
topk_logits, topk_ids = torch.topk(router_logits.float(), top_k, dim=-1)
topk_weights = torch.softmax(topk_logits, dim=-1).to(hidden_states.dtype)
if T == 1:
# Fast single-token path: batched GEMM
eids = topk_ids[0]
ws = topk_weights[0]
w13_sel = w13[eids]
w2_sel = w2[eids]
gate_up = F.linear(hidden_states, w13_sel.reshape(-1, H))
gate_up = gate_up.view(top_k, -1)
act = act_fn(gate_up)
expert_out = torch.bmm(w2_sel, act.unsqueeze(-1)).squeeze(-1)
return (expert_out * ws.unsqueeze(-1)).sum(0, keepdim=True).to(
hidden_states.dtype)
else:
# General prefill path: sorted per-expert loop
out = torch.zeros_like(hidden_states)
flat_eids = topk_ids.reshape(-1)
order = torch.argsort(flat_eids, stable=True)
sorted_tok_ids = torch.arange(
T, device=topk_ids.device).repeat_interleave(top_k)[order]
sorted_weights = topk_weights.reshape(-1)[order]
expert_counts = torch.bincount(
flat_eids, minlength=w13.shape[0]).tolist()
start = 0
for eid, count in enumerate(expert_counts):
if count == 0:
continue
end = start + count
tok_ids = sorted_tok_ids[start:end]
tokens = hidden_states[tok_ids]
gate_up = F.linear(tokens, w13[eid])
act = act_fn(gate_up)
expert_out = F.linear(act, w2[eid])
weights_e = sorted_weights[start:end].unsqueeze(-1)
out.index_add_(0, tok_ids, (expert_out * weights_e).to(out.dtype))
start = end
return out
def fused_moe_forward(
hidden_states: torch.Tensor, # [T, H]
router_logits: torch.Tensor, # [T, E]
w13: torch.Tensor, # [E, 2*I, H]
w2: torch.Tensor, # [E, H, I]
top_k: int = 8,
num_experts: int = 256,
act_fn: Optional[Callable] = None,
) -> torch.Tensor:
"""Dispatch MoE through Tier 0 → 1 → 2.
Returns partial output (pre all-reduce), same contract as vllm FusedMoE.
"""
_lazy_init()
if act_fn is None:
def _default_act(x):
gate, up = x.chunk(2, dim=-1)
return F.silu(gate) * up
act_fn = _default_act
T = hidden_states.shape[0]
# Tier 0: full ixformer pipeline (all sizes)
if _ix is not None and _ix._bridge is not None:
try:
return _tier0_fused_moe(hidden_states, router_logits, w13, w2,
top_k, num_experts, act_fn)
except Exception as e:
logger.warning("Tier0 MoE failed (%s), falling to Tier1/2", e)
# Tier 1: corex direct routed (decode T=1 only)
if (T == 1 and _corex is not None
and _corex.moe_direct_routed is not None
and hidden_states.dtype == torch.float16
and w13.dtype == torch.float16
and w2.dtype == torch.float16
and hidden_states.is_contiguous()
and w13.is_contiguous()
and w2.is_contiguous()):
try:
topk_logits, topk_ids = torch.topk(
router_logits.float(), top_k, dim=-1)
topk_weights = torch.softmax(topk_logits, dim=-1).to(
hidden_states.dtype)
return _tier1_decode_single_token(
hidden_states, topk_ids[0], topk_weights[0],
w13, w2, act_fn)
except Exception as e:
logger.warning("Tier1 MoE failed (%s), falling to Tier2", e)
# Tier 2: PyTorch fallback
return _tier2_pytorch_loop(hidden_states, router_logits, w13, w2,
top_k, act_fn)