使用 Megatron-LM 訓練語言模型

概述

Megatron-LM 是由 NVIDIA 的 Applied Deep Learning Research 團隊開發的高效能框架,旨在於 GPU 上高效預訓練大型 transformer 模型。雖然相較於 Hugging Face Trainer API 或 accelerate 套件,設定上較為複雜,但 Megatron-LM 透過專門的資料處理與硬體層面的最佳化,提供顯著的加速效果。

Megatron-LM 的技術最佳化

Megatron-LM 透過兩項主要機制——最佳化的資料載入與 kernel 融合,達成相較於標準 PyTorch 迴圈更優的訓練效率。

高效的 DataLoader

Megatron-LM 使用一個在訓練開始前即完成斷詞與資料洗牌的 DataLoader。它會一次計算帶索引的序列編號,並根據訓練參數決定 epoch 數量。此方式避免在每個新 epoch 前必須遍歷整個資料集,從而使學習曲線更平滑、訓練時間縮短。

融合的 CUDA Kernels

為了降低記憶體開銷,Megatron-LM 採用融合的 CUDA kernels。標準的 PyTorch 操作會在每個獨立運算中從記憶體讀取資料、計算後再寫回。Kernel 融合將相似的運算合併為單一硬體操作,使中間結果能保留在 GPU 暫存器中,而不必回寫至主記憶體。

此外,Megatron-LM 使用來自 NVIDIA Apex 套件的融合版 AdamW 實作,其效能較標準的 PyTorch 實作更快。

實作工作流程

在 Megatron-LM 中訓練模型(例如 GPT-2)包含四個步驟的流程:環境設定、資料前處理、訓練以及模型轉換。

環境設定

建議的設定方式是使用 NGC 提供的 NVIDIA PyTorch Container,該容器已包含所有必要的相依套件。若自行安裝,則需確保 PyTorch、CUDA、NCCL、NVIDIA Apex 以及 nltk 套件皆為最新版本。使用者還必須在 Megatron-LM 目錄中提供斷詞器的 vocab.jsonmerges.txt 檔案。

資料前處理

資料在處理前必須先轉換為寬鬆的 JSON 格式(每行一筆樣本)。接著 Megatron-LM 會使用 tools/preprocess_data.py 進行斷詞、洗牌,並將資料轉換為二進位格式。此過程會產生 .idx.bin 檔案,這些檔案是訓練階段所必需的。dataset-impl 可設定為 'lazy'、'cached' 或 'mmap'。

訓練執行

訓練透過 torch.distributed.launch 啟動。以一個 110M 參數的模型(如 CodeParrot-small)為例,在 8 顆 GPU 上的訓練時間約為 12 小時,使用以下設定:

  • Architecture: 12 層、768 隱藏單元大小、12 個注意力頭。
  • Sequence Length: 1024。
  • Optimizer: 使用 AdamW 並搭配餘弦學習率衰減。
  • Batch Size: 微批次大小為 12,整體批次大小為 192。

模型平行化策略

對於無法容納於單一 GPU 的大型模型,Megatron-LM 支援兩種模型平行化方式:

  1. Tensor Parallelism: 透過 tensor-model-parallel-size 參數,將單一 transformer 模組的執行分散至多個 GPU。
  2. Pipeline Parallelism: 透過 pipeline-model-parallel-size 參數,將 transformer 模組切分為等大小的階段,分配至多個 GPU。

與 Hugging Face Transformers 的整合

若要將 Megatron-LM 訓練的模型用於評估或上線,必須將其轉換為 transformers 套件支援的格式。這可透過使用 convert_megatron_gpt2_checkpoint.py 腳本,將 model_optim_rng.pt 檢查點檔案轉換為 pytorch_model.bin 檔案來完成。

轉換完成後,可使用 AutoModelForCausalLM 載入模型。對於使用模型平行化訓練的極大型模型,可透過 device_map="auto" 參數,結合 accelerate 套件,自動將權重分配至可用的 GPU 與 CPU 記憶體。

框架選擇指引

由於具備高度最佳化,Megatron-LM 最適合用於大型模型的預訓練或長時間的微調。然而,它在前處理與模型轉換階段會產生額外的開銷。若僅需對中等規模模型進行較短的微調,建議使用 Hugging Face Trainer API 與 accelerate 套件,因為它們與裝置無關且具更高的彈性。

Sources