feat: xllm MoE CUDA kernels — fused_topk + compute_index + combine
3 MoE kernel files adapted for corex: moe_fused_topk.cu: LOG(FATAL)→TORCH_CHECK, +torch/extension.h moe_compute_index.cu: CHECK_LE→TORCH_CHECK, uses cub::BlockScan (corex CUB) moe_combine.cu: fixed duplicate include, +torch/extension.h New pybind binding: xllm_moe_bind.cpp → moe_fused_topk(gating, topk, renormalize, bias, scoring_func) → moe_compute_index(expert_id, num_experts) → moe_combine_result(gemm2, weights, N, topk) AST verification added for all 3 functions
This commit is contained in:
34
ex_engine/xllm_kernels/cuda/bindings/xllm_moe_bind.cpp
Normal file
34
ex_engine/xllm_kernels/cuda/bindings/xllm_moe_bind.cpp
Normal file
@@ -0,0 +1,34 @@
|
||||
// xllm_moe_bind.cpp — pybind11 for MoE CUDA kernels
|
||||
#include <torch/extension.h>
|
||||
#include <optional>
|
||||
#include <tuple>
|
||||
|
||||
namespace xllm::kernel::cuda {
|
||||
std::tuple<torch::Tensor, torch::Tensor> moe_fused_topk(
|
||||
torch::Tensor& gating_output, int64_t topk, bool renormalize,
|
||||
const std::optional<torch::Tensor>& correction_bias,
|
||||
const std::string& scoring_func);
|
||||
|
||||
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> moe_compute_index(
|
||||
const torch::Tensor& expert_id, int64_t num_experts);
|
||||
|
||||
torch::Tensor moe_combine_result(
|
||||
const torch::Tensor& gemm2, const torch::Tensor& reduce_weight,
|
||||
int64_t N, int32_t topk);
|
||||
}
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("moe_fused_topk", &xllm::kernel::cuda::moe_fused_topk,
|
||||
"MoE fused topk (softmax or sigmoid routing)",
|
||||
py::arg("gating_output"), py::arg("topk"),
|
||||
py::arg("renormalize") = true,
|
||||
py::arg("correction_bias") = py::none(),
|
||||
py::arg("scoring_func") = "softmax");
|
||||
m.def("moe_compute_index", &xllm::kernel::cuda::moe_compute_index,
|
||||
"MoE compute permutation index (histogram + prefix_sum + place)",
|
||||
py::arg("expert_id"), py::arg("num_experts"));
|
||||
m.def("moe_combine_result", &xllm::kernel::cuda::moe_combine_result,
|
||||
"MoE combine (reorder + weighted sum)",
|
||||
py::arg("gemm2"), py::arg("reduce_weight"),
|
||||
py::arg("N"), py::arg("topk"));
|
||||
}
|
||||
Reference in New Issue
Block a user