1. DeepGEMM与Hopper架构的MoE实现解析
DeepGEMM是deepseek团队开源的高性能GEMM(通用矩阵乘法)算子库,特别针对NVIDIA Hopper架构进行了优化。该库不仅支持常规的GEMM运算,还针对混合专家模型(MoE)中的两种典型场景进行了专门优化:group contiguous模式用于prefill阶段,group masked模式用于decode阶段。
在Hopper架构上,DeepGEMM充分利用了新一代GPU的几个关键特性:
- 第三代张量内存加速器(TMA)
- 异步warp级矩阵乘法累加(WGMMA)
- 增强的共享内存(SMEM)管理
- 集群范围的同步机制
2. 核心架构设计与实现原理
2.1 基础概念与符号约定
在深入代码前,我们需要明确几个关键概念:
- Tile划分:GEMM运算D = A × B被划分为多个tile进行处理
- 内存层级:
- 逻辑上的A矩阵tile称为tileA
- SMEM中A矩阵的tile称为SA
- 类似地定义tileB/SB和tileC/SC
2.2 基本用法与测试案例
从test_fp8.py测试文件可以看到库的基本用法。test_gemm函数测试常规GEMM场景,通过enumerate_normal遍历各种参数组合:
python复制def test_gemm() -> None:
for kernel_type, m, n, k, major_a, major_b, accumulate, out_dtype in enumerate_normal(torch.float8_e4m3fn):
major_opt = 'N' if major_a.is_k_major() else 'T'
major_opt += 'T' if major_b.is_k_major() else 'N'
out_opt = 'FP32' if out_dtype == torch.float else 'BF16'
acc_opt = f'acc={int(accumulate)}'
kernel_opt = f'1D1D' if kernel_type.is_1d1d() else '1D2D'
use_ue8m0 = get_ue8m0_usage(kernel_type)
disable_ue8m0_cast = not use_ue8m0
recipe = (1, 1, 128) if kernel_type.is_1d1d() and accumulate else None
a, b, c, d, ref_d = generate_normal(m, n, k, major_a, major_b, accumulate, out_dtype, kernel_type, use_ue8m0=use_ue8m0)
func_name = f'fp8_gemm_{major_opt.lower() if test_alias else "nt"}'
getattr(deep_gemm, func_name)(a, b, d, c=c, disable_ue8m0_cast=disable_ue8m0_cast, recipe=recipe)
这段代码展示了几个关键点:
- 支持多种数据类型组合(FP8输入,BF16/FP32输出)
- 支持不同的量化方式(1D1D和1D2D)
- 支持累加模式
- 自动选择最优的kernel配置
2.3 数据生成与量化处理
generate_normal函数负责生成测试数据并进行FP8量化:
python复制def generate_normal(...):
a = torch.randn((m, k), device='cuda', dtype=torch.bfloat16)
b = torch.randn((n, k), device='cuda', dtype=torch.bfloat16)
d = torch.randn((m, n), device='cuda', dtype=out_dtype) * 32 if accumulate else \
torch.empty((m, n), device='cuda', dtype=out_dtype)
c = d if accumulate else None
ref_d = (a.float() @ b.float().t() + (c if accumulate else 0)).to(out_dtype)
a_fp8 = per_token_cast_to_fp8(a, use_ue8m0=use_ue8m0)
b_fp8 = per_token_cast_to_fp8(b, use_ue8m0=use_ue8m0) if kernel_type.is_1d1d() and accumulate \
else per_block_cast_to_fp8(b, use_ue8m0=use_ue8m0)
a_fp8 = a_fp8 if major_a.is_k_major() else (a_fp8[0].T.contiguous().T, a_fp8[1])
b_fp8 = b_fp8 if major_b.is_k_major() else (b_fp8[0].T.contiguous().T, b_fp8[1])
return a_fp8, b_fp8, c, d, ref_d
量化处理是性能优化的关键,DeepGEMM实现了两种量化方式:
- Per-token量化:对A矩阵沿K方向,每个token中连续128个元素进行量化
- Per-block量化:对B矩阵中每个128×128的block进行量化
per_token_cast_to_fp8函数的实现展示了量化细节:
python复制def per_token_cast_to_fp8(x: torch.Tensor, use_ue8m0: bool) -> Tuple[torch.Tensor, torch.Tensor]:
m, n = x.shape
padded_n = align(n, 128)
x_padded = torch.empty((m, padded_n), dtype=x.dtype, device=x.device).fill_(0)
x_padded[:, :n] = x
x_view = x_padded.view(m, -1, 128)
x_amax = x_view.abs().float().amax(dim=2).view(m, -1).clamp(1e-4)
sf = x_amax / 448.0
sf = ceil_to_ue8m0(sf) if use_ue8m0 else sf
return (x_view * (1.0 / sf.unsqueeze(2))).to(torch.float8_e4m3fn).view(m, padded_n)[:, :n].contiguous(), sf
3. 核心计算流程解析
3.1 GEMM函数入口
fp8_gemm_nt是计算的主要入口:
cpp复制static void fp8_gemm_nt(...) {
if (not recipe.has_value())
recipe = get_default_recipe(a.second.scalar_type(), b.second.scalar_type());
DG_HOST_ASSERT(recipe.value() == std::make_tuple(1, 1, 128) or recipe.value() == std::make_tuple(1, 128, 128));
const auto& sfa = layout::transform_sf_into_required_layout(a.second, m, k, recipe.value(), std::nullopt, true, disable_ue8m0_cast);
const auto& sfb = layout::transform_sf_into_required_layout(b.second, n, k, recipe.value(), std::nullopt, false, disable_ue8m0_cast);
const auto& arch_major = device_runtime->get_arch_major();
if (arch_major == 9 and sfa.scalar_type() == torch::kFloat) {
if (std::get<1>(recipe.value()) == 1) {
sm90_fp8_gemm_1d1d(a.first, sfa, b.first, sfb, c, d, m, n, k, major_a, major_b, compiled_dims);
} else {
const auto& major_sfb = get_major_type_ab(sfb);
sm90_fp8_gemm_1d2d(a.first, sfa, b.first, sfb, c, d, m, n, k, major_a, major_b, major_sfb, compiled_dims);
}
...
}
这里有几个关键处理:
- 默认使用{1, 128, 128}的量化layout
- 对scale因子进行布局转换以适应TMA要求
- 根据架构和参数选择不同的kernel实现
3.2 配置自动优化
DeepGEMM会自动选择最优的执行配置:
cpp复制static GemmConfig get_best_config() {
for (const auto& block_m: block_ms) {
for (const auto& block_n: block_ns) {
const int& num_waves = get_num_waves(block_m, block_n);
const auto& last_util = get_last_wave_util(block_m, block_n);
if (not ArchSpec::is_block_size_legal(kernel_type, major_a, major_b, ab_dtype, cd_dtype, m, n, k, block_m, block_n, block_k))
continue;
bool success = false;
if (best_block_m == 0 or best_block_n == 0 or num_waves < best_num_waves) {
success = true;
} else if (num_waves == best_num_waves) {
// 检查最后一个wave的利用率
success = last_util > best_last_util;
if (last_util == best_last_util) {
// Case 1: same `block_m`, smaller `block_n` (wasted)
success |= block_m == best_block_m and block_n < best_block_n;
// Case 2: same `block_n`, smaller `block_m` (wasted)
success |= block_n == best_block_n and block_m < best_block_m;
// Case 3: different for both `block_m` and `block_n`, larger `block_n` is better
success |= block_m != best_block_m and block_n > best_block_n
and block_n <= n and block_m <= m;
}
}
if (success) {
best_block_m = block_m, best_block_n = block_n;
best_num_waves = num_waves, best_last_util = last_util;
}
}
}
}
优化策略包括:
- 优先选择wave数最少的配置(更大的block尺寸)
- wave数相同时选择最后一个wave利用率更高的配置
- 进一步考虑计算资源的浪费情况
3.3 多播(Multicast)优化
Hopper架构支持TMA多播,可以显著减少内存带宽消耗:
cpp复制static GemmConfig get_best_config() {
// 决定TMA多播数量和广播方向
MulticastConfig best_multicast_config = {1, false};
const auto& [is_legal_on_a, is_legal_on_b] = ArchSpec::get_multicast_legality(
gemm_type, num_groups, m, n, best_block_m, best_block_n, num_sms);
const bool is_legal[2] = {is_legal_on_b, is_legal_on_a};
bool order[2] = {false, true};
if (best_block_m > best_block_n)
std::swap(order[0], order[1]);
for (const bool& is_multicast_on_a: order) {
if (m >= 512 and is_legal[static_cast<int>(is_multicast_on_a)]) {
best_multicast_config = {2, is_multicast_on_a};
break;
}
}
}
多播配置策略:
- 默认不启用多播(num_multicast=1)
- 当矩阵尺寸足够大(m≥512)且符合架构限制时启用
- 优先对较大的矩阵维度进行广播以节省带宽
4. 内存管理与同步机制
4.1 共享内存布局
DeepGEMM精心设计了共享内存的布局以最大化利用Hopper的SMEM:
cpp复制__global__ __launch_bounds__(kNumTMAThreads + kNumMathThreads, 1) void
sm90_fp8_gemm_1d2d_impl(...) {
static constexpr uint32_t SMEM_D_SIZE = constexpr_align(BLOCK_M * BLOCK_N * static_cast<uint32_t>(sizeof(__nv_bfloat16)), 1024u);
static constexpr uint32_t SMEM_A_SIZE_PER_STAGE = BLOCK_M * BLOCK_K * sizeof(__nv_fp8_e4m3);
static constexpr uint32_t SMEM_B_SIZE_PER_STAGE = BLOCK_N * BLOCK_K * sizeof(__nv_fp8_e4m3);
static constexpr uint32_t SMEM_SFA_SIZE_PER_STAGE = BLOCK_M * sizeof(float);
static constexpr uint32_t ALIGNED_SMEM_SFA_SIZE_PER_STAGE = constexpr_align(SMEM_SFA_SIZE_PER_STAGE, 128u);
const uint32_t& shape_k_scales = ceil_div(shape_k, BLOCK_K);
const uint32_t& shape_n_sfb = ceil_div(shape_n, BLOCK_K);
const uint32_t& smem_sfb_size = align<uint32_t>(shape_k_scales * (kMustUseUniformedScaleB ? 1 : 2) * sizeof(float), sizeof(Barrier));
const uint32_t num_total_k_blocks = ceil_div(shape_k, BLOCK_K);
}
SMEM中按顺序存储了:
- 输出矩阵D
- 输入矩阵A和B的tile
- A和B的scale因子
- 同步用的barrier
4.2 同步屏障设计
DeepGEMM使用了Hopper的集群事务屏障来实现高效的线程间同步:
cpp复制// 初始化barrier
if (warp_idx == kNumMathThreads / 32 + 1 and cute::elect_one_sync()) {
#pragma unroll
for (uint32_t i = 0; i < kNumStages; ++ i) {
full_barriers[i]->init(1);
empty_barriers[i]->init(kNumTMAMulticast * kNumMathThreads / 32);
}
cutlass::arch::fence_barrier_init();
}
(kNumTMAMulticast > 1) ? cute::cluster_sync() : __syncthreads();
屏障设计特点:
- full barrier初始计数为1(TMA完成后触发)
- empty barrier初始计数与math warp数和多播数相关
- 多播场景使用集群级同步
5. 调度器设计与优化
5.1 调度器核心逻辑
DeepGEMM实现了智能的调度器来优化L2缓存利用率:
cpp复制template <GemmType kGemmType,
uint32_t BLOCK_M, uint32_t BLOCK_N,
uint32_t kNumGroups,
uint32_t kNumMulticast, bool kIsMulticastOnA,
uint32_t kNumSMs,
uint32_t SF_K_ALIGNMENT = 512u,
uint32_t kNum1DBlocksPerGroup = get_num_1d_blocks_per_group<kGemmType, BLOCK_M, BLOCK_N, kNumSMs, kIsMulticastOnA>()>
struct Scheduler {
int current_iter = -1;
uint32_t num_blocks;
uint32_t num_m_blocks;
uint32_t num_n_blocks;
uint32_t num_blocks_in_group;
bool is_peer_cta_alive = true;
}
调度器通过分组计算来优化数据局部性,减少对全局内存的访问。
5.2 分组策略
分组大小通过启发式算法确定:
cpp复制template <GemmType kGemmType, uint32_t BLOCK_M, uint32_t BLOCK_N, uint32_t kNumSMs, bool kIsMulticastOnA>
static constexpr uint32_t get_num_1d_blocks_per_group() {
uint32_t num_best_blocks = 0, min_usage = cute::numeric_limits<uint32_t>::max();
for (const auto& candidate: {8u, 16u}) {
const auto& usage = kIsMulticastOnA ?
candidate * BLOCK_N + constexpr_ceil_div(kNumSMs, candidate) * BLOCK_M:
candidate * BLOCK_M + constexpr_ceil_div(kNumSMs, candidate) * BLOCK_N;
if (usage < min_usage)
min_usage = usage, num_best_blocks = candidate;
}
return num_best_blocks;
}
策略要点:
- 候选分组大小为8或16
- 计算每种分组大小的"usage"指标(所需加载的数据量)
- 选择usage最小的分组方案
6. 线程分工与执行流程
6.1 TMA线程
负责通过TMA加载数据:
cpp复制if (warp_idx >= kNumMathThreads / 32) {
cutlass::arch::warpgroup_reg_dealloc<kNumTMARegisters>();
if (warp_idx == kNumMathThreads / 32 + 2 and cute::elect_one_sync()) {
while (scheduler.get_next_block(m_block_idx, n_block_idx)) {
for (uint32_t k_block_idx = 0; k_block_idx < num_total_k_blocks; advance_pipeline(k_block_idx)) {
empty_barriers[stage_idx]->wait(phase ^ 1);
const bool is_tma_multicast_valid = scheduler.is_tma_multicast_valid(m_block_idx);
const uint32_t num_tma_multicast_a = (kIsTMAMulticastOnA and is_tma_multicast_valid) ? kNumTMAMulticast : 1;
const uint32_t num_tma_multicast_b = (not kIsTMAMulticastOnA and is_tma_multicast_valid) ? kNumTMAMulticast : 1;
tma_copy<BLOCK_K, BLOCK_M, kSwizzleAMode>(&tensor_map_a, &full_barrier,
smem_a[stage_idx], k_idx, scheduler.get_global_idx<kWithGroupOffsetA>(shape_m, BLOCK_M, m_block_idx),
num_tma_multicast_a);
full_barrier.arrive_and_expect_tx(SMEM_A_SIZE_PER_STAGE + SMEM_B_SIZE_PER_STAGE + SMEM_SFA_SIZE_PER_STAGE);
}
}
}
}
TMA线程的关键职责:
- 等待empty barrier确保SMEM可用
- 根据调度决定是否使用多播
- 执行TMA加载操作
- 触发full barrier通知计算线程
6.2 计算线程
负责实际的矩阵乘法计算:
cpp复制else {
cutlass::arch::warpgroup_reg_alloc<kNumMathRegisters>();
const auto math_wg_idx = __shfl_sync(0xffffffff, threadIdx.x / 128, 0);
auto a_desc = make_smem_desc(smem_a[0] + math_wg_idx * WGMMA::M * BLOCK_K, 1);
auto b_desc = make_smem_desc(smem_b[0], 1);
const uint32_t a_desc_lo = __shfl_sync(0xffffffff, a_desc.reg32_[0], 0);
const uint32_t b_desc_lo = __shfl_sync(0xffffffff, b_desc.reg32_[0], 0);
while (scheduler.get_next_block(m_block_idx, n_block_idx)) {
for (uint32_t k_block_idx = 0; k_block_idx < num_total_k_blocks; advance_pipeline(k_block_idx)) {
full_barriers[stage_idx]->wait(phase ^ 1);
// 加载B的scale因子
if (k_block_idx == 0) {
const auto& n_start = n_block_idx * BLOCK_N;
const auto& n_end = min(shape_n, n_start + BLOCK_N);
const auto& k_start = k_block_idx * BLOCK_K;
const auto& k_end = min(shape_k, k_start + BLOCK_K);
load_sfb(smem_sfb, n_start, n_end, k_start, k_end);
}
// 执行WGMMA
wgmma.mma_async.sync.aligned.m64n8k32.f32.e4m3.e4m3(
a_desc_lo, b_desc_lo, d0, d1, d2, d3, scale_D);
empty_barriers[stage_idx]->arrive();
}
}
}
计算线程的关键步骤:
- 准备WGMMA所需的描述符
- 等待full barrier确保数据就绪
- 加载B的scale因子
- 执行异步WGMMA操作
- 触发empty barrier通知TMA线程
7. 性能优化技巧与经验分享
在实际使用DeepGEMM进行开发时,以下几点经验值得注意:
-
量化策略选择:
- Per-token量化适合A矩阵(通常是激活值)
- Per-block量化适合B矩阵(通常是权重)
- 448.0的magic number来自FP8(E4M3)的最大可表示值
-
TMA使用技巧:
- 确保全局内存地址和步长都是16字节对齐的
- 对频繁访问的数据使用prefetch.tensormap预取描述符
- 合理设置swizzle模式以优化bank冲突
-
WGMMA优化:
- 尽量使用更大的tile尺寸以减少wave数量
- 保持矩阵描述符在warp内一致以节省寄存器
- 利用异步执行隐藏内存延迟
-
同步最佳实践:
- 多播场景下,empty barrier的计数需要乘以多播数量
- 使用轻量级的额外同步确保资源安全释放
- 集群同步比全局同步更高效
-
调试技巧:
- 使用CUDA-GDB可以单步调试WGMMA指令
- NSight Compute可以分析TMA和WGMMA的性能
- 通过cutlass::arch::ClusterTransactionBarrier提供的接口可以调试屏障状态
DeepGEMM的这些优化技巧不仅适用于MoE场景,也可以为其他高性能GEMM实现提供参考。特别是在处理不规则矩阵乘法时,其灵活的分组策略和智能调度算法展现了出色的适应性。
