Grouped GEMM
也称为 grouped GEMM、grouped matmul、GroupedGEMM、分组矩阵乘法、ragged matmul
先用大白话
grouped GEMM 在一次启动中执行多个小矩阵乘法,每个专家一个,让 MoE 层不必为每个专家单独付出开销。
技术定义
grouped GEMM 是在一次启动中执行一组行数各不相同的独立矩阵乘法的 kernel,用于 MoE 层中每个专家接收一组不规则路由 token 的场景。
工程细节
MoE 层为每个专家产生不规则的 token 分组,这些 ragged 分组必须重排成矩阵单元能高效消费的形状。TPU 上第二版 grouped matmul 移除了冗余的 tile 计算,按有效行数而非补齐后的最大值确定传输量,对专家权重做三缓冲让下一组在当前组计算时已在途,并把分组元数据的生成融合进 kernel。不规则的 token permute 迁移到了 SparseCore。小 batch 下,专用路径构建 one-hot 矩阵并用普通矩阵乘法完成 token 的 permute 和 unpermute,因为通用 ragged 路径的开销超过了它所整理的计算本身。在 DP attention 与专家并行下,kernel 先 all-gather token 激活值和路由元数据,再用 reduce-scatter 把加权输出送回各自的 attention rank。
为什么重要
任何加速器上的 MoE 效率都归结为 grouped GEMM 对不规则分组大小的容忍度,以及能隐藏多少路由工作。TPU 上多出的约束是 MXU tile 几何,因此专家宽度和 token 数也需要填满 256 宽的 tile。
如何在 InferenceX 中解读
SparseCore permute 重写在 Ironwood 上使 8k1k 吞吐量提高 12%,小 batch one-hot 路径在并发 64 和 128 下分别提高 7.3% 和 5.1%,合并路由 all-gather 每层节省约 80 微秒,在 DeepSeek-V3 的 58 层上每次前向约节省 4.64 毫秒。