在 PyTorch 中的效能分析(第 2 部分):從 nn.Linear 到融合的 MLP

在 PyTorch 中的效能分析(第 2 部分):從 nn.Linear 到融合的 MLP

Hugging Face 詳細說明了在 PyTorch 中優化多層感知器(MLP)的過程,展示了如何從標準的 nn.Linear 層轉換為融合的 kernel,以減少 CPU 開銷和 GPU 記憶體流量。主要的結論是,雖然 torch.compile 能將點對點運算融合成單一的 Triton kernel,從而避免昂貴的 HBM 往返,但手工調校的 kernel 能在不產生編譯延遲與形狀專門化開銷的情況下提供相似的效能提升。

nn.Linear 的機制與 Kernel 後置程式

在 PyTorch 中,nn.Linear 是矩陣乘法與加法的包裝器。當 bias=True 時,運算會以 y = x @ w.T + b 方式執行。

透過後置程式折疊偏差

效能分析顯示,偏差加法並沒有單獨的 aten::add kernel。相反地,偏差加法會透過 後置程式 (epilogue) 折疊進矩陣乘法 kernel 中。後置程式是由 GEMM(General Matrix Multiply)kernel 在將結果寫回高頻寬記憶體 (HBM) 前執行的一小段計算。透過將偏差加法整合到 matmul kernel 的寫回階段,PyTorch 免除了第二次載入與寫入 HBM 的需求,從而減少記憶體流量。

CPU 調度與轉置視圖

aten::t(轉置)運算會出現在 eager CPU 調度鏈中,位於 aten::addmm 之前。然而,這不會啟動 GPU kernel。aten::t 只是在 CPU 上重新寫入張量的元資料(形狀與步幅),以建立相同原始資料的新視圖。

當對單一的 nn.Linear 層使用 torch.compile 時,所使用的 GPU kernel 不會改變——仍執行相同的 cuBLAS GEMM kernel。相反地,它會在編譯時硬編碼轉置視圖的步幅,從而消除 CPU 調度轉置視圖的開銷。

GeGLU MLP 的效能分析

GeGLU MLP 包含三個線性投影(gate_projup_projdown_proj)以及一個點對點激活序列(GeLU 後接乘法)。

eager 執行效能

在 eager 模式下,GeGLU MLP 的前向傳播會啟動五個不同的 GPU kernel:三個 GEMM 與兩個點對點 kernel(GeLU 與乘法)。GeLU 運算產生的中間張量必須寫入 HBM,然後再被乘法 kernel 讀回,形成記憶體瓶頸。

GEMM Kernel 的差異

即使 FLOP 數相同,並非所有 GEMM kernel 都相同。例如,down_proj 可能比 gate_proj 更快,因為 cuBLAS 會根據輸入形狀選擇不同的 tile 大小(例如 128x256128x128)以及管線深度,以最佳化資料重用。

透過 torch.compile 與融合 Kernel 進行最佳化

torch.compile 的影響

torch.compile 透過將 GeLU、乘法與 reshape 操作合併為單一的融合 Triton kernel(例如 triton_poi_fused__unsafe_view_gelu_mul_0)來最佳化 MLP。

此融合透過將中間值保留在暫存器中而非寫入 HBM,帶來顯著的效能提升。然而,此最佳化也伴隨成本:編譯器會加入「前置作業」如 TorchDynamo guard 與前置開銷,且會針對特定輸入形狀專門化 kernel。若輸入形狀改變,模型必須重新追蹤與重新編譯。

使用 kernels 函式庫的手工調校 Kernel

為避免編譯延遲與形狀專門化,可透過 Hugging Face 的 kernels 函式庫使用手工調校的 kernel。例如,LigerGEGLUMLP 層使用預先構建的 Triton kernel,將 GeLU 與乘法的融合內建於其中。

編譯 vs 手工調校 Kernel 的比較:

特性 torch.compile (Inductor) Hand‑Tuned (Liger)
融合 動態融合點對點運算 內建融合
CPU 開銷 Dynamo guard 與前置程式 無編譯前置作業
形狀處理 針對靜態形狀專門化(對單一形狀最快) 通用啟動參數(對形狀變化具韌性)
部署 需要本地編譯/Triton 透過 HF Hub 提供的預建二進位檔

雖然編譯後的 kernel 可能因極度專門化而在特定靜態形狀上稍快,但手工調校的 kernel 更具韌性,且可避免 PyTorch 編譯管線相關的風險與延遲。

最佳化階段摘要

設定 GPU 變更 CPU 變更
eager nn.Linear 偏差加法折疊進 GEMM 後置程式 基線調度
compiled nn.Linear 無變更(相同 cuBLAS kernel) 移除 aten::t 視圖的帳務管理
eager MLP 5 個 kernel;中間結果寫入 HBM 基線調度
compiled MLP GeLU + 乘法融合為單一 Triton kernel 加入 Dynamo/guard 前置作業
Liger MLP 與 compiled 相同的融合;調校的啟動參數 無編譯前置作業;無重新編譯風險

摘要: Hugging Face 說明了如何透過分析 GPU kernel 融合、torch.compile 的影響,以及使用 kernels 函式庫的手工調校 kernel,來最佳化 PyTorch 中的多層感知器(MLP)。

Sources