PyTorch 性能分析(第 2 部分):从 nn.Linear 到融合的 MLP
PyTorch 性能分析(第 2 部分):从 nn.Linear 到融合的 MLP
Hugging Face 详细阐述了在 PyTorch 中优化多层感知机(MLP)的过程,展示了如何从标准的 nn.Linear 层转向融合内核,以降低 CPU 开销和 GPU 内存流量。主要结论是,虽然 torch.compile 可以将点状操作融合为单个 Triton 内核,从而避免昂贵的 HBM 往返,但手工调优的内核可以在不产生编译延迟和形状专化开销的情况下提供类似的性能提升。
nn.Linear 的工作原理与内核尾部(Epilogues)
在 PyTorch 中,nn.Linear 是矩阵乘法和加法的包装器。当 bias=True 时,操作执行为 y = x @ w.T + b。
通过尾部(Epilogue)实现偏置折叠
性能分析显示,偏置加法没有单独的 aten::add 内核。相反,偏置加法被折叠进矩阵乘法内核,使用 epilogue(尾部)实现。epilogue 是在 GEMM(通用矩阵乘法)内核将结果写回高速带宽内存(HBM)之前执行的一小段计算。通过将偏置加法整合到 matmul 内核的写回阶段,PyTorch 避免了第二次加载和写入 HBM,从而降低了内存流量。
CPU 调度与转置视图
aten::t(转置)操作出现在 eager CPU 调度链中 aten::addmm 之前。然而,这并不会启动 GPU 内核。aten::t 仅在 CPU 上重写张量的元数据(形状和步幅),以创建对相同原始数据的新视图。
当对单个 nn.Linear 层使用 torch.compile 时,它并不会改变所使用的 GPU 内核——仍然执行相同的 cuBLAS GEMM 内核。相反,它通过在编译时硬编码得到的步幅,消除了调度转置视图的 CPU 开销。
对 GeGLU MLP 的性能分析
GeGLU MLP 由三个线性投影(gate_proj、up_proj 和 down_proj)以及一个点状激活序列(GeLU 后接乘法)组成。
Eager 执行性能
在 eager 模式下,GeGLU MLP 的前向传播会启动五个不同的 GPU 内核:三个 GEMM 和两个点状内核(GeLU 和乘法)。GeLU 操作产生的中间张量必须写入 HBM,然后被乘法内核读取,导致内存瓶颈。
GEMM 内核差异
即使 FLOP 计数相同,GEMM 内核也不完全相同。例如,down_proj 可能比 gate_proj 更快,因为 cuBLAS 会根据输入形状选择不同的 tile 大小(例如 128x256 与 128x128)和流水线深度,以优化数据复用。
通过 torch.compile 与融合内核进行优化
torch.compile 的影响
torch.compile 通过将 GeLU、乘法和 reshape 操作合并为单个融合的 Triton 内核(例如 triton_poi_fused__unsafe_view_gelu_mul_0)来优化 MLP。
这种融合通过将中间值保留在寄存器中而不是写入 HBM,带来了显著的性能提升。然而,这种优化也有代价:编译器会引入诸如 TorchDynamo guard 与前置开销等 “pre-ops”,并且会针对特定输入形状进行内核专化。如果输入形状改变,模型必须重新追踪并重新编译。
使用 kernels 库的手工调优内核
为了避免编译延迟和形状专化,可以通过 Hugging Face 的 kernels 库使用手工调优的内核。例如,LigerGEGLUMLP 层使用预构建的 Triton 内核,将 GeLU 与乘法的融合内置其中。
编译内核与手工调优内核的比较:
| 特性 | torch.compile (Inductor) |
手工调优 (Liger) |
|---|---|---|
| 融合 | 点状操作的动态融合 | 内置融合 |
| CPU 开销 | Dynamo guard 与前置开销 | 无编译前置操作 |
| 形状处理 | 针对静态形状专化(对单一形状最快) | 通用启动参数(对形状变化更稳健) |
| 部署 | 需要本地编译/Triton | 通过 HF Hub 提供的预构建二进制文件 |
虽然编译内核由于极端专化在特定静态形状下可能略快,但手工调优的内核更为稳健,且避免了 PyTorch 编译流水线带来的风险和延迟。
优化阶段概览
| 设置 | GPU 变化 | CPU 变化 |
|---|---|---|
Eager nn.Linear |
偏置加法折叠进 GEMM 尾部 | 基线调度 |
Compiled nn.Linear |
无变化(相同 cuBLAS 内核) | 移除 aten::t 视图记录 |
| Eager MLP | 5 个内核;中间结果写入 HBM | 基线调度 |
| Compiled MLP | GeLU + 乘法融合为一个 Triton 内核 | 添加 Dynamo/guard 前置操作 |
| Liger MLP | 与编译相同的融合;调优的启动参数 | 无编译前置操作;无重新编译风险 |