一文读懂MARS框架:为什么方差 reduction 是大模型训练的关键?
一文读懂MARS框架:为什么方差 reduction 是大模型训练的关键?
【免费下载链接】MARSThe official implementation of MARS: Unleashing the Power of Variance Reduction for Training Large Models项目地址: https://gitcode.com/gh_mirrors/mars11/MARS
MARS(Make vAriance Reduction Shine)是一个专为解决大模型训练挑战设计的统一优化框架。传统自适应梯度方法如Adam和AdamW常受高随机梯度方差困扰,而方差减少技术在深度学习中一直难以获得实际影响。MARS通过结合预处理梯度方法与方差减少技术,实现了两者的优势,加速了优化中临界点的搜索。
为什么方差 reduction 对大模型训练至关重要?
在大模型训练过程中,随机梯度的高方差会导致训练不稳定、收敛速度慢以及最终性能不佳。方差 reduction 技术通过降低梯度估计的波动性,能够有效改善这些问题,使模型更快收敛到更好的解。MARS框架正是围绕这一核心思想构建,旨在充分释放方差 reduction 在大模型训练中的潜力。
MARS框架的核心组件
MARS框架主要包含两个关键部分:
- 缩放随机递归动量:提供全梯度的方差减少估计器,以获得更好的梯度复杂度。
- 预处理更新:近似二阶牛顿法,以获得更好的每迭代复杂度。
MARS的三种实例化方式
在MARS框架下,基于不同的Hessian矩阵近似,提供了三种实例化方式:
MARS-AdamW
(通过在mars.py中设置mars_type="mars-adamw"启用)
Hessian矩阵近似定义为: $$\mathbf{v}t =\beta_2 \mathbf{v}{t-1}+(1-\beta_2) \big(\nabla f(\mathbf{x}_t, \mathbf{\xi}_t)\big)^2$$ $$\mathbf{H}_t := \sqrt{\text{diag}\Big(\mathbf{v}_t\Big)}\cdot \frac{1 - \beta_1^t}{\sqrt{1 - \beta_2^t}}$$
MARS-Lion
(通过在mars.py中设置mars_type="mars-lion"启用)
Hessian矩阵近似定义为: $$\mathbf{H}_t := \sqrt{\text{diag}(\mathbf{m}_t^2)}$$
MARS-Shampoo
(通过在mars.py中设置mars_type="mars-shampoo"启用)
预处理器可视为正交映射算子: $$\mathbf{U}_t, \mathbf{\Sigma}_t, \mathbf{V}_t = \text{SVD}(\mathbf{G}t),\qquad \mathbf{x}{t+1} =\mathbf{x}_t-\eta_t\mathbf{U}_t\mathbf{V}_t^\top$$
MARS的性能表现
在OpenWebText上的实验结果
MARS在各种GPT-2模型上始终优于AdamW和Muon优化器。以下是GPT-2 large模型在OpenWebText数据集上的验证损失对比:
从图中可以看出,MARS(红色和绿色曲线)的验证损失明显低于AdamW(黄色曲线)和Muon(蓝色曲线),尤其是在训练后期,差距更加明显。
在CIFAR-10上的实验结果
MARS在视觉任务上也表现出色。在CIFAR-10数据集上,MARS的测试准确率显著高于AdamW和Muon:
红色曲线代表MARS,绿色曲线代表AdamW,蓝色曲线代表Muon。可以看到,MARS不仅收敛速度更快,而且最终的测试准确率也最高。
MARS的效率优势
MARS算法不仅在相同训练步数内表现更好,而且在相同训练时间内也能取得更优结果:
图中展示了GPT-2 large模型在32xH100上的验证损失随时间变化情况。红色曲线代表MARS,绿色曲线代表AdamW,蓝色曲线代表Muon。MARS在相同时间内能够达到更低的验证损失,证明了其高效性。
如何开始使用MARS?
安装依赖
$ pip install torch==2.1.2 transformers==4.33.0 datasets tiktoken numpy==1.26.4 wandb获取代码
$ git clone https://gitcode.com/gh_mirrors/mars11/MARS数据准备
按照nanoGPT的方法准备OpenWebText数据:
$ python data/openwebtext/prepare.py开始训练
要使用MARS优化器训练模型,运行以下命令:
$ torchrun --standalone --nproc_per_node=8 MARS/train_mars.py config/${your_config_file}此命令使用MARS优化器在OpenWebText数据集上启动GPT-2模型的训练。所有相关超参数(训练、模型和优化器)都在配置文件(${your_config_file})中指定。这些参数可以直接在配置文件中调整,也可以通过bash脚本调整。
总结
MARS框架通过创新性地结合方差 reduction 技术和预处理梯度方法,为大模型训练提供了一个高效、稳定的优化解决方案。无论是在语言模型还是视觉任务上,MARS都展现出了优异的性能和效率。如果你正在从事大模型训练相关工作,不妨尝试MARS框架,体验方差 reduction 带来的训练加速和性能提升。
通过合理设置MARS的超参数,特别是学习率,你可以进一步优化模型性能。MARS的灵活性和强大性能使其成为大模型训练的理想选择,值得在各种深度学习任务中尝试和应用。
【免费下载链接】MARSThe official implementation of MARS: Unleashing the Power of Variance Reduction for Training Large Models项目地址: https://gitcode.com/gh_mirrors/mars11/MARS
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
