⚠️ Alpha内测版本警告:此为早期内部构建版本,尚不完整且可能存在错误,欢迎大家提Issue反馈问题或建议。
Skip to content

第13章 综合实战:Fused RMSNorm

本章导读

RMSNorm 很适合作为 Part 2 的结业题:平方是 Element-Wise,均值是 Reduction,归一化和权重缩放又是 Element-Wise,而高效实现希望把这些步骤融合在一次读写中。

本章不再堆叠新术语,而是完整走一遍“固定语义 → 建立基线 → 融合 → 正确性 → 计时 → profiling → 记录负结果”的工作流。HIP/Triton 边界检查、3 个独立正式进程和逐实现 trace 已在 RX 9070 XT 上完成。

13.1 从 LayerNorm 到 RMSNorm

对一行向量 x 和同长度权重 w,RMSNorm 定义为:

text
mean_square = sum(x[i]²) / N
inv_rms     = 1 / sqrt(mean_square + epsilon)
y[i]        = x[i] * inv_rms * w[i]

与 LayerNorm 相比,它不减去均值,也没有 (x-mean)² 这条路径。它仍然需要跨整行归约,因此不是纯逐元素算子。

手算 x=[3,4]、w=[1,1]、epsilon=0:平方均值是 (9+16)/2=12.5,RMS 是 sqrt(12.5),两个输出都除以同一个行级尺度。这个“先得到一个标量,再广播回整行”的结构,就是本章要优化的数据流。

13.2 先锁定实验契约

首版固定:

  • 输入、权重和输出均为 FP32;
  • 对最后一维逐行归约;
  • epsilon=1e-5
  • CPU/PyTorch reference 使用同一公式;
  • kernel-only GPU event 计时;
  • 必须覆盖列数不是二次幂的尾部;
  • 不发布尚未在实验机验证的性能数字。

建议正确性形状:

RowsCols目的
11最小输入
313mask/尾部
33257跨 wavefront/block 边界
10244096教学主形状

误差阈值不能只看一个固定常数。换成 FP16/BF16 后,应根据累加 dtype、列数和 reference 精度重新制定 atol/rtol

13.3 HIP serial:最短正确基线

code/part2-kernels/chapter13/rmsnorm_hip.hipserial 版本让一个 GPU 线程处理一整行:先顺序累加平方,再顺序写回归一化结果。

bash
./rmsnorm_hip --version serial --rows 1024 --cols 4096 \
  --block 256 --epsilon 1e-5

它的并行度只有“行之间并行”,长行内部没有协作。这通常不是最终性能方案,但有三个教学价值:

  1. 控制流与公式几乎一一对应;
  2. 不需要 LDS 和同步,容易定位数值错误;
  3. 为 block 协作版本提供 GPU 侧基线。

基线首先要可靠,不需要故意把它写得很糟,也不能拿 CPU 时间冒充 GPU kernel 时间。

13.4 HIP block:协作归约并融合写回

block 版本让一个 block 处理一行。每个线程读取多个列并得到局部平方和,然后在 LDS 中做树形归约:

cpp
for (int col = tid; col < cols; col += blockDim.x) {
    float value = input[row * cols + col];
    square_sum += value * value;
}
shared[tid] = square_sum;
// LDS tree reduction
float inverse_rms = rsqrtf(shared[0] / cols + epsilon);

归约结束后,同一批线程直接读取输入、乘 inverse_rms 和 weight、写回输出。中间的 数组和行级统计量都不写到全局内存。

bash
./rmsnorm_hip --version block --rows 1024 --cols 4096 \
  --block 256 --epsilon 1e-5

这里仍有可优化空间:用 wavefront shuffle 收尾、一次加载后在寄存器中复用 x、向量化读写,以及为不同列数选择 block。第一版只改变“行内协作”这一项,便于和 serial 版本比较。

13.5 Triton:一行一个 program

rmsnorm_triton.py 让一个 program 处理一行,列方向扩展到下一个二次幂并用 mask 保护尾部:

python
offsets = tl.arange(0, BLOCK_SIZE)
mask = offsets < cols
values = tl.load(input_ptr + row * cols + offsets, mask=mask, other=0.0)
mean_square = tl.sum(values * values, axis=0) / cols
inverse_rms = tl.rsqrt(mean_square + epsilon)
tl.store(output_ptr + row * cols + offsets,
         values * inverse_rms * weights, mask=mask)

t0t1 使用同一数学 kernel,分别用 4 和 8 个 warps:

bash
python rmsnorm_triton.py --version all --rows 1024 --cols 4096

这是一项受控配置实验,而不是“8 warps 一定优于 4 warps”。列数变大时,单个 program 的向量状态可能增加寄存器或 scratch 使用;应让 profiler 和实测时间回答。

13.6 融合到底省掉了什么

假设分步实现先写 square=x²,再归约得到 mean_square,最后启动新 kernel 做 normalize:

text
x -> square buffer -> row statistic -> reread x -> y

融合实现的数据流是:

text
x -> local square sum -> block/program reduction -> y
                   weight ----------------------^

它省掉了 square buffer 的分配、写入和读取,也减少 dispatch。但 x 是否只从显存读取一次,要看具体实现和编译器生成结果:HIP 第一版在归约后会再次从 input 读取 x;Triton 源码看似只 load 一次,是否发生 spill/重载仍应通过 profiling 判断。不要从源码表象直接推出物理流量。

13.7 一次完整的运行与检查

在 Part 2 环境中执行:

bash
cd code/part2-kernels
bash chapter13/run_all.sh

脚本会:

  1. 激活本篇 .venv 和 ROCm 工具链;
  2. 编译 HIP;
  3. 3×13 检查尾部正确性;
  4. 1024×4096 跑教学主形状;
  5. HIP/Triton 每个版本都打印 RESULT 行。

首版输出中的时间适合判断脚本是否工作,但正式比较仍应重复独立进程、记录环境、保留中位数和范围。

13.8 Profiling 与单变量实验

HIP:

bash
hipcc --offload-arch=gfx1201 -O3 -std=c++17 \
  chapter13/rmsnorm_hip.hip -o /tmp/rmsnorm_hip
rocprofv3 --kernel-trace -- \
  /tmp/rmsnorm_hip --version block --rows 1024 --cols 4096 \
  --warmup 0 --repeat 10

Triton:

bash
rocprofv3 --kernel-trace -- \
  python chapter13/rmsnorm_triton.py --version t1 \
  --rows 1024 --cols 4096 --warmup 0 --repeat 10

建议按以下顺序实验:

  1. HIP block 只改 128/256/512
  2. 再把归约收尾替换成 wavefront shuffle;
  3. 再尝试 float4 读写,并保留 scalar tail;
  4. Triton 只改 num_warps
  5. 最后增加 FP16 输入、FP32 累加。

每一步记录 shape、假设、唯一改动、正确性、时间、VGPR/LDS/scratch 和结论。变慢的版本也要留下,因为它说明资源或 shape 边界在哪里。

13.9 HIP 与 Triton 对照

层次HIP blockTriton row program
行映射一个 block一个 program
局部平方和每线程寄存器program 向量
行归约LDS 树 + 同步tl.sum
广播标量shared[0]标量 SSA 值
尾部col < colsmask
配置变量block sizeblock size / num warps

两边的核心算法完全相同。HIP 需要显式决定线程如何合作;Triton 把同一工作提升为 program 内的向量和归约。理解映射关系比背诵某一份最终代码更重要。

13.10 独立优化记录模板

完成结业实验时,建议保留如下记录:

text
目标 shape / dtype:
硬件与软件:
reference 与误差阈值:
baseline:
本轮假设:
唯一改动:
正确性结果:
进程级计时与范围:
profiling 字段:
是否接受本轮改动:
负结果与回退点:

“最快版本”不是唯一产物。能解释为什么某个 block 在 257 列有效、在 4096 列反而出现 scratch,才是一份可迁移的优化记录。

13.11 迁移到新题目

拿到新的算子规格时,按以下顺序处理:

  1. 写数学语义和 CPU/PyTorch reference;
  2. 标记 Element-Wise、Reduction、GEMM 和可融合边界;
  3. 画输入、中间量和输出的数据生命周期;
  4. 做最短正确 HIP/Triton 基线;
  5. 固定 shape、计时和误差口径;
  6. 每轮只改变一项机制;
  7. 用 profiling 解释结果,而不是靠 kernel 名字解释。

LeetGPU 或其他平台题目可以作为扩展练习,但隐藏 shape、评分口径和硬件可能不同。平台成绩不能替代本书实验机上的可复跑记录。

正式实验结果

Chapter 13 RMSNorm 性能对比

主 shape 为 1024×4096 FP32。HIP serial 为 1.50486 ms,block 协作版为 0.058440 ms;Triton t0/t1 为 0.033360/0.030081 ms。串行行归约是明确负基线,而 t0/t1 的差距只适用于当前列数和资源配置。

完整记录见 code/part2-kernels/chapter13/EXPERIMENT.md

本章小结

  • RMSNorm 把逐元素平方、行归约、广播和权重缩放组合在一个小而完整的算子中。
  • 相对 LayerNorm,RMSNorm 去掉的是 re-centering(减均值),仍保留基于均方根的 re-scaling,并可带学习权重。
  • HIP serial 建立最短基线,HIP block 和 Triton row program 表达行内并行归约。
  • 融合主要减少中间数组和 dispatch;真实物理读写仍需 profiler 验证。
  • 当前实现已经提供 HIP、Triton、正确性 reference、GPU event、边界检查、3 个独立正式进程与逐实现 profiling 证据。

延伸阅读