数据截至 (上游 commit 1fe27b1b53f3)
02 · 手写 Triton kernel:RMSNorm / RoPE / SwiGLU
这章回答:一个手写的 Triton kernel 到底比框架默认实现快在哪、省在哪。用三个代表性 kernel 看同一套打法。
1. 小问题
transformers 里 RMSNorm、RoPE、SwiGLU 都是逐元素 / 按行的小运算,但 PyTorch 默认会把 它们拆成 多个 kernel launch,每个中间结果都写回显存再读回来:
- 慢在显存带宽:GPU 算得快、搬得慢,小运算全是搬运;
- 废在 autograd 图:默认反向要为每个中间量建图、存激活。
Unsloth 的答案是:每个运算写成一个 Triton kernel(前向)+ 一个 Triton kernel(手工反 向),中间量不出显存寄存器,反向不要 autograd 图。
2. 统一打法(三个 kernel 的共性)
所有 kernel 共享一个骨架,记牢它再看细节:
- 包成
torch.autograd.Function:手写forward启动 Triton kernel,手写backward启动另一个(或同一个)kernel; - 按行 / 按块并行:
tl.program_id(0)拿到行号,一次tl.load进寄存器,算完tl.store回去,显存只进出一次; - fp32 算、原 dtype 存:进 kernel 先
.to(tl.float32)累加,最后 cast 回权重 dtype—— 注释标明 "Exact copy from HF"(unsloth/kernels/rms_layernorm.py:57); - 能写回输入就写回输入(原地),反向省一张激活;
@torch.compiler.disable:这些函数刻意不让 torch.compile 碰 (rms_layernorm.py:244、rope_embedding.py:265注释 "[TODO] Unsure why ... not torch.compiling properly")。
BLOCK_SIZE 等 launch 参数由 calculate_settings 按列数取 2 的幂算出
(unsloth/kernels/utils.py:114)。
3. RMSNorm:一行一 program,inv_var 顺手存给反向
直觉
RMSNorm 对每行算 x / sqrt(mean(x²) + eps) * w。默认实现要 launch 平方、求均值、除、乘
好几个 kernel;融合后一次读 X、一次写 Y,顺便把每行的 inv_var 存下来——反向要用,
省得重算。
真实实现
前向 kernel _rms_layernorm_forward(unsloth/kernels/rms_layernorm.py:22-59):
row_idx = tl.program_id(0),每个 program 处理一行(:40);- 加载后转 fp32 算方差,
tl.math.rsqrt得inv_var并tl.store(r, inv_var)(:51-55); normed.to(W_row.dtype)再乘权重(:57-58)——顺序和 HF 逐位一致。
反向 _rms_layernorm_backward(:62-112)按 LayerNorm 反向的标准闭合式直接写 dX:
# 示意,非源码 —— RMSNorm 反向的闭合式
# 已知 inv_var(前向存的),输入 dY、X、W
normed = X * inv_var
dY_W = dY * W
dX = inv_var / n * (n * dY_W - normed * sum(dY_W * normed))
两个细节:
- 省显存的原地写:非 Gemma 时
dX = dY(:95),直接把梯度写回 dY 的缓冲区,不新开 张量(Gemma 因为 dY 还要复用才单独分配,:218); - Gemma 变体单独一个 kernel:Gemma 的权重语义是
(1 + W),所以前向有_gemma_rms_layernorm_forward(:123-159),反向用triton.heuristics按GEMMA标志切换分支(:116-120)。
对外包装是 fast_rms_layernorm(:245-255);模块级补丁 patch_rms_layernorm 见
01 章 第 7 节。
4. RoPE:原地旋转,反向 = 同一个 kernel 取负 sin
直觉
RoPE 把每对维度按位置旋转:Q1' = Q1*cos - Q2*sin,Q2' = Q2*cos + Q1*sin。它是可逆
的逐对变换:反向不需要存任何激活,把同一个旋转「倒着转」(sin 取负)施加到梯度上 即可。
真实实现
kernel _rope_embedding(unsloth/kernels/rope_embedding.py:104-156):
- 原地写 Q:
tl.load读出Q1, Q2,算完tl.store回同一块内存(:151-155); - 4 个头一组(
ROPE_GROUP_SIZE = 4,:121):一个 program 循环处理 4 个头,摊薄 cos/sin 的加载——注释注明这是 PR#238 的 10% 提速(:145); BACKWARD_PASS开关:kernel 里只有一行差别sin1 = -sin1(:137-139),triton.heuristics按参数生成两份机器码(:159-163)。
Fast_RoPE_Embedding.backward(:222-264)因此对梯度把同一 kernel 再打一遍,什么都不
用 save(ctx 只存 cos/sin 和 launch 参数,:214-218)。入口 fast_rope_embedding
在 :266-280:Q、K 各转置后过 kernel;有 rope_embedding_indices(变长 packing)时走
Fast_RoPE_Embedding_QK(:283)。