Qwen 用於 MoE LLM 訓練的 Global-Batch 負載平衡
Global-batch 負載平衡提升 MoE 性能與專家專業化
Qwen 為訓練 Mixture-of-Experts (MoE) 大型語言模型 (LLMs) 引入了一種 global-batch 負載平衡損失方法。透過將負載平衡損失的計算從 micro-batch 層級轉移到 global-batch 層級,模型在各種任務中實現了更好的性能,並促進了專家之間顯著的領域專業化,解決了 micro-batch 平衡往往阻礙專家在特定數據領域專業化的限制。
Micro-batch 負載平衡的局限性
大多數現有的 MoE 訓練框架(例如 Megatron-core)是在 micro-batch 層級實現負載平衡損失 ($L_{\text{balance}}$)。這意味著損失是在每個 micro-batch 內計算,然後在整個 global batch 中取平均值。
當一個 micro-batch 缺乏多樣化數據時,這種方法會產生問題。例如,如果一個 micro-batch 僅包含代碼數據,micro-batch 層級的損失會迫使 router 將這些 code tokens 均勻地分配給所有專家。這阻礙了專家在特定領域(如 code)的專業化,並可能損害模型的整體性能。由於在 LLM 訓練中,單個 micro-batch 中的數據通常來自同一個領域,這解釋了為什麼許多開源 MoE 模型缺乏顯著的專家專業化。
實現 Global-Batch 負載平衡
為了修復這個問題,Qwen 提出在 global-batch 層級計算負載平衡損失 ($L_{\text{global}}$)。這是透過以下三個步驟實現的:
- 同步所有並行組 (parallel groups) 的專家選擇頻率 ($f_{i}$)。
- 計算每個並行組中的負載平衡損失(例如,在單個 GPU 上)。
- 聚合所有 micro-batches 的損失。
由於專家選擇頻率是一個很小的向量(維度等於專家數量),在 micro-batches 之間同步這些數據的計算成本非常低,使得該實現幾乎是「免費」的。
性能增益與專家專業化
Qwen 在三種 MoE 配置(3.4B 激活 0.6B、15B 激活 2.5B、43B 激活 6.6B)和兩種數據配置(120B 和 400B tokens)下測試了該方法。
提升模型效能
與 micro-batch 層級的損失相比,global batch 方法在所有測試設置中都取得了更好的性能,包括所有模型規模、數據量和任務。
增強領域專業化
Global batch 平衡使專家能夠實現專業化。雖然 micro-batch 平衡會導致專家無論領域如何都被均勻激活,但 global batch 平衡允許特定專家被特定領域頻繁激活,展現出明顯的專業化。
平衡 Batch Size 的影響
對 3.4B 模型(激活 0.6B)的實驗顯示,隨著 balance batch size 從 2 增加到 128,ptr-training PPL (perplexity) 迅速下降,並在 128 處達到飽和。這突顯了 global 方法的重要性,因為主流的 MoE 框架對於較大型模型通常使用 8 到 16 之間的 balance batch size。
平衡效率與效能
雖然 global batch 平衡可以提高模型質量,但它可能會導致 micro-batch 平衡的退化,進而對計算效率產生負面影響。
為了緩解這一問題,Qwen 嘗試在 global-batch 平衡損失之上添加 micro-batch 平衡損失(使用 global batch 損失 0.01 的常數權重)。這種混合方法在幾乎不影響模型效能的情況下,提高了更新步長速度(從每次更新步長 1.64 秒提高到 1.59 秒)。
結論
透過實施 global-batch 負載平衡,Qwen 解決了 MoE 訓練中的一個關鍵限制,從而實現了性能更強大且更專業的專家。這種方法為優化 MoE 模型提供了新的視角,特別是在跨多個領域擴展至更龐大且專業的 MoE 模型的能力方面。