Richael
← 返回博客

2026-08-16

RMSNorm:前向、反向,以及它为什么重要

第五个 kernel:RMSNorm,前向加反向。动机来自这条线的论文那一半——最近在读 Kimi K3 的报告,RMSNorm 反复作为报告里比较大的架构收益出现,就紧挨着注意力残差。感觉与其只是读到,不如动手写一遍。

前向:正确(最大差值 0.00195,属于正常的 fp16 容差),约 230 GB/s,对比未融合的 PyTorch 基线约 20 GB/s。和几周前 softmax 那次是同一个融合故事。

反向是更难的一半——得真的把梯度推出来(dx = rrms * (dxnorm - x_norm * mean(dxnorm * x_norm))),而且权重梯度需要跨每一行求和,不只是行内。这里用了 atomic_add,能跑通(dx 差值 0.0078,dweight 差值 0.0625,考虑到 dweight 的数值尺度更大,两个都没问题),但它的吞吐曲线是锯齿状的,不像前向那条那么干净——是真实的原子操作争用,不是 bug。不管怎么说,都还是比 PyTorch 快约 10 倍。

下一步:要么用一个正经的两阶段归约把原子操作修掉,要么直接进第 6 周。

折线图,比较 Triton 与 PyTorch 在 1024 到 15872 列宽下 RMSNorm 反向的吞吐(GB/s)。Triton 陡升后落在大致 108–150 GB/s 之间的锯齿带里,有明显的下陷,PyTorch 一直平在 11 GB/s 附近。