当前位置: 首页 > news >正文

一文读懂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框架主要包含两个关键部分:

  1. 缩放随机递归动量:提供全梯度的方差减少估计器,以获得更好的梯度复杂度。
  2. 预处理更新:近似二阶牛顿法,以获得更好的每迭代复杂度。

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),仅供参考

http://www.cnnetsun.cn/news/3649590.html

相关文章:

  • 计算机Python毕设实战-基于 Python Web 的学生日常考勤信息系统 班级学生考勤登记与异常报备系统设计【完整源码+LW+部署说明+演示视频,全bao一条龙等】
  • LZHAM新手入门:从安装到压缩第一个文件的完整教程
  • Kimi K3编程实力远超GLM5.2:一个Trae复赛证据
  • 解锁SwiftUI 5新特性:Metal Shader Collection中的滚动增强与视觉效果
  • 【管理科学】第五十六篇 企业管理层的权-责-利益分析及权利-人性-资源限制分析01
  • Netty在HuLa-Server中的应用:高性能WebSocket连接管理与消息推送
  • 5分钟掌握OBS专业虚拟背景:零绿幕AI抠图完全指南
  • Buzz项目管理:高效管理平台开发的方法与工具
  • Java后端面试7天冲刺:从八股文到实战的系统复习指南
  • Qwest兼容性处理:IE8+与现代浏览器适配方案
  • 从AI高考数学148分看大模型推理能力:原理、瓶颈与编程实战指南
  • n8n与RAG技术在钉钉机器人中的智能客服应用
  • GitHub_Trending/cla/claude-skills异常检测技能:识别系统异常行为的终极指南
  • AI如何变革学术专著创作:工具链与效率提升实战
  • 孪生网络原理与应用:从相似性度量到工业实践
  • PyOfficeRobot定时群发攻略:解放双手的微信营销神器
  • Node.js网站下载器环境配置与实战指南
  • TkinterMapView核心功能详解:标记、多边形与路径绘制的完整教程
  • Splunk Attack Data高级技巧:选择性拉取数据集节省90%存储空间
  • gh_mirrors/core109/core安全最佳实践:保护用户数据与API接口的终极指南
  • SpatialClaw代码接口:空间智能体的精确控制与工程实践
  • 10分钟上手Gemini-OpenAI-Proxy:开发者必看的API调用示例与参数说明
  • 5分钟掌握专业网络测速:iperf3 Windows版终极指南
  • SpERT完全指南:Span-based Entity and Relation Transformer如何彻底改变实体关系抽取
  • AI内容生产革命(豆包×剪映深度耦合实战手册):实测效率提升417%,92%新手3天即达专业级交付水准
  • 网络安全从零开始学习CTF——CTF基本概念
  • 从甲骨文到数字孪生:AI驱动的历史记忆范式革命(全球首份跨文明记忆强度对比报告首发)
  • Unity2D拖尾渲染器性能优化全攻略:从原理到实战解决卡顿与渲染问题
  • 解决AirplaneJS常见问题:设备连接失败、信号弱与地图加载异常处理方案
  • tinker-manager常见问题解答:新手入门必看