從運算到核心函式
你呼叫的每個 PyTorch 運算——`x.relu()`、`a + b`、一次 softmax——最終都會啟動一個或多個核心函式:在 GPU 數千條執行緒上跑的程式。麻煩在於一連串運算會啟動一連串核心,每一個都從 HBM 讀取輸入、再把輸出寫回。一個由五個逐元素運算組成的殘差區塊,幾乎沒做什麼算術,卻碰了十次記憶體——這是教科書等級的頻寬受限災難。解方是讓 GPU 在兩次記憶體往返之間做更多工作。
融合:槓桿最高的一步
核心函式融合(kernel fusion)把數個運算合併成單一核心,讓中間值留在暫存器或共享記憶體中,永遠不寫回 HBM。把上面那五個逐元素運算融合起來,你只讀輸入一次、在晶片上算完所有東西、再寫結果一次——記憶體流量從十次往返降到兩次,算術強度也按同樣倍數提升。融合是深度學習中最有效的單一頻寬最佳化,這也是為什麼如此多的編譯器與核心工程,本質上都是在設法更積極地融合。
CUDA 執行模型
要手寫一個融合核心,你必須面對 GPU 真正的機器模型,而 CUDA 核心最佳化就是把你的問題妥善映射到它上面的功夫。執行緒以 32 條為一組、步調一致地執行,稱為執行緒束(warp);多個 warp 組成共享一塊快速暫存區共享記憶體(shared memory)的執行緒區塊(thread block);區塊再被排程到串流多處理器上。三個槓桿主導效能。佔用率(occupancy):保持足夠多的 warp 駐留,以掩蓋記憶體延遲。合併存取(coalescing):安排成讓一個 warp 中的 32 條執行緒讀取 32 個連續位址,把許多小讀取變成一次寬交易。共享記憶體分塊(tiling):把一塊資料暫存在晶片上、在執行緒間重複使用,而不是重讀 HBM。
# tiled matmul, conceptually
for tile_k in range(0, K, TILE):
load A[:, tile_k:tile_k+TILE] into shared
load B[tile_k:tile_k+TILE, :] into shared
__syncthreads() # barrier: tile is ready
accumulate partial products from shared, not HBM
__syncthreads()分块为何提高复用,量化来看:一个 B×B 的数据块只加载一次就能供给约 B³ 次乘加,因此算术强度随块尺寸 B 线性增长。
Triton:你真的維護得了的核心函式
手寫 CUDA 既強大又難以維護。Triton 命中了甜蜜點:你寫的是對整塊tile(元素區塊)操作的 Python,而編譯器替你處理每執行緒映射、合併存取與共享記憶體配置。你仍然以 tile 和強度來思考,但不必再手動管理 32 條執行緒的 warp。實務上,一個 Triton 核心能用零頭的程式碼達到專家級 CUDA 的相當大一部分效能,這也是它成為研究者交付自訂核心的預設方式的原因。
@triton.jit
def fused_gelu_kernel(x_ptr, y_ptr, n, BLOCK: tl.constexpr):
off = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
mask = off < n
x = tl.load(x_ptr + off, mask=mask) # one HBM read
y = x * 0.5 * (1.0 + tl.erf(x * 0.7071)) # compute on-chip
tl.store(y_ptr + off, y, mask=mask) # one HBM write案例研究與何時收手
FlashAttention 是經典的回報。它把整個注意力運算——分數、softmax、以值加權的總和——融合成一個核心,沿序列分塊,使 N×N 分數矩陣從不碰 HBM,並用線上 softmax 技巧來處理歸約而不具現化整列。FlashAttention-2 接著重新調校跨 warp 的工作切分,把佔用率推得更高。結果是又更快、又更省記憶體,讓上下文能長得多。
FlashAttention 融合进单个分块核函数的全部计算——打分、softmax 以及按值加权求和,全程不实体化整张 N×N 矩阵。