fix: wire MoE topk via ixformer C++ bridge + disable broken flash_qla GDN

Two call chain breaks fixed:

1. MoE routing (2304 calls/token):
   BEFORE: torch.softmax + torch.topk (3 Python GPU ops, no ixformer)
   AFTER:  ix_bridge.py → ix_moe_bridge.cpp → ixformer::infer::topk_softmax()
   Source: upstream_ref/xllm/xllm/core/kernels/ilu/fused_moe.cpp line 46
   The C++ API exists in base image SDK (ixformer.h declares it),
   only the Python binding (ixformer.functions) was missing.

2. GDN prefill (4 layers, 99.98% NaN):
   BEFORE: flash_qla SM70 kernel → abs mean=inf → nan_to_num → zeros
   AFTER:  skip flash_qla, use _pytorch_forward directly
   Source: upstream_ref/xllm qwen3_gated_delta_net_base.cpp uses
   identical PyTorch chunked logic (no flash_qla).
   Sub168 (working build) never deployed flash_qla either.

Files:
- ex_engine/csrc/ix_moe_bridge.cpp: torch C++ extension calling ixformer C++ API
- ex_engine/python/ix_bridge.py: JIT-compile loader with PyTorch fallback
- qwen3_5.py: import ix_bridge for MoE, disable flash_qla for GDN
- patch_ops.sh: deploy ix_bridge .cpp + .py into vllm model dir
This commit is contained in:
EX Engine
2026-08-10 03:00:24 +00:00
parent 8e6adf20e6
commit d21b2505bb
4 changed files with 231 additions and 30 deletions

View File

@@ -0,0 +1,95 @@
// ix_moe_bridge.cpp — Bridge to ixformer C++ topk_softmax
//
// Problem: ixformer Python (ixformer.functions) lacks vllm_moe_topk_softmax
// Solution: Call ixformer::infer::topk_softmax() directly via C++ torch extension
//
// Source: upstream_ref/xllm/xllm/core/kernels/ilu/ixformer.h declares:
// void topk_softmax(torch::Tensor&, torch::Tensor&, torch::Tensor&,
// torch::Tensor&, bool);
// Source: upstream_ref/xllm/xllm/core/kernels/ilu/fused_moe.cpp shows usage:
// infer::topk_softmax(reduce_weight, topk_indices, token_expert_indices, input_, false);
#include <torch/extension.h>
// Forward-declare ixformer C++ API (from ixformer.h in base image SDK)
namespace ixformer {
namespace infer {
void topk_softmax(torch::Tensor& topk_weights,
torch::Tensor& topk_indices,
torch::Tensor& token_expert_indices,
torch::Tensor& gating_output,
bool renormalize);
void moe_compute_token_index_api(
torch::Tensor& topk_ids,
torch::Tensor& src_dst,
torch::Tensor& dst_src,
torch::Tensor& expert_sizes_gpu,
const c10::optional<torch::Tensor>& expert_mask,
const c10::optional<torch::Tensor>& expert_sizes_cpu,
const c10::optional<torch::Tensor>& expand_tokens_gpu,
int64_t start_expert_id,
int64_t end_expert_id,
int64_t num_experts);
void moe_expand_input(torch::Tensor outputs,
torch::Tensor inputs,
torch::Tensor dst_to_src,
const c10::optional<torch::Tensor>& src_to_dst,
int64_t dst_tokens,
int64_t expand_factor);
void moe_w16a16_group_gemm(torch::Tensor output,
torch::Tensor inputs,
torch::Tensor weights,
torch::Tensor tokens_per_experts,
const c10::optional<torch::Tensor>& dst_to_src,
const c10::optional<torch::Tensor>& bias,
std::string format,
int64_t persistent,
int64_t output_n);
void moe_output_reduce_sum(torch::Tensor outputs,
torch::Tensor inputs,
const c10::optional<torch::Tensor>& mul_weight,
const c10::optional<torch::Tensor>& mask,
const c10::optional<torch::Tensor>& extra_residual,
double scaling_factor);
void silu_and_mul(torch::Tensor& input, torch::Tensor& output);
} // namespace infer
} // namespace ixformer
// Python-callable wrappers
std::tuple<torch::Tensor, torch::Tensor> ix_moe_topk_softmax(
torch::Tensor gating_output, // (num_tokens, num_experts) float32
int64_t topk,
bool renormalize) {
auto input = gating_output.to(torch::kFloat32).contiguous();
int64_t num_tokens = input.size(0);
auto topk_weights = torch::empty({num_tokens, topk},
torch::dtype(torch::kFloat32).device(input.device()));
auto topk_indices = torch::empty({num_tokens, topk},
torch::dtype(torch::kInt32).device(input.device()));
auto token_expert_indices = torch::empty({num_tokens, topk},
torch::dtype(torch::kInt32).device(input.device()));
ixformer::infer::topk_softmax(
topk_weights, topk_indices, token_expert_indices, input, renormalize);
// Renormalize if not done by kernel (match xllm behavior)
if (!renormalize) {
auto row_sum = topk_weights.sum(-1, /*keepdim=*/true);
topk_weights = topk_weights / row_sum;
}
return std::make_tuple(topk_weights, topk_indices);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("topk_softmax", &ix_moe_topk_softmax,
"Fused topk+softmax via ixformer C++ API (bypasses missing Python binding)",
py::arg("gating_output"), py::arg("topk"), py::arg("renormalize") = true);
}

View File

@@ -0,0 +1,77 @@
"""
ix_bridge.py — Load ix_moe_bridge C++ extension at runtime.
Calls ixformer::infer::topk_softmax() via C++ torch extension,
bypassing the missing Python binding in ixformer.functions.
Build: JIT-compiled on first import via torch.utils.cpp_extension.load()
(same mechanism as flash_qla_sm70 GDN kernel — proven to work on BI-V100)
"""
import os
import logging
import torch
logger = logging.getLogger("ex_engine.ix_bridge")
_ix_bridge = None
_ix_bridge_available = False
def _load_bridge():
"""JIT-compile and load ix_moe_bridge.so"""
global _ix_bridge, _ix_bridge_available
if _ix_bridge is not None:
return _ix_bridge_available
csrc_dir = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "csrc")
cpp_file = os.path.join(csrc_dir, "ix_moe_bridge.cpp")
if not os.path.exists(cpp_file):
# Try deployed path (inside vllm model dir)
alt_dir = os.path.dirname(os.path.abspath(__file__))
cpp_file = os.path.join(alt_dir, "ix_moe_bridge.cpp")
if not os.path.exists(cpp_file):
logger.warning("ix_moe_bridge.cpp not found at %s", cpp_file)
_ix_bridge_available = False
return False
try:
from torch.utils.cpp_extension import load
logger.info("JIT-compiling ix_moe_bridge.cpp ...")
_ix_bridge = load(
name="ix_moe_bridge",
sources=[cpp_file],
extra_cflags=["-O2"],
verbose=False,
)
_ix_bridge_available = True
logger.info("ix_moe_bridge loaded successfully: %s", dir(_ix_bridge))
return True
except Exception as e:
logger.warning("ix_moe_bridge JIT compile failed: %s", e)
_ix_bridge_available = False
return False
def topk_softmax(gating_output: torch.Tensor, topk: int, renormalize: bool = True):
"""
Fused topk+softmax via ixformer C++ API.
Args:
gating_output: (num_tokens, num_experts) router logits
topk: number of experts to select
renormalize: whether to renormalize weights
Returns:
(topk_weights, topk_indices) — both (num_tokens, topk)
"""
if not _ix_bridge_available:
if not _load_bridge():
# Fallback to pure PyTorch
probs = torch.softmax(gating_output.float(), dim=-1)
topk_w, topk_ids = torch.topk(probs, topk, dim=-1)
if renormalize:
topk_w = topk_w / topk_w.sum(dim=-1, keepdim=True)
return topk_w, topk_ids.to(torch.int32)
return _ix_bridge.topk_softmax(gating_output, topk, renormalize)