diff --git a/qwen3_6_scripts/flash_qla_sm70/build/.ninja_deps b/qwen3_6_scripts/flash_qla_sm70/build/.ninja_deps deleted file mode 100644 index e5675ec1..00000000 Binary files a/qwen3_6_scripts/flash_qla_sm70/build/.ninja_deps and /dev/null differ diff --git a/qwen3_6_scripts/flash_qla_sm70/build/.ninja_log b/qwen3_6_scripts/flash_qla_sm70/build/.ninja_log deleted file mode 100644 index ee3aa016..00000000 --- a/qwen3_6_scripts/flash_qla_sm70/build/.ninja_log +++ /dev/null @@ -1,5 +0,0 @@ -# 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 deleted file mode 100644 index e9e140ce..00000000 --- a/qwen3_6_scripts/flash_qla_sm70/build/build.ninja +++ /dev/null @@ -1,31 +0,0 @@ -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 deleted file mode 100755 index a042bfe9..00000000 Binary files a/qwen3_6_scripts/flash_qla_sm70/build/flash_qla_sm70_gdn_strided.so and /dev/null 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 deleted file mode 100644 index daecc042..00000000 Binary files a/qwen3_6_scripts/flash_qla_sm70/build/gdn_forward.cuda.o and /dev/null differ diff --git a/vllm/corex_moe.py b/vllm/corex_moe.py deleted file mode 100644 index 3de7cf28..00000000 --- a/vllm/corex_moe.py +++ /dev/null @@ -1,233 +0,0 @@ -""" -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 deleted file mode 100644 index 7bcc78ae..00000000 --- a/vllm/corex_so_loader.py +++ /dev/null @@ -1,178 +0,0 @@ -"""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 deleted file mode 100644 index 50f87a61..00000000 --- a/vllm/ix_unified.py +++ /dev/null @@ -1,343 +0,0 @@ -"""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 deleted file mode 100755 index 1612b635..00000000 Binary files a/vllm/ix_unified_bridge.so and /dev/null differ diff --git a/vllm/moe_fused_dispatch.py b/vllm/moe_fused_dispatch.py deleted file mode 100644 index 7030a5b0..00000000 --- a/vllm/moe_fused_dispatch.py +++ /dev/null @@ -1,236 +0,0 @@ -"""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)