Compare commits
4
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f67c1ad263 | ||
|
|
46ef0dfe19 | ||
|
|
8c6f498c22 | ||
|
|
f03d7b96f4 |
@@ -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)
|
||||
|
||||
@@ -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 (
|
||||
|
||||
Reference in New Issue
Block a user