GaLore:在消費級硬體上推進大型模型訓練
GaLore 透過大幅降低優化器狀態的記憶體占用,使得在消費級硬體(例如 NVIDIA RTX 4090)上訓練最高可達 70 億參數的大型語言模型(LLM)成為可能。此舉讓 AI 研究更加民主化,使實踐者無需高端工業級計算資源即可訓練大規模模型。
透過低秩梯度投影提升記憶體效率
GaLore 透過在梯度送入優化器之前,將其投影至較低維度的子空間,以降低記憶體消耗。此方法利用深度神經網路中梯度天然的低秩結構,將訓練過程中需儲存與操作的資料量降至最小。
對於像 Adam 這樣的自適應優化演算法,優化器狀態通常佔據記憶體的大部分。透過此投影,GaLore 在訓練期間報告稱可將儲存優化器狀態的記憶體需求降低超過 82.5%。
動態子空間切換
為了保有完整參數學習的能力,並避免模型被限制在參數空間的某一小部分,GaLore 採用動態子空間切換機制。此機制使模型在訓練過程中能在不同的低秩子空間之間切換。
切換的頻率經過平衡,以確保優化路徑的一致性,同時適應梯度不斷演變的低秩結構。這使得在記憶體效率與優化效能之間的取捨上能進行細緻的控制。
與 8 位元優化器的整合
將 GaLore 與 8 位元精度的優化器結合,可透過量化優化器狀態進一步提升記憶體效率。此協同效應使得在相同硬體限制下,能訓練更大的模型或使用更大的批次大小,而不會影響模型的準確度或收斂速度。
GaLore 搭配 8 位元優化的演算法流程
- 梯度投影:全精度梯度使用投影矩陣投射至低秩子空間,然後量化為 8 位元格式。
- 量化:投影後的梯度、模型權重以及優化器狀態(例如 Adam 的動量平均)皆從 32 位元浮點數量化為 8 位元整數表示。
- 優化器更新:對 8 位元量化的梯度進行更新;此過程包括將梯度去量化回浮點數、套用更新規則,並將更新後的優化器狀態重新量化為 8 位元。
- 去量化與權重更新:將權重去量化為浮點數以進行處理。GaLore 隨後使用最終投影,將去量化的低秩更新映射回原始參數空間,然後再套用權重更新。
與 Hugging Face Transformers 的實作
GaLore 已整合至 Hugging Face 的 transformers 套件(版本 4.39.0 或以上)以及 galore-torch 套件。使用者可在 TrainingArguments 中指定優化器,例如 galore_adamw、galore_adamw_8bit 或 galore_adafactor,並透過 optim_target_modules 定義目標模組,以實作 GaLore。
分層更新
為了進一步降低記憶體占用,GaLore 支援分層權重更新。優化器不再於反向傳播後一次性更新所有層,而是使用 PyTorch 的後累加鉤子逐層更新權重。此功能可透過在優化器名稱後加上 _layerwise 來啟用(例如 galore_adamw_layerwise)。