究竟為什麼要自動化融合
一個 transformer 有數百個運算、以數十種形狀出貨;手動融合每條路徑是沒指望的。運算子融合編譯器會機械化地完成它。它接過你的框架本來就會建構的計算圖,找出可以安全融合的運算鏈,為每個區域產生一個核心,再排程整件事。你免費得到上一篇大部分的頻寬收益——而且關鍵在於,來自自動微分的反向圖也會被融合,而樸素框架正是在這裡漏掉最多記憶體流量。
Transformer 块示意图,将层归一化、多头注意力、残差连接和前馈网络显示为一连串算子。
XLA:提前編譯
XLA(Accelerated Linear Algebra,加速線性代數)會提前、針對一組固定的輸入形狀,把整張圖編譯成一個緊密融合、最佳化過的程式,供目標加速器使用。因為它一開始就看見整個運算,它能做出逐核心執行的執行期所無法做的全域決策——布局指派、融合邊界、緩衝區重用。代價是形狀改變時要重新編譯,這也是為什麼以 XLA 為基礎的技術堆疊倚賴固定或分桶(bucketed)的形狀與填補(padding)。它是 JAX 的骨幹、也是 TensorFlow 的一條主要路徑,並能從同一張圖瞄準 GPU 與 TPU。
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 regionCUDA Graphs:消滅啟動開銷
即使是完美融合的核心,每次啟動仍要付一筆固定的 CPU 成本——幾微秒的驅動程式工作。當核心很小(小批次、大型語言模型一次解碼一個 token)時,啟動開銷可能超過運算本身,GPU 在兩次啟動之間閒置、受 CPU 限制。CUDA Graphs 把整串啟動擷取一次、凍結成一張可重播的圖,讓 GPU 以單次提交就發射整串序列。每核心的 CPU 成本因此消失。它與上面的編譯器自然搭配——擷取融合後的程式,然後重播它。
Eager 模式的启动开销随核函数数量 N 线性增长;捕获的 CUDA Graph 用一次启动重放全部 N 个核函数,将 N·ℓ 压缩为单个 ℓ。
編譯器做不到的事
誠實面對極限。編譯器只在你給的邊界內融合:動態形狀會觸發重新編譯或 graph break;依賴張量值的控制流抗拒擷取;而且編譯器最佳化的是單一裝置的程式,不是裝置之間的通訊。這些工具沒有一個會把一個位元組送過網路。要超越單張加速器,你需要集體通訊與一張網路布料——這是下一篇。