跳到主要内容

数据截至 (上游 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 共享一个骨架,记牢它再看细节:

  1. 包成 torch.autograd.Function:手写 forward 启动 Triton kernel,手写 backward 启动另一个(或同一个)kernel;
  2. 按行 / 按块并行:tl.program_id(0) 拿到行号,一次 tl.load 进寄存器,算完 tl.store 回去,显存只进出一次;
  3. fp32 算、原 dtype 存:进 kernel 先 .to(tl.float32) 累加,最后 cast 回权重 dtype—— 注释标明 "Exact copy from HF"(unsloth/kernels/rms_layernorm.py:57);
  4. 能写回输入就写回输入(原地),反向省一张激活;
  5. @torch.compiler.disable:这些函数刻意不让 torch.compile 碰 (rms_layernorm.py:244rope_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.rsqrtinv_vartl.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.heuristicsGEMMA 标志切换分支(:116-120)。

对外包装是 fast_rms_layernorm(:245-255);模块级补丁 patch_rms_layernorm01 章 第 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)。

5. SwiGLU:前向一次融合,反向一个 kernel 吐三个梯度

直觉

SwiGLU MLP 是 h = silu(X@G) * (X@U); out = h@W。矩阵乘交给 cuBLAS,但中间的 silu(e) * g 和它的反向是纯逐元素运算——正是带宽黑洞。Unsloth 把它们各融成一个 kernel。

真实实现

  • 前向 swiglu_fg_kernel_fg_kernel(unsloth/kernels/swiglu.py:28-47): 读 e, g,fp32 算 f = e*sigmoid(e),cast 后 h = f*g 写出。一次读完一次写。
  • 反向 swiglu_DWf_DW_dfg_kernel(:68-109):一个 kernel 原地吐出三份东西—— h(留给下游算 dW)、df = DW*fde(silu 的导数链),分别写回 DW、e、g 三块输入缓冲区(:107-109),零新分配。

反向公式 kernel 头部用注释完整给出(:69-76),方便核对。

这两个 kernel 用在哪

它们是 04 章 LoRA_MLP 手工 autograd 的「点火塞」:前向的 _forward_function 和反向的 _backward_function 就是这两个 kernel (unsloth/kernels/fast_lora.py:261-262)。

6. 坑

  • 手写反向 = 数学必须对。每个 kernel 文件尾部自带数值测试(如 test_rms_layernorm,rms_layernorm.py:301-326)和 HF 实现对拍梯度,容差 0.05。
  • torch.compile 被刻意排除:Triton kernel 与 Dynamo 图相互踩,@torch.compiler.disable 是显式决定,不是遗漏。
  • kernel 收益随硬件变is_cdna()/is_rdna()(unsloth/kernels/utils.py:80-111) 针对 AMD 架构调 num_warps;同一 kernel 在 T4 上甚至可能放不下(交叉熵章会再遇到)。