diff --git a/qwen3_6_scripts/flash_qla_sm70/build/.ninja_deps b/qwen3_6_scripts/flash_qla_sm70/build/.ninja_deps new file mode 100644 index 00000000..e5675ec1 Binary files /dev/null and b/qwen3_6_scripts/flash_qla_sm70/build/.ninja_deps differ diff --git a/qwen3_6_scripts/flash_qla_sm70/build/.ninja_log b/qwen3_6_scripts/flash_qla_sm70/build/.ninja_log new file mode 100644 index 00000000..ee3aa016 --- /dev/null +++ b/qwen3_6_scripts/flash_qla_sm70/build/.ninja_log @@ -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 diff --git a/qwen3_6_scripts/flash_qla_sm70/build/build.ninja b/qwen3_6_scripts/flash_qla_sm70/build/build.ninja new file mode 100644 index 00000000..e9e140ce --- /dev/null +++ b/qwen3_6_scripts/flash_qla_sm70/build/build.ninja @@ -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 diff --git a/qwen3_6_scripts/flash_qla_sm70/build/flash_qla_sm70_gdn_strided.so b/qwen3_6_scripts/flash_qla_sm70/build/flash_qla_sm70_gdn_strided.so new file mode 100755 index 00000000..a042bfe9 Binary files /dev/null and b/qwen3_6_scripts/flash_qla_sm70/build/flash_qla_sm70_gdn_strided.so differ diff --git a/qwen3_6_scripts/flash_qla_sm70/build/gdn_forward.cuda.o b/qwen3_6_scripts/flash_qla_sm70/build/gdn_forward.cuda.o new file mode 100644 index 00000000..daecc042 Binary files /dev/null and b/qwen3_6_scripts/flash_qla_sm70/build/gdn_forward.cuda.o differ diff --git a/qwen3_6_scripts/prebuilt/corex-3.2.3-ivcore10/SHA256SUMS b/qwen3_6_scripts/prebuilt/corex-3.2.3-ivcore10/SHA256SUMS index f2915aa5..7360f63e 100644 --- a/qwen3_6_scripts/prebuilt/corex-3.2.3-ivcore10/SHA256SUMS +++ b/qwen3_6_scripts/prebuilt/corex-3.2.3-ivcore10/SHA256SUMS @@ -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 diff --git a/qwen3_6_scripts/prebuilt/corex-3.2.3-ivcore10/corex_moe_topk_softmax.so b/qwen3_6_scripts/prebuilt/corex-3.2.3-ivcore10/corex_moe_topk_softmax.so new file mode 100755 index 00000000..f88e0e5c Binary files /dev/null and b/qwen3_6_scripts/prebuilt/corex-3.2.3-ivcore10/corex_moe_topk_softmax.so differ diff --git a/vllm/corex_moe.py b/vllm/corex_moe.py new file mode 100644 index 00000000..3de7cf28 --- /dev/null +++ b/vllm/corex_moe.py @@ -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 diff --git a/vllm/corex_so_loader.py b/vllm/corex_so_loader.py new file mode 100644 index 00000000..7bcc78ae --- /dev/null +++ b/vllm/corex_so_loader.py @@ -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() diff --git a/vllm/ix_unified.py b/vllm/ix_unified.py new file mode 100644 index 00000000..50f87a61 --- /dev/null +++ b/vllm/ix_unified.py @@ -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() diff --git a/vllm/ix_unified_bridge.so b/vllm/ix_unified_bridge.so new file mode 100755 index 00000000..1612b635 Binary files /dev/null and b/vllm/ix_unified_bridge.so differ diff --git a/vllm/moe_fused_dispatch.py b/vllm/moe_fused_dispatch.py new file mode 100644 index 00000000..7030a5b0 --- /dev/null +++ b/vllm/moe_fused_dispatch.py @@ -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)