JOVANA
Explore Library Glossary Getting Started Three Levels Fields How it works Mission
Join the mission
All guides

會融合的編譯器:XLA、torch.compile 與 CUDA Graphs

你無法手寫每一個核心函式。機器學習編譯器擷取運算圖、自動融合,並以近乎零的啟動開銷重播它。

究竟為什麼要自動化融合

一個 transformer 有數百個運算、以數十種形狀出貨;手動融合每條路徑是沒指望的。運算子融合編譯器會機械化地完成它。它接過你的框架本來就會建構的計算圖,找出可以安全融合的運算鏈,為每個區域產生一個核心,再排程整件事。你免費得到上一篇大部分的頻寬收益——而且關鍵在於,來自自動微分反向圖也會被融合,而樸素框架正是在這裡漏掉最多記憶體流量。

Transformer 區塊的眾多運算子——正規化、注意力、殘差、前饋網路——正是融合編譯器必須擷取並機械式融合的長圖。

Transformer 區塊示意圖,將層正規化、多頭注意力、殘差連接與前饋網路顯示為一連串運算子。

XLA:提前編譯

XLA(Accelerated Linear Algebra,加速線性代數)會提前、針對一組固定的輸入形狀,把整張圖編譯成一個緊密融合、最佳化過的程式,供目標加速器使用。因為它一開始就看見整個運算,它能做出逐核心執行的執行期所無法做的全域決策——布局指派、融合邊界、緩衝區重用。代價是形狀改變時要重新編譯,這也是為什麼以 XLA 為基礎的技術堆疊倚賴固定或分桶(bucketed)的形狀與填補(padding)。它是 JAX 的骨幹、也是 TensorFlow 的一條主要路徑,並能從同一張圖瞄準 GPU 與 TPU。

Q_{\text{unfused}} \approx 2kn \;\gg\; 2n \approx Q_{\text{fused}}

XLA 的緊密融合程式為何快:未融合時,k 個逐元素運算子的每個中間結果都要往返 HBM(約 2kn 位元組);融合後將其保留在暫存器中,只搬運約 2n。

torch.compile:擷取即時執行的運算圖

PyTorch 是即時的——運算隨著 Python 執行而執行——所以沒有圖可編譯。torch.compile 製造出一張。TorchDynamo 在執行期追蹤你的 Python,把它擷取成圖,並在遇到無法追蹤的東西時優雅地回退(fall back)到即時執行(即所謂的 graph break)。擷取到的圖交給一個後端——預設是 Inductor,它融合逐點運算、為 GPU 產生 Triton 核心,並對大型 GEMM 呼叫廠商函式庫。你保留 Python 的彈性,又在要緊的部分得到編譯後的速度。

model = torch.compile(model)        # first call traces + compiles
for batch in loader:                # later calls hit the cached graph
    loss = model(batch).loss
    loss.backward()
# watch for 'graph break' logs: each break splits the fused region
一行即可啟用;收益來自讓融合區域保持大塊,所以要盡量減少 graph break。

CUDA Graphs:消滅啟動開銷

即使是完美融合的核心,每次啟動仍要付一筆固定的 CPU 成本——幾微秒的驅動程式工作。當核心很小(小批次、大型語言模型一次解碼一個 token)時,啟動開銷可能超過運算本身,GPU 在兩次啟動之間閒置、受 CPU 限制。CUDA Graphs 把整串啟動擷取一次、凍結成一張可重播的圖,讓 GPU 以單次提交就發射整串序列。每核心的 CPU 成本因此消失。它與上面的編譯器自然搭配——擷取融合後的程式,然後重播它。

T_{\text{eager}} \approx N\ell + \sum_{k=1}^{N} t_k \;\xrightarrow{\text{CUDA Graph}}\; T_{\text{graph}} \approx \ell + \sum_{k=1}^{N} t_k

Eager 模式的啟動開銷隨核函式數量 N 線性增長;擷取的 CUDA Graph 以單次啟動重放全部 N 個核函式,將 N·ℓ 壓縮為單一 ℓ。

編譯器做不到的事

誠實面對極限。編譯器只在你給的邊界內融合:動態形狀會觸發重新編譯或 graph break;依賴張量值的控制流抗拒擷取;而且編譯器最佳化的是單一裝置的程式,不是裝置之間的通訊。這些工具沒有一個會把一個位元組送過網路。要超越單張加速器,你需要集體通訊與一張網路布料——這是下一篇。