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

把各軸疊起來:三維平行與專家混合

沒有任何單一切分能擴展到上兆參數。前沿做法是把資料、張量與管線平行組合成一張網格——再為稀疏專家加上第四個軸。

為何是組合而非二選一

每個平行軸單獨都會撞牆。張量平行需要最快的連結,超過單一伺服器就不再有效率。管線平行對頻寬便宜,但若拉得太薄,氣泡就會變大。資料平行能擴展批次卻擴不了模型大小。三維平行 把這三者組成單一裝置網格:在伺服器做張量平行、跨少數幾台伺服器做管線平行以橫跨模型深度、再跨由此得到的模型副本做資料平行以消化資料集。一場上兆參數的訓練可能用 TP=8、PP=16、DP=128——而這些的乘積就是 GPU 的數量。

神经缩放定律正是你要组合每一种并行轴、而非二选一的原因:随着规模推向万亿参数,损失持续下降,因此单一切分方式远远不够。

对数-对数坐标下的神经缩放定律曲线,显示损失随模型与数据规模增大而下降。

  1. 選 TP 讓單層的張量塞進一台伺服器,以節點內連結為上限。
  2. 選 PP,使模型的完整深度除以 TP 後,能在節點間連結上以可接受的氣泡攤開。
  3. 讓 DP 吸收剩下的 GPU,把全域批次擴向梯度雜訊尺度的目標。
  4. 在 DP 軸上加入激活值重算與 ZeRO 式優化器分片,補上任何剩餘的記憶體缺口。
N_{\text{GPU}} = d_{p}\cdot t_{p}\cdot p_{p}

总设备数等于数据并行、张量并行与流水线并行规模的乘积——这正是你要组合而非二选一的三维网格。

重疊是讓網格保持效率的關鍵

當三、四個集體通訊在不同軸上同時運行,唯有它們的流量藏在計算背後,叢集才保持高利用率。通訊與計算重疊 在此成了跨軸的排程問題:資料平行的 reduce-scatter 與反向傳播重疊,張量平行的 all-reduce 什麼都不能重疊(它們在關鍵路徑上,這正是 TP 保持小而本地的原因),管線的點對點傳送則與相鄰微批次的計算重疊。模型 FLOPs 利用率(MFU)——你真正轉化為有用訓練的峰值算力比例——是判斷你的三維佈局是否合理的單一數字;調校良好的兆級訓練回報的 MFU 落在 40–55% 區間。

第四個軸:專家平行

稀疏專家混合(MoE) 模型用許多專家子網路取代密集的前饋層,由路由器(router)為每個詞元只啟動其中一兩個。這把參數的數量與每詞元的計算量解耦——但也意味著每一層持有數十個專家,可能放不進單一裝置。專家平行 把專家分散到各裝置,且由於每個詞元都必須到達其路由器所選的專家,這一層的前後便被一個 all-to-all 集體通訊包夾,把詞元洗到它們的專家再洗回來。專家平行作為又一個軸嵌入三維網格之中,而那個 all-to-all 便成了 MoE 層的主要通訊成本。

y = \sum_{i \in \operatorname{TopK}(g(x))} g_i(x)\, E_i(x)

稀疏 MoE 层的输出,是仅对路由器为该 token 选中的前 k 个专家的加权求和。

更快的算術:以 FP8 訓練

這些平行化分配了工作;FP8 訓練 則讓每單位工作更便宜。現代加速器執行 8 位元浮點矩陣乘的吞吐量約為 bf16 的兩倍,因此把大型線性層轉成 FP8 直接換來速度。難處在動態範圍:尾數與指數位元只有寥寥數個,梯度與激活值可能上溢或下溢。標準對策是逐張量縮放(per-tensor scaling)——為每個張量追蹤一個縮放因子,使其數值落進 FP8 可表示的範圍——施加於矩陣乘的輸入,同時以較高精度累加並保留權重的主副本。FP8 能與每一個平行軸乾淨組合,因為它只改變本地計算的數值,不動通訊型態。

x_{\text{fp8}} = \operatorname{cast}\!\left(\frac{x}{s}\right),\quad s = \frac{\max|x|}{448}

逐张量 amax 缩放在更廉价的 8 位矩阵乘法前,把每个张量重新缩放到 FP8 的 E4M3 范围(最大约 448)。