Copied from upstream_ref (NOT rewritten — exact upstream code):
ixformer C++ API (the authoritative header):
include/ixformer.h — ixformer::infer namespace: topk_softmax,
moe_compute_token_index_api, moe_w16a16_group_gemm, moe_expand_input,
moe_output_reduce_sum, silu_and_mul, rms_norm, xllm_paged_attention, etc.
include/ilu_ops_api.h — xllm::kernel::ilu namespace: moe_active_topk,
moe_gen_idx, moe_expand_input, group_gemm, moe_combine_result,
batch_prefill, batch_decode, rms_norm, matmul, act_and_mul, etc.
ILU kernel wrappers (call ixformer::infer directly):
csrc/ilu_kernel_fused_moe.cpp — topk routing + gen_idx + expand + combine
csrc/ilu_kernel_group_gemm.cpp — batched expert GEMM
csrc/ilu_kernel_{activation,norm,rope,matmul,attention}.cpp
ILU layer implementations (full pipeline):
csrc/ilu_layer_fused_moe.{cpp,h} — 797 lines, the complete MoE pipeline
that competitor 168 ran as corex_moe.py
csrc/ilu_layer_attention.{cpp,h} — prefill/decode attention dispatch
CUDA MoE kernels (from xllm + ds_vllm):
csrc/moe/moe_topk_softmax_kernels.cuh — CUB BlockReduce + warp topk
csrc/moe/moe_topk_sigmoid_kernels.cuh — sigmoid scoring variant
csrc/moe/moe_topk.cuh + moe_fused_topk.cu — entry points
csrc/moe/moeTopKFuncs.cuh — TRT-LLM derived vllm-compatible topk
csrc/moe/moe_ops.h + moe_align_sum_kernels.cu — alignment kernels
Common layer headers:
csrc/common_fused_moe{,_base}.h + common_moe_fused_topk.{cpp,h}
72 lines
2.7 KiB
C++
72 lines
2.7 KiB
C++
/* Copyright 2026 The xLLM Authors. All Rights Reserved.
|
|
|
|
Licensed under the Apache License, Version 2.0 (the "License");
|
|
you may not use this file except in compliance with the License.
|
|
You may obtain a copy of the License at
|
|
|
|
https://github.com/jd-opensource/xllm/blob/main/LICENSE
|
|
|
|
Unless required by applicable law or agreed to in writing, software
|
|
distributed under the License is distributed on an "AS IS" BASIS,
|
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
See the License for the specific language governing permissions and
|
|
limitations under the License.
|
|
==============================================================================*/
|
|
|
|
#include "moe_fused_topk.h"
|
|
|
|
#include "kernels/ops_api.h"
|
|
|
|
namespace xllm {
|
|
namespace layer {
|
|
|
|
MoEFusedTopkImpl::MoEFusedTopkImpl(const ModelArgs& model_args,
|
|
const QuantArgs& quant_args,
|
|
const torch::TensorOptions& options)
|
|
: topk_(model_args.num_experts_per_tok()),
|
|
num_expert_group_(model_args.n_group()),
|
|
topk_group_(model_args.topk_group()),
|
|
route_scale_(model_args.routed_scaling_factor()),
|
|
hidden_size_(model_args.hidden_size()),
|
|
renormalize_(model_args.norm_topk_prob()),
|
|
scoring_func_(model_args.scoring_func()) {
|
|
const std::string& topk_method = model_args.topk_method();
|
|
if (topk_method == "noaux_tc") {
|
|
e_score_correction_bias_ = register_parameter(
|
|
"e_score_correction_bias",
|
|
torch::empty({model_args.n_routed_experts()}, options),
|
|
false);
|
|
}
|
|
}
|
|
|
|
// select the experts and return the reduce_weight and expert_id
|
|
std::tuple<torch::Tensor, torch::Tensor> MoEFusedTopkImpl::forward(
|
|
torch::Tensor& router_logits) {
|
|
std::optional<torch::Tensor> e_score_correction_bias = std::nullopt;
|
|
if (e_score_correction_bias_.defined()) {
|
|
e_score_correction_bias = e_score_correction_bias_;
|
|
}
|
|
|
|
xllm::kernel::MoeFusedTopkParams moe_active_topk_params;
|
|
moe_active_topk_params.input = router_logits;
|
|
moe_active_topk_params.topk = topk_;
|
|
moe_active_topk_params.num_expert_group = num_expert_group_;
|
|
moe_active_topk_params.topk_group = topk_group_;
|
|
moe_active_topk_params.normalize = renormalize_;
|
|
moe_active_topk_params.normed_by = "topk_logit";
|
|
moe_active_topk_params.scoring_func = scoring_func_;
|
|
moe_active_topk_params.route_scale = route_scale_;
|
|
moe_active_topk_params.e_score_correction_bias = e_score_correction_bias;
|
|
|
|
return xllm::kernel::moe_active_topk(moe_active_topk_params);
|
|
}
|
|
|
|
void MoEFusedTopkImpl::load_state_dict(const StateDict& state_dict) {
|
|
if (e_score_correction_bias_.defined() &&
|
|
!e_score_correction_bias_is_loaded_) {
|
|
LOAD_WEIGHT(e_score_correction_bias);
|
|
}
|
|
}
|
|
} // namespace layer
|
|
} // namespace xllm
|