[perf] Optimize w4a8 MoE for glm5.2 on H200 (#38220)

Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
This commit is contained in:
Benjamin Truong
2026-09-19 20:34:25 +08:00
committed by GitHub
co-authored by Xiaoyu Zhang
parent 7fac84b639
commit 83e29d6c5a
@@ -1,7 +1,9 @@
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cudaTypedefs.h>
#include <torch/all.h>
#include <string_view>
#include <type_traits>
#include "cutlass/cutlass.h"
@@ -85,6 +87,70 @@ inline void invoke_gemm(
chunk_size)
#define INVOKE_GEMM_WITH_CONFIG(Config) INVOKE_GEMM_WITH_CONFIG_HELPER Config
// H200-tuned GLM-5.2 configurations. Return false for untuned shapes.
bool try_dispatch_w4a8_moe_mm_h200(
torch::Tensor& d_tensors,
torch::Tensor const& a_tensors,
torch::Tensor const& b_tensors,
torch::Tensor const& a_scales,
torch::Tensor const& b_scales,
torch::Tensor const& expert_offsets,
torch::Tensor const& problem_sizes,
torch::Tensor const& a_strides,
torch::Tensor const& b_strides,
torch::Tensor const& d_strides,
torch::Tensor const& s_strides,
int64_t chunk_size,
int64_t topk) {
uint32_t const m = a_tensors.size(0) / topk;
uint32_t const n = d_tensors.size(1);
uint32_t const k = a_tensors.size(1);
if (n == 6144 && k == 256) {
// GLM-5.2 group gemm 2 at TP-8
if (m <= 256) {
INVOKE_GEMM_WITH_CONFIG((SM90_CO<256, 16, 128, 1, 1, 1>));
} else if (m <= 512) {
INVOKE_GEMM_WITH_CONFIG((SM90_CO<256, 32, 128, 2, 1, 1>));
} else if (m <= 4096) {
INVOKE_GEMM_WITH_CONFIG((SM90_CO<256, 64, 128, 1, 1, 1>));
} else {
INVOKE_GEMM_WITH_CONFIG((SM90_CO<128, 128, 128, 1, 1, 1>));
}
} else if (n == 512 && k == 6144) {
// GLM-5.2 group gemm 1 at TP-8
if (m <= 256) {
INVOKE_GEMM_WITH_CONFIG((SM90_CO<128, 16, 512, 2, 1, 1>));
} else if (m <= 2048) {
INVOKE_GEMM_WITH_CONFIG((SM90_CO<128, 32, 512, 1, 1, 1>));
} else {
INVOKE_GEMM_WITH_CONFIG((SM90_CO<128, 64, 512, 1, 1, 1>));
}
} else if (n == 4096 && k == 6144) {
// GLM-5.2 group gemm 1 at EP-8 (no DeepEP)
if (m <= 256) {
INVOKE_GEMM_WITH_CONFIG((SM90_CO<128, 16, 512, 2, 1, 1>));
} else if (m <= 1024) {
INVOKE_GEMM_WITH_CONFIG((SM90_CO<128, 32, 512, 1, 1, 1>));
} else {
INVOKE_GEMM_WITH_CONFIG((SM90_CO<128, 64, 512, 1, 1, 1>));
}
} else if (n == 6144 && k == 2048) {
// GLM-5.2 group gemm 2 at EP-8 (no DeepEP)
if (m <= 256) {
INVOKE_GEMM_WITH_CONFIG((SM90_CO<128, 16, 512, 2, 1, 1>));
} else if (m <= 1024) {
INVOKE_GEMM_WITH_CONFIG((SM90_CO<128, 32, 512, 1, 1, 1>));
} else {
INVOKE_GEMM_WITH_CONFIG((SM90_CO<128, 64, 512, 1, 1, 1>));
}
} else {
return false;
}
return true;
}
// Legacy SM90 configuration table, retained for all devices as the fallback.
void dispatch_w4a8_moe_mm_sm90(
torch::Tensor& d_tensors,
torch::Tensor const& a_tensors,
@@ -197,6 +263,29 @@ void cutlass_w4a8_moe_mm_sm90(
torch::Tensor const& s_strides,
int64_t chunk_size,
int64_t topk) {
const c10::cuda::CUDAGuard device_guard(a_tensors.device());
// Detect H200 devices, including names such as "NVIDIA H200 SXM".
const std::string_view device_name(at::cuda::getCurrentDeviceProperties()->name);
const auto model_pos = device_name.find("H200");
const bool is_h200 = model_pos != std::string_view::npos && (model_pos == 0 || device_name[model_pos - 1] == ' ') &&
(model_pos + 4 == device_name.size() || device_name[model_pos + 4] == ' ');
if (is_h200 && try_dispatch_w4a8_moe_mm_h200(
d_tensors,
a_tensors,
b_tensors,
a_scales,
b_scales,
expert_offsets,
problem_sizes,
a_strides,
b_strides,
d_strides,
s_strides,
chunk_size,
topk)) {
return;
}
dispatch_w4a8_moe_mm_sm90(
d_tensors,
a_tensors,