高中生开发者连续 9 周编写 Triton kernel 并与 PyTorch 基准对比,覆盖融合 softmax、交叉熵、RoPE、Flash Attention 等,详解内存带宽约束下的融合策略与正确性校验方法。
我是 Richael (dh8116),奥克兰一名 11 年级学生。自从 7 月底以来,我每周写一个 Triton kernel,并与 PyTorch 做基准测试对比,即使 PyTorch 经常胜出,我也公布测试数据。以下是九周教会我的东西,一个 kernel 接一个 kernel。
融合 softmax(第 2 周)。朴素 softmax 需要三次遍历内存。融合成一次后,从约 55 GB/s 提升到约 230 GB/s,4 倍的跳跃,而且比 PyTorch 自己的 kernel 表现更稳定。经验一:对于内存密集型操作,那些你不需要执行的遍历就是加速。
融合交叉熵(第 7 周)。前向和反向一次完成,梯度直接写入 logits 缓冲区,这样 PyTorch 的第二个 [N, V] 张量永远不会存在。在 T4 上,词表 131,072、fp16 精度:15.90 ms vs 24.03 ms(1.51 倍),峰值内存减少 1.67 倍。
这个基准测试的第一个版本显示 2.5 倍。那是错误的:我在将融合后的 kernel 与未融合的 PyTorch 基线做比较。融合对融合的比较让我发现了自己的基线设置有偏差。
RoPE(第 6 周)。不是讲加速,而是讲优雅:反向传播复用了与前向完全相同的 kernel,只需要把 sin 的符号翻转一下。
Flash Attention(第 3 周)。正确,但在我的 T4 上比 PyTorch 慢约 25 倍。
矩阵乘法(第 4 周)。输出匹配,但吞吐量始终停留在约 1 TFLOPS,而 cuBLAS 攀升到 38。读取编译后的 PTX 发现零条 tensor-core 指令:编译器从未为我的布局生成这些指令。
融合线性层 + 交叉熵(第 8 周)。将 lm_head 投影分块到 loss 计算中意味着 [batch, vocab] 的 logits 永远不会存在。那大约是 1/10 的激活内存,但在每个尺寸上都更慢。有时候内存才是你买的东西,而你用时间为之付费。
融合 SwiGLU MLP(第 9 周)。时间上与 torch.compile 持平。赢在内存:它在 forward 和 backward 之间保留两个大激活张量,而 eager 保留四个。编译器会帮你做融合,但它不会丢掉激活。
基准测试要融合对融合。 未融合的基线会让每个 kernel 看起来都像赢了。
数字不对时读 PTX。 矩阵乘法的答案一直就藏在编译输出里。
内存也是结果。 我两个最好的 kernel 在时间上"输"了、在内存上赢了,而这个权衡决定了 batch size。
公布失败的结果。 它们教会我更多,而且这些帖子人们才会真正来讨论。
代码:https://github.com/dh8116/triton-kernels 所有详细文章:https://dh8116.github.io/blog
下一个计划是在 Ampere 上重新测试 flash attention,因为 T4 的结果值得再看一下。