核心概念
這是 Hugging Face 《PyTorch Profiling》系列第二篇,聚焦在「算子融合(kernel fusion)」如何消除不必要的記憶體搬移,以實際 GeGLU MLP 為例,逐步說明從 eager 模式到 torch.compile 再到手工 Triton kernel(Liger)的效能演化。請先閱讀 PyTorch Profiling 入門:torch.profiler 追蹤解讀指南 了解 profiler 基礎。
nn.Linear 的底層真相
y = linear(x) 在 GPU 上只觸發一個 cuBLAS GEMM kernel(aten::addmm),原因有二:
轉置(aten::t)不是真正的 GPU 計算:PyTorch 只重寫張量的 metadata(shape + stride),完全不搬移資料、不啟動 GPU kernel。Stride 例如 (2, 1) 表示「移動一列需跨 2 個元素,移動一行需跨 1 個元素」,轉置只是把這對數字交換。
Bias 加法折入 epilogue:cuBLAS GEMM kernel 在寫出結果到 HBM(GPU 高頻寬記憶體)前,有一個「epilogue」階段可以插入小型計算(加 bias、套 activation、縮放等)。因此 addmm(bias, x, weight) 等於 x @ weight.T + bias,只需一個 kernel 完成。
使用 torch.compile 後,單一 Linear 的 GEMM kernel 字節完全相同——compile 只消除了 CPU dispatch 的 aten::t 開銷,GPU 端無任何改變。
GeGLU MLP 的五個 GPU Kernel
以 SimpleGeGLUMLP(含 gate_proj / up_proj / down_proj 三個 Linear + GeLU 激活 + element-wise 乘法)為例,eager 模式下一次 forward 觸發五個 GPU kernel:
| 操作 | GPU Kernel 類型 |
|---|---|
| gate_proj | cuBLAS GEMM(128×128 tile) |
| up_proj | cuBLAS GEMM(128×128 tile) |
| gelu(g) | vectorized elementwise |
| h × u | vectorized elementwise |
| down_proj | cuBLAS GEMM(128×256 tile) |
有趣的是 down_proj 雖然 FLOP 數與另外兩個 GEMM 相同(約 38.7 GFLOP),卻快約 10%:cuBLAS 為不同矩陣形狀選擇不同的 tile 策略(128×256 / stages_64×3),下投影矩陣形狀讓 cuBLAS 能做更好的資料重用。
torch.compile 的關鍵融合
compile 後三個 GEMM kernel 完全不變(kernel 名稱字節相同),但兩個 pointwise kernel 消失,合併為一個 Triton kernel:
triton_poi_fused__unsafe_view_gelu_mul_0
poi= pointwise kernelfused__unsafe_view_gelu_mul= 融合了 reshape、GeLU、mul 三個操作0= 圖內唯一 ID
為何融合能加速:eager 模式下,GeLU 會先把中間結果 h = gelu(g) 寫入 HBM(shape [8192, 3072] bf16 ≈ 50 MB),然後 mul 立刻再把它讀回來——一次完整的 HBM 往返。融合後,kernel 直接在暫存器(on-chip memory)裡計算 gelu(g) * u,完全不碰 HBM,消除整個往返開銷。
Liger 手工 Triton Kernel
kernels 套件提供 LigerGEGLUMLP,一個 drop-in 替換,無須 torch.compile:
from kernels import get_kernel
kernels_layers = get_kernel("kernels-community/liger-kernels", version=1).layers
liger_mlp = kernels_layers.LigerGEGLUMLP(Config()).to(device, dtype=torch.bfloat16)
Liger 的優勢:
- 融合直接寫進 Triton kernel(
_geglu_tanh_forward_kernel),不需 Dynamo 編譯 - launch 參數根據
calculate_settings針對實際硬體調優 - 版本鎖定(
version=1),CI 預編譯完成,不同環境行為一致
效能對比(NVIDIA A100-SXM4-80GB):
| 實作 | Kernel 時間 | 特點 |
|---|---|---|
| Compiled MLP | 89.4 µs | 靜態 shape 最快,shape 變動需重新編譯 |
| LigerGEGLUMLP | 92.8 µs | 稍慢 3%,任意 shape 無需重編 |
兩者效能差距極小,選擇取決於是否需要動態 shape 支援。
關鍵要點
- 轉置是純 metadata 操作:
aten::t、aten::view、aten::reshape等在 profiler 的 CUDA time 欄位顯示 0µs,不啟動任何 GPU kernel,只重寫 stride - Bias add = GEMM epilogue:不存在獨立的 add kernel,cuBLAS 把它折入 GEMM 末段
- torch.compile 不改 GEMM,改的是 pointwise:GEMM kernel 名稱在 eager 與 compile 模式下字節完全一致,compile 的貢獻是融合 elementwise 操作和消除 CPU dispatch 開銷
- HBM 往返是效能殺手:融合的意義不在於少執行幾次計算,而在於避免把中間張量寫入 HBM 再讀回;50 MB 的往返在現代 GPU 上也需要數十微秒
- cuBLAS 自動選 tile:同樣 FLOP 數的矩陣,因形狀不同可能選到不同的 tile 策略,效能差距可達 10%
實務應用
Profiling 工作習慣(文章強調的方法論):
- 先寫下「預期 trace 中會看到什麼」
- 打開 trace 對照
- 把任何不符預期的發現當成最有價值的線索
- 這個習慣能建立直覺、揭示隱性效能問題
選擇融合策略:
- 靜態 shape 推論(batch size 固定)→
torch.compile最省力 - 動態 shape 推論 → 考慮 Liger 等手工 kernel,避免反覆重編
- 需要跨架構相容 →
kernels套件的版本鎖定確保 CI 到生產環境一致
何時不需要擔心融合:如果 bottleneck 是 GEMM 本身(通常是大 batch 的情況),融合 elementwise 的效益邊際化;先確認瓶頸在哪再優化。
延伸觀點
根據 PyTorch、CUDA 優化相關文獻的交叉驗證:
Epilogue fusion 是 cuBLAS 長期存在的設計:不只是 bias,activation(ReLU、sigmoid)、scaling、residual add 都可以折入 epilogue,這是 cuBLAS 為 transformer 優化的重要機制,在 cuBLASLt(cuBLAS Lightning)的介面文件中有明確說明。
Triton 替代 CUDA C 的趨勢:Hugging Face、Meta、Liger 等專案越來越多用 Triton 寫自訂 kernel,原因是 Triton 可以在 Python 中直接表達 tile 級別的邏輯,且 torch.compile 的 Inductor 後端也以 Triton 作為 GPU codegen 目標。手工 Triton 與 compile 生成的 Triton 在效能上趨近,差距通常在 5% 以內。
記憶體頻寬是現代 GPU 真正的瓶頸:對 LLM 推論來說,GEMM 通常是 compute-bound,但 attention softmax、layernorm、activation 等 elementwise 操作是 memory-bound。融合優先順序應針對後者——這與本文的 GeGLU 案例完全吻合。
反向連結
以下頁面引用了本頁: