核心概念

這是 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 kernel
  • fused__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::taten::viewaten::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 工作習慣(文章強調的方法論):

  1. 先寫下「預期 trace 中會看到什麼」
  2. 打開 trace 對照
  3. 把任何不符預期的發現當成最有價值的線索
  4. 這個習慣能建立直覺、揭示隱性效能問題

選擇融合策略

  • 靜態 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 案例完全吻合。

反向連結

以下頁面引用了本頁: