feat(EX): wire xllm CUB topk_softmax kernel into MoE routing
Upstream: xllm/kernels/cuda/moe/moe_topk_softmax_kernels.cuh (Apache 2.0)
Adapted: CHECK→TORCH_CHECK, include path fix, cuda/functional guard, pybind11
Call chain now:
qwen3_5.py:_pure_pytorch_experts()
→ _ex_moe_topk_softmax (fused CUB kernel, 1 launch)
→ fallback: torch.softmax + torch.topk (3 launches)
Files:
ex_engine/csrc/moe/moe_topk_softmax_kernels.cuh — xllm kernel (adapted)
ex_engine/csrc/moe/device_utils.cuh — xllm device utils
ex_engine/csrc/moe/moe_topk_softmax_ext.cu — pybind11 wrapper
ex_engine/python/moe_topk.py — JIT loader (same pattern as flash_qla_sm70)
qwen3_5.py — import + use in _pure_pytorch_experts()
patch_ops.sh — deploy kernel sources for JIT
This commit is contained in:
@@ -127,6 +127,21 @@ try:
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# EX Engine: fused MoE topk_softmax CUDA kernel (xllm CUB-based)
|
||||
_ex_moe_topk_softmax = None
|
||||
_ex_moe_topk_available = False
|
||||
try:
|
||||
from ex_engine.python.moe_topk import moe_topk_softmax as _ex_moe_topk_softmax
|
||||
_ex_moe_topk_available = True
|
||||
logger.info("EX Engine MoE topk_softmax kernel available")
|
||||
except ImportError:
|
||||
try:
|
||||
from vllm.model_executor.models.ex_engine.moe_topk import moe_topk_softmax as _ex_moe_topk_softmax
|
||||
_ex_moe_topk_available = True
|
||||
logger.info("EX Engine MoE topk_softmax kernel available (vllm path)")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ixformer-accelerated ops (drop-in replacements for torch ops)
|
||||
@@ -1065,9 +1080,21 @@ class Qwen3_5MoeSparseBlock(nn.Module):
|
||||
Output is partial (pre-all-reduce), same contract as FusedMoE
|
||||
with reduce_results=False.
|
||||
"""
|
||||
# Routing: fused topk+softmax via ixformer C++ bridge (if available)
|
||||
# Falls back to PyTorch softmax → topk → renormalize
|
||||
if _ix_bridge_available:
|
||||
# Routing: fused topk+softmax dispatch chain
|
||||
# Tier 1: EX Engine CUB kernel → Tier 2: ix_bridge → Tier 3: PyTorch
|
||||
if _ex_moe_topk_available:
|
||||
T_tok = router_logits.shape[0]
|
||||
topk_weights = torch.empty(T_tok, self.top_k, dtype=torch.float32,
|
||||
device=router_logits.device)
|
||||
topk_ids = torch.empty(T_tok, self.top_k, dtype=torch.int32,
|
||||
device=router_logits.device)
|
||||
token_expert_indices = torch.empty(T_tok, self.top_k, dtype=torch.int32,
|
||||
device=router_logits.device)
|
||||
_ex_moe_topk_softmax(topk_weights, topk_ids, token_expert_indices,
|
||||
router_logits.float(), True)
|
||||
topk_ids = topk_ids.to(torch.long)
|
||||
topk_weights = topk_weights.to(hidden_states.dtype)
|
||||
elif _ix_bridge_available:
|
||||
topk_weights, topk_ids = _ix_topk_softmax(
|
||||
router_logits, self.top_k, renormalize=True)
|
||||
topk_weights = topk_weights.to(hidden_states.dtype)
|
||||
|
||||
Reference in New Issue
Block a user