PyTorch 性能分析实战:从 nn.Linear 到融合 MLP 内核

很多人以为深度学习模型的性能瓶颈都在 GPU 计算上,一遇到速度慢就无脑上 torch.compile。但真相是,**对于单个 nn.Linear 层,torch.compile 几乎什么都不做**,因为偏置加法已经被折叠进 cuBLAS 的 addmm 内核里了。这篇文章通过实际分析 PyTorch profiler 的 trace,把这个常见误区解释得清清楚楚。

核心洞察在于区分两种场景。在单层 Linear 下,编译器唯一的作用是消除 CPU 端的 view 调度开销,GPU 内核完全一样。但到了 MLP 模块,情况就不同了:**torch.compile 能够将 GeLU 和逐元素乘法融合成单个 Triton 内核**,让中间结果留在寄存器里,省掉一次 HBM 读写。文章还引入了 Liger 手写内核,它不需要动态重编译就能拿到同样的融合效果,代价是在特定静态形状上可能比 Inductor 微秒级更慢。

这篇文章最大的启发是教会你如何像工程师一样思考:在分析 trace 之前先做预测,然后对照验证。不要盲目信任任何优化手段。对于更复杂的模型如 Attention Block 和完整模型,**阅读 trace 的能力将成为你理解性能瓶颈最重要的习惯**。这种对硬件底层行为的直觉,是用任何黑盒优化器都无法替代的。

Profiling in PyTorch (Part 2): From nn.Linear to a Fused MLP

查看原文