Files
project_6/ex_engine/csrc/moe/moe_topk_softmax_ext.cu
EngineX 7839982707 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
2026-08-10 03:10:58 +00:00

56 lines
2.2 KiB
Plaintext

// ex_engine/csrc/moe/moe_topk_softmax_ext.cu
//
// Torch extension wrapper for xllm's topk_gating_softmax kernel.
// Compiles via torch.utils.cpp_extension.load() on BI-V100.
//
// Interface matches vllm's _custom_ops.topk_softmax():
// topk_softmax(topk_weights, topk_ids, token_expert_indices, gating_output)
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
// Include the kernel (adapted from xllm, CHECK→TORCH_CHECK)
#include "moe_topk_softmax_kernels.cuh"
// ---------------------------------------------------------------------------
// Python-facing wrapper: matches _custom_ops.topk_softmax signature exactly
// ---------------------------------------------------------------------------
void topk_softmax_ext(
torch::Tensor& topk_weights, // [num_tokens, topk] float32 output
torch::Tensor& topk_ids, // [num_tokens, topk] int32 output
torch::Tensor& token_expert_indices, // [num_tokens, topk] int32 output
torch::Tensor& gating_output, // [num_tokens, num_experts] input
bool renormalize = false
) {
// Call the xllm kernel
xllm::kernel::cuda::topk_softmax(
topk_weights,
topk_ids,
gating_output,
renormalize,
0.0, // moe_softcapping (unused for Qwen3.5)
std::nullopt // correction_bias
);
// Fill token_expert_indices: flatten assignment
// token_expert_indices[i][j] = i * topk + j
const int num_tokens = topk_weights.size(0);
const int topk = topk_weights.size(1);
auto arange_tokens = torch::arange(num_tokens, topk_ids.options().dtype(torch::kInt32));
auto arange_topk = torch::arange(topk, topk_ids.options().dtype(torch::kInt32));
token_expert_indices.copy_(
arange_tokens.unsqueeze(1) * topk + arange_topk.unsqueeze(0)
);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("topk_softmax", &topk_softmax_ext,
"Fused softmax + topk for MoE routing (xllm CUB kernel)",
py::arg("topk_weights"),
py::arg("topk_ids"),
py::arg("token_expert_indices"),
py::arg("gating_output"),
py::arg("renormalize") = false);
}