機器學習系統與基礎設施
GPU 核心融合(kernel fusion)
在 GPU 上跑一連串基本運算時——例如先做矩陣乘法,再加偏置,再過一個 GELU——最天真的做法是每個運算各啟動一個核心(kernel),而每個核心都從全域記憶體讀進輸入、算完再把輸出寫回去。對這類便宜的逐元素運算來說,這些往返慢速晶片外記憶體的成本反而主導了整體開銷。核心融合把整條鏈縫進單一核心,讓中間值完全不離開晶片,一直留在暫存器或共享記憶體裡,直到最後結果只被寫出一次。
融合主要在受限於記憶體(memory-bound)的運算上划算,因為它同時省下搬動的位元組數與每個核心的啟動開銷。最典型的樣式,是把逐元素的尾段(偏置、激活、dropout、殘差相加)融進前面矩陣乘法或正規化的尾巴。它受晶片上資源所限:融合核心必須把所有存活的中間值塞進暫存器與共享記憶體,所以融得太兇就會溢出(spill)回記憶體或拉低佔用率(occupancy)。歸約(reduction)會讓融合更棘手,因為它引入跨執行緒的依賴,是單純的逐元素融合器吸收不了的。
在 Transformer 的訓練與推論裡,把正規化、注意力尾段、MLP 激活融在一起,能救回原本浪費在記憶體間搬運激活值的算術強度;flash attention 本身就是一個被積極手工融合的核心。現代編譯器已能自動完成其中許多工作,但生產環境中價值最高的那些融合,往往仍是手寫的。
融合越多並非越好:晶片上塞太多中間值會壓低佔用率或被迫溢出,所以最佳點是一個平衡,而非極大化。
又称
另见