AI News字數 3103閱讀時長8 分鐘

Marin 開始公開訓練 535B 參數 MoE 模型

Marin 已開始公開進行一項使用 18.75 萬億個 token 的 535B 參數 MoE 訓練,並公開程式碼、日誌及擴展預測。

一、一項公開進行的 535B 參數訓練

Marin 已開始訓練 Marin 535B-A23B,這是一個混合專家語言模型,總參數約為 5,350 億,每個 token 啟用 230 億個參數。該項目計劃在約三個月內,使用 11 套 NVIDIA GB200 NVL72 系統處理 18.75 萬億個 token。

已公布的計劃將 80% 的 token 預算分配予預訓練,20% 分配予中期訓練。其後將進行後訓練,但 Marin 尚未公布最終的後訓練方案或時間表。該項目估計,主要訓練將需要約 \(2.7 \times 10^{24}\) 次浮點運算。

每套 GB200 NVL72 均為機櫃級系統,內含 72 枚 Blackwell GPU 及 36 枚 Grace CPU。11 套這類系統合共構成 792 枚 GPU 的硬件規模,儘管 Marin 的專家並行實作報告描述的是每個訓練機櫃內一個由 64 枚 GPU 組成的專家並行域。

這次訓練並非模型發布。截至 8 月 24 日,訓練仍在進行,因此尚未有最終權重、基準測試結果或後訓練評估可供判斷。其當前的重要性,在於選擇在這種規模的訓練進行期間公開相關資訊,而非只在挑選出成功結果後才發布技術報告。

Marin 為這次訓練設立的公開 issue 於 8 月 18 日開啟。其中載有運作計劃、工程風險、擴展方法、延長上下文方案及應變程序。連結的 Weights & Biases 報告提供了該項目的即時追蹤介面。

公告列明為 18.75 萬億個 token,而 GitHub issue 標題則將這次訓練簡稱為「18T tokens」。較精確的數字見於公開 voyage 計劃;該項目並未將較短的標題表述為修訂後的預算。

二、稀疏模型的組織方式

Marin 535B-A23B 使用 48 個 transformer block。每個 block 都結合注意力分支與一個稀疏 MoE 分支,後者包含 384 個經路由選取的專家。路由器會為每個 token 選取八個專家,同時有兩個共享專家保留於本地,並獨立於路由路徑處理每個 token。

模型狀態的寬度為 6,144 個值。在經路由的啟用值於 GPU 之間交換前,潛在投影會將其壓縮至 3,072 個值。Marin 表示,這會將穿過專家並行 all-to-all 操作的啟用值流量寬度減半。輸出會投影回模型寬度,然後才與共享專家路徑結合。

這個傳輸問題相當重大。384 個經路由專家分布於一個由 64 枚 GPU 組成的專家並行域,每枚 GPU 承載六個經路由專家。若為每個專家傳送獨立調整大小的 buffer,將帶來棘手的記憶體需求及動態通訊模式。

Marin 因此為 JAX 及 XLA 開發了固定 pooled-wave all-to-all 實作。發送端會為每個目標 GPU 建立一個固定 pool,而非為每個專家建立一個 buffer。傳輸分三個連續 wave 進行,並使用相同的 array shape;專家識別碼則封裝在啟用值 payload 內。這避免了另行交換 token 數量或路由 metadata。

該實作採用兩項容量限制。發送端容量因子為 1.10,用以限制由一個來源發送至一個目標的流量;接收端容量因子為 1.15,限制分配予單一本地專家的 row 數量。任何超出任一固定 buffer 的分配都會被捨棄,並另行記錄。

一項使用一個機櫃、歷時 20 個 step 的 gate 已完成,沒有發生記憶體不足故障,並在第 2 至第 19 個 step 錄得每秒 250,691 個 token 的中位吞吐量。Marin 明確警告,這項短測試並不能確立最終的 token 捨棄率。這是工程資格驗證結果,而非訓練品質基準,也不是對完整 11 個機櫃訓練的量度。

該報告亦記錄了失敗的配置。直接固定 expert-cell 設計得出的 XLA 記憶體估算為 192.65 GiB,並在一項 123.49 GiB CUDA 配置時失敗。一個六專家接收端 bank 亦出現記憶體不足,而將工作拆分為三個 wave,則令該 bank 每次只維持部分處於啟用狀態。這些負面結果亦是該項目公開設計紀錄的一部分。

三、主訓練前先建立擴展階梯

在啟動 535B 模型前,Marin 訓練了一個四級擴展階梯。它由一個總參數為 16 億、啟用 6,100 萬個參數,並以 480 億個 token 訓練的 MoE 開始。最大的級別包含 277 億個總參數、啟用 12 億個參數,並處理了 9,260 億個 token。

這個階梯同時用作預測及診斷參考。Marin 可將主訓練的 loss、gradient norm、token 捨棄及評估走勢,與較小規模下觀察到的模式比較。若出現實質偏離,便可在完整訓練耗用數月運算資源前啟動調查。

根據該項目,這個階梯的成本約為主訓練運算量的 1%。Marin 將這項開支視為一種風險控制:測試所選架構、資料混合、optimizer 設定及訓練範圍,會否隨模型規模增加而保持一致的行為。

該項目表示,較早的階梯曾發現,隨著 token 範圍擴展,gradient norm 會增長至四以上。這項發現促成採用 logit z-loss。其後的消融研究據報顯示,部分高 batch 配置否則可能會在訓練期間發散。

較小規模的訓練並不能保證 535B 模型會依循其預測。外推仍是這項實驗其中一項核心不確定性。其實際價值在於,Marin 已公布一個基線,讓外界可據此評估在約 100 日訓練期間作出的介入。

Marin 亦記錄了一項針對基礎設施延誤或模型 FLOP 利用率低於預期的應變方案。在 token 預算約首四分之一期間,預設回應是縮短 token 範圍、調整資料混合,並重新安排線性學習率衰減,使其仍能在修訂後的終點降至峰值學習率的 5%。這意味著所公布的時間表仍是一項運作計劃,而非不可變更的規格。

四、長上下文取決於解決 token 捨棄問題

模型以 4,096-token 序列長度開始預訓練。Marin 先前的大型訓練由 8K 開始,延伸至 65K 並持續一萬億個 token,並已安排稍後延伸至 262K。回到 4K 可令每個 batch 包含的獨立序列數量較 8K 多一倍,理應使 token 在專家之間分布得更平均。

這項選擇針對現時專家並行實作的一個弱點。在較早測試中,token 捨棄率由 4K 上下文的約 7%,升至 65K 的約 40%。Marin 表示,其較新的 pooled-wave 設計在 4K 時約捨棄 3%,但預期在 65K 時比率仍可能過高。

被捨棄的分配不一定表示整個 token 消失。兩個共享專家仍會處理每個 token,而八個被選取的經路由專家是共享路徑以外的額外部分。Marin 表示,當經路由分配超出可用容量時,共享專家提供了較密集的骨幹;但高捨棄率仍可能減少稀疏專家的效益,並改變訓練行為。

該項目計劃在訓練開始約 10 至 20 日後,進行一次為期一至兩日的早期 cooldown。此分支旨在為強化學習實驗提供一個完整規模的 checkpoint,並測試較長上下文如何影響路由,同時不改變主要訓練走勢。

若該實驗穩定,暫定的上下文時間表將在訓練中段由 4K 升至 8K,在約 95% 進度時由 8K 升至 65K,然後在接近尾聲時進入目標為 262K 的階段。這些均屬有條件的目標,而非已完成模型的確認能力。

若 token 捨棄率仍然過高,Marin 列出三個替代方案:採用無捨棄的 ragged all-to-all 實作、提高容量因子並承擔隨之而來的記憶體成本,或引入序列層級平衡。最後一項方案可能迫使專家在訓練期間重新專門化,因此該項目目前將其視為後備方案。

五、何以這次訓練是公開開發的里程碑

Marin 區分公開開發與訓練後發布權重。其標準工作流程由一個記錄實驗假設及目標的 GitHub issue 開始。實作會以可供審閱的程式碼提交,執行過程會連結至公開遙測資料,而分析結果——包括失敗嘗試——則會回寫至 issue。

就 535B 訓練而言,公開紀錄已包括 hero-run issue、即時追蹤報告、原始碼儲存庫,以及專家傳輸實作的詳細說明。傳輸報告把量度結果與工程判斷分開,並指出哪些結論只建基於短期 profiling 訓練。

這份紀錄讓研究人員可評估擴展預測能否在完整模型上成立。它亦公開了一些最終模型卡通常只會濃縮成數行的決定:為何訓練由 4K 上下文開始、token 分配在何處被捨棄、記憶體限制如何改變傳輸設計,以及團隊若訓練進度落後會如何處理。

公開文件不會令主實驗變得便宜易於重現。重現一項使用 11 個機櫃的訓練,仍超出大多數獨立研究人員的資源範圍。較易取得的產物是程式碼、較小的擴展訓練、失敗分析、配置選擇及可供檢視或以縮小規模測試的即時量度資料。

因此,該項目的狀態必須準確而有限地描述:Marin 已開始一項有公開文件記錄、前沿規模的訓練。它尚未展示最終模型的品質、長上下文行為、後訓練表現或最終產物發布。

常見問題

「535B-A23B」是甚麼意思?

該模型總參數約為 5,350 億,而每個 token 約有 230 億個參數處於啟用狀態。稀疏路由會在每個 transformer block 的 384 個經路由專家中選取八個。

現在可以下載完成後的模型嗎?

不可以。訓練仍在進行,Marin 尚未發布這次訓練的最終權重或後訓練評估。

為何訓練只以 4K 上下文開始?

4K 序列長度可在每個 batch 放入更多獨立序列,改善專家平衡。Marin 正在測試較長上下文,因為較早的實作在 65K 時出現明顯更高的 token 捨棄率。

研究人員如何追蹤這次訓練?

Marin 提供公開 GitHub issue、原始碼儲存庫、專家並行設計報告,以及連結的 Weights & Biases 追蹤頁面。

262K 上下文長度是否有保證?

沒有。擬議延伸至 8K、65K 及 262K,取決於早期 cooldown 及 token 捨棄實驗的結果。

參考來源

Share

分享這篇文章