Compare commits

...
Author SHA1 Message Date
Jee Jee Li f67c1ad263 FIX
Signed-off-by: Jee Jee Li <jeejeelee@inferact.ai>
2026-07-04 13:45:35 +00:00
Jee Jee Li 46ef0dfe19 Merge remote-tracking branch 'origin/main' into glm5--fp32-router 2026-07-03 15:40:18 +00:00
Jee Jee Li 8c6f498c22 Add bench config
Signed-off-by: Jee Jee Li <jeejeelee@inferact.ai>
2026-07-03 11:25:06 +00:00
Jee Jee Li f03d7b96f4 Init
Signed-off-by: Jee Jee Li <jeejeelee@inferact.ai>
2026-07-03 10:27:36 +00:00
7 changed files with 47 additions and 158 deletions
+6 -3
View File
@@ -18,9 +18,10 @@ DSV3_SUPPORTED_HIDDEN_SIZES = [7168]
GPT_OSS_SUPPORTED_NUM_EXPERTS = [32, 128]
GPT_OSS_SUPPORTED_HIDDEN_SIZES = [2880]
# Dimensions supported by the fp32 specialized kernel (MiniMax-M2)
# Dimensions supported by the fp32 specialized kernel
# (3072, 256) -> MiniMax-M2/M2.5, (6144, 256) -> GLM-5
FP32_SUPPORTED_NUM_EXPERTS = [256]
FP32_SUPPORTED_HIDDEN_SIZES = [3072]
FP32_SUPPORTED_HIDDEN_SIZES = [3072, 6144]
FP32_MAX_TOKENS = 32
@@ -33,6 +34,7 @@ def get_model_params(config):
"DeepseekV2ForCausalLM",
"DeepseekV3ForCausalLM",
"DeepseekV32ForCausalLM",
"GlmMoeDsaForCausalLM",
):
num_experts = config.n_routed_experts
hidden_size = config.hidden_size
@@ -107,7 +109,8 @@ def get_benchmark(model, max_batch_size, trust_remote_code):
if provider == "torch":
def runner():
if allow_fp32_router_gemm:
if is_fp32_router_model:
# fp32 weights: reference always computes in fp32
F.linear(mat_a.float(), mat_b)
elif has_bias:
F.linear(mat_a, mat_b, bias)
+4 -1
View File
@@ -176,7 +176,8 @@ void invokeFp32RouterGemm(float* output, InputT const* mat_a,
// ---------------------------------------------------------------------------
// Explicit instantiations: M=1..32, for both input types, for the supported
// (E, H) pairs: (256, 3072) [MiniMax-M2/M2.5] and (128, 6144) [MiniMax-M3].
// (E, H) pairs: (256, 3072) [MiniMax-M2/M2.5], (128, 6144) [MiniMax-M3]
// and (256, 6144) [GLM-5].
// ---------------------------------------------------------------------------
#define INSTANTIATE(T, M, E, H) \
@@ -221,6 +222,8 @@ INSTANTIATE_ALL(float, 256, 3072)
INSTANTIATE_ALL(__nv_bfloat16, 256, 3072)
INSTANTIATE_ALL(float, 128, 6144)
INSTANTIATE_ALL(__nv_bfloat16, 128, 6144)
INSTANTIATE_ALL(float, 256, 6144)
INSTANTIATE_ALL(__nv_bfloat16, 256, 6144)
#undef INSTANTIATE_ALL
#undef INSTANTIATE
@@ -25,10 +25,12 @@ inline int getSMVersion() {
static constexpr int FP32_MAX_TOKENS = 32;
// Supported (hidden_dim, num_experts) pairs (must match the instantiations in
// fp32_router_gemm.cu): (3072, 256) for MiniMax-M2/M2.5, (6144, 128) for M3.
// fp32_router_gemm.cu): (3072, 256) for MiniMax-M2/M2.5, (6144, 128) for M3,
// (6144, 256) for GLM-5.
static inline bool fp32_router_gemm_supported(int hidden_dim, int num_experts) {
return (hidden_dim == 3072 && num_experts == 256) ||
(hidden_dim == 6144 && num_experts == 128);
(hidden_dim == 6144 && num_experts == 128) ||
(hidden_dim == 6144 && num_experts == 256);
}
// Forward declarations — 4 template params must match fp32_router_gemm.cu
@@ -77,6 +79,9 @@ void dispatchFp32RouterGemm(int num_experts, int hidden_dim, int num_tokens,
} else if (num_experts == 128 && hidden_dim == 6144) {
Fp32LoopUnroller<InputT, 128, 6144, 1, FP32_MAX_TOKENS>::unroll(
num_tokens, output, mat_a, mat_b, stream);
} else if (num_experts == 256 && hidden_dim == 6144) {
Fp32LoopUnroller<InputT, 256, 6144, 1, FP32_MAX_TOKENS>::unroll(
num_tokens, output, mat_a, mat_b, stream);
} else {
throw std::invalid_argument(
"fp32_router_gemm: unsupported (hidden_dim, num_experts) pair");
@@ -111,7 +116,7 @@ void fp32_router_gemm(
STD_TORCH_CHECK(
fp32_router_gemm_supported(hidden_dim, num_experts),
"fp32_router_gemm: supported (hidden_dim, num_experts) pairs are "
"(3072, 256) and (6144, 128)");
"(3072, 256), (6144, 128) and (6144, 256)");
STD_TORCH_CHECK(num_tokens <= FP32_MAX_TOKENS,
"fp32_router_gemm: num_tokens must be in [0, 32]");
STD_TORCH_CHECK(
@@ -286,52 +286,3 @@ template void invokeRouterGemmBf16Output<__nv_bfloat16, 15, 384, 7168>(
template void invokeRouterGemmBf16Output<__nv_bfloat16, 16, 384, 7168>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
// Template instantiations for GLM-5 (DEFAULT_NUM_EXPERTS, hidden_dim=6144)
template void invokeRouterGemmBf16Output<__nv_bfloat16, 1, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 2, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 3, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 4, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 5, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 6, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 7, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 8, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 9, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 10, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 11, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 12, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 13, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 14, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 15, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmBf16Output<__nv_bfloat16, 16, 256, 6144>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
@@ -41,7 +41,6 @@ inline int getSMVersion() {
static constexpr int DEFAULT_NUM_EXPERTS = 256;
static constexpr int KIMI_K2_NUM_EXPERTS = 384;
static constexpr int DEFAULT_HIDDEN_DIM = 7168;
static constexpr int GLM_5_HIDDEN_DIM = 6144;
template <typename T, int kNumTokens, int kNumExperts, int kHiddenDim>
void invokeRouterGemmFloatOutput(float* output, T const* mat_a, T const* mat_b,
@@ -122,21 +121,14 @@ void dsv3_router_gemm(
STD_TORCH_CHECK(mat_a.size(1) == mat_b.size(1),
"mat_a and mat_b must have the same hidden_dim");
STD_TORCH_CHECK(
hidden_dim == DEFAULT_HIDDEN_DIM || hidden_dim == GLM_5_HIDDEN_DIM,
"Expected hidden_dim=", DEFAULT_HIDDEN_DIM,
" or hidden_dim=", GLM_5_HIDDEN_DIM, ", but got hidden_dim=", hidden_dim);
STD_TORCH_CHECK(hidden_dim == DEFAULT_HIDDEN_DIM,
"Expected hidden_dim=", DEFAULT_HIDDEN_DIM,
", but got hidden_dim=", hidden_dim);
STD_TORCH_CHECK(
num_experts == DEFAULT_NUM_EXPERTS || num_experts == KIMI_K2_NUM_EXPERTS,
"Expected num_experts=", DEFAULT_NUM_EXPERTS,
" or num_experts=", KIMI_K2_NUM_EXPERTS,
", but got num_experts=", num_experts);
// KIMI_K2_NUM_EXPERTS is only instantiated for the default hidden_dim.
STD_TORCH_CHECK(
hidden_dim == DEFAULT_HIDDEN_DIM || num_experts == DEFAULT_NUM_EXPERTS,
"hidden_dim=", GLM_5_HIDDEN_DIM,
" only supports num_experts=", DEFAULT_NUM_EXPERTS,
", but got num_experts=", num_experts);
STD_TORCH_CHECK(num_tokens >= 1 && num_tokens <= 16,
"currently num_tokens must be less than or equal to 16 for "
"router_gemm");
@@ -165,42 +157,30 @@ void dsv3_router_gemm(
if (output.scalar_type() == torch::headeronly::ScalarType::Float) {
float* out_ptr = reinterpret_cast<float*>(output.mutable_data_ptr());
if (hidden_dim == DEFAULT_HIDDEN_DIM) {
if (num_experts == DEFAULT_NUM_EXPERTS) {
LoopUnroller<1, 16, DEFAULT_NUM_EXPERTS,
DEFAULT_HIDDEN_DIM>::unroll_float_output(num_tokens,
out_ptr, a_ptr,
b_ptr, stream);
} else {
LoopUnroller<1, 16, KIMI_K2_NUM_EXPERTS,
DEFAULT_HIDDEN_DIM>::unroll_float_output(num_tokens,
out_ptr, a_ptr,
b_ptr, stream);
}
} else { // GLM_5_HIDDEN_DIM
if (num_experts == DEFAULT_NUM_EXPERTS) {
LoopUnroller<1, 16, DEFAULT_NUM_EXPERTS,
GLM_5_HIDDEN_DIM>::unroll_float_output(num_tokens, out_ptr,
a_ptr, b_ptr, stream);
DEFAULT_HIDDEN_DIM>::unroll_float_output(num_tokens, out_ptr,
a_ptr, b_ptr,
stream);
} else {
LoopUnroller<1, 16, KIMI_K2_NUM_EXPERTS,
DEFAULT_HIDDEN_DIM>::unroll_float_output(num_tokens, out_ptr,
a_ptr, b_ptr,
stream);
}
} else if (output.scalar_type() == torch::headeronly::ScalarType::BFloat16) {
__nv_bfloat16* out_ptr =
reinterpret_cast<__nv_bfloat16*>(output.mutable_data_ptr());
if (hidden_dim == DEFAULT_HIDDEN_DIM) {
if (num_experts == DEFAULT_NUM_EXPERTS) {
LoopUnroller<1, 16, DEFAULT_NUM_EXPERTS,
DEFAULT_HIDDEN_DIM>::unroll_bf16_output(num_tokens,
out_ptr, a_ptr,
b_ptr, stream);
} else {
LoopUnroller<1, 16, KIMI_K2_NUM_EXPERTS,
DEFAULT_HIDDEN_DIM>::unroll_bf16_output(num_tokens,
out_ptr, a_ptr,
b_ptr, stream);
}
} else { // GLM_5_HIDDEN_DIM
if (num_experts == DEFAULT_NUM_EXPERTS) {
LoopUnroller<1, 16, DEFAULT_NUM_EXPERTS,
GLM_5_HIDDEN_DIM>::unroll_bf16_output(num_tokens, out_ptr,
a_ptr, b_ptr, stream);
DEFAULT_HIDDEN_DIM>::unroll_bf16_output(num_tokens, out_ptr,
a_ptr, b_ptr,
stream);
} else {
LoopUnroller<1, 16, KIMI_K2_NUM_EXPERTS,
DEFAULT_HIDDEN_DIM>::unroll_bf16_output(num_tokens, out_ptr,
a_ptr, b_ptr,
stream);
}
}
}
@@ -286,52 +286,3 @@ template void invokeRouterGemmFloatOutput<__nv_bfloat16, 15, 384, 7168>(
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 16, 384, 7168>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
// Template instantiations for GLM-5 (DEFAULT_NUM_EXPERTS, hidden_dim=6144)
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 1, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 2, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 3, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 4, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 5, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 6, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 7, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 8, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 9, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 10, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 11, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 12, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 13, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 14, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 15, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
template void invokeRouterGemmFloatOutput<__nv_bfloat16, 16, 256, 6144>(
float*, __nv_bfloat16 const*, __nv_bfloat16 const*, cudaStream_t);
@@ -14,9 +14,8 @@ from vllm.utils.torch_utils import direct_register_custom_op
class GateLinear(ReplicatedLinear):
"""MoE gate linear layer with multi-tier GEMM dispatch:
1. DSV3 specialized kernel (SM90+, M<=16, H=7168 E=256/384, H=6144 E=256)
2. fp32 specialized kernel (SM90+, bf16/fp32 in, fp32 out,
M<=32, H=3072, E=256)
1. DSV3 specialized kernel (SM90+, M<=16, H=7168 E=256/384)
2. fp32 specialized kernel (SM90+, bf16/fp32 in, fp32 out, M<=32)
3. cuBLAS bf16×bf16→fp32 (SM90+ + bf16 weight + fp32 out_dtype)
4. F.linear via ReplicatedLinear (ultimate fallback)
@@ -27,16 +26,14 @@ class GateLinear(ReplicatedLinear):
# Dimensions supported by the DSV3 specialized kernel.
# Valid (hidden_size, num_experts) combinations:
# (7168, 256) -> DeepSeek-V3, (7168, 384) -> Kimi-K2,
# (6144, 256) -> GLM-5
# (7168, 256) -> DeepSeek-V3, (7168, 384) -> Kimi-K2
DSV3_SUPPORTED_NUM_EXPERTS = [256, 384]
DSV3_SUPPORTED_HIDDEN_SIZES = [7168, 6144]
# num_experts=384 is only instantiated for hidden_size=7168.
DSV3_UNSUPPORTED_SHAPES = {(6144, 384)}
DSV3_SUPPORTED_HIDDEN_SIZES = [7168]
# (hidden_size, num_experts) pairs with an instantiated fp32 kernel:
# (3072, 256) -> MiniMax-M2/M2.5, (6144, 128) -> MiniMax-M3
FP32_SUPPORTED_SHAPES = {(3072, 256), (6144, 128)}
# (3072, 256) -> MiniMax-M2/M2.5, (6144, 128) -> MiniMax-M3,
# (6144, 256) -> GLM-5
FP32_SUPPORTED_SHAPES = {(3072, 256), (6144, 128), (6144, 256)}
FP32_MAX_TOKENS = 32
def __init__(
@@ -77,7 +74,6 @@ class GateLinear(ReplicatedLinear):
and self.weight.dtype == torch.bfloat16
and output_size in self.DSV3_SUPPORTED_NUM_EXPERTS
and input_size in self.DSV3_SUPPORTED_HIDDEN_SIZES
and (input_size, output_size) not in self.DSV3_UNSUPPORTED_SHAPES
)
# See https://github.com/vllm-project/vllm/pull/44217
# for more details.
@@ -128,7 +124,7 @@ class GateLinear(ReplicatedLinear):
)
return output, None
# Tier 2: fp32 specialized kernel (H=3072, E=256, M<=32)
# Tier 2: fp32 specialized kernel (M<=32)
# Dispatch is wrapped in a custom op so that torch.compile/CUDA-graph
# capture does not freeze the runtime num_tokens branch.
if self.allow_fp32_router_gemm and x.dtype in (