FlashOptim:面向内存高效训练的优化器

FlashOptim: Optimizers for Memory Efficient Training

arXiv: 2602.23349v1

论文信息

标题: FlashOptim: Optimizers for Memory Efficient Training

作者: Jose Javier Gonzalez Ortiz, Abhay Gupta, Chris Renard, et al.

发布日期: 2026-02-26

arXiv ID: 2602.23349v1

PDF 链接: 下载 PDF

3 分钟速览

  • 研究问题:在大模型训练中,每个参数不仅要存储模型权重,还要存储梯度与优化器状态,普通 AdamW 每参数要占用约 16 字节显存,这使 7B 模型的训练至少需要 112 GB 加速器内存,严重限制了资源有限的研发者。
  • 核心方法:提出 FlashOptim,包含两项关键技术——改进的权重复分裂(基于单位最后位置的误差校正,实现 24 位有效精度)和带压缩扩展的优化器状态 8 位量化,将 AdamW 的每参数存储降到 7 字节(如再释放梯度则降至 5 字节)。
  • 关键结果:在 Llama-3.1-8B 微调中,FlashOptim 将峰值显存从 175 GiB 降至 113 GiB(降低约 36%),而在多种视觉和语言基准上收敛与精度无退化。
  • 主要局限:该方法主要针对参数相关内存,对于激活内存占主导的浅层卷积网络收益有限;某些极端数据分布仍可能对量化敏感,需要谨慎处理。
  • 适合读者:从事大型神经网络训练、受显存限制的研究人员和工程师,特别是需要复用标准优化器语义但希望节省显存的用户。

论文背景和研究动机

现代深度学习的进步很大程度上依赖模型规模的扩大,但随之而来的是巨大的加速器内存需求。在典型的混合精度训练中,每个参数不仅要存储 4 字节的 FP32 主权重,还要存储 FP32 的梯度、一阶动量和二阶方差(Adam 系列),总体需要约 16 字节(表 1)。对于一个 70 亿参数的大语言模型,仅参数相关内存就已超过 112 GB,远超许多研究者的可用硬件上限。

已有的缓解手段各有不足:分布式张量切分(如 ZeRO)需要多个加速器,离线策略(CPU offloading)会引入额外通信与复杂度,参数高效微调(如 LoRA)改变了训练动态,可能牺牲表达能力。因此,论文的目标是在不改变优化器语义、不增加训练负担的前提下,显著压缩参数相关内存。

FlashOptim 正是一套面向 SGD、AdamW、Lion 等常用优化器的压缩方案,通过权重复分裂和优化器状态压缩两个核心手段,将半精度形式的主权重和低精度误差项紧密耦合,并对动量/方差进行非线性压缩扩展后再做线性量化,最终以极小的精度损失换取超过 50% 的显存下降。

核心方法和技术细节

权重复分裂:基于单位最后位置的误差校正

在混合精度训练中,通常用 FP32 维护主权重,但前向、反向传播使用 FP16/BF16 权重。这实际上存储了冗余信息:FP16 权重可以通过向下取整从 FP32 得到,而误差信息被浪费。权重复分裂的思想是将 FP32 主权重分解为低精度权重 θ′\theta' 和误差修正项 ρ\rho,从 (θ′,ρ)( \theta', \rho ) 恢复出高精度的 θ^\hat{\theta}。

论文分析指出,传统直接存储 θ−θ′\theta - \theta' 为 BF16 的做浪费了指数位,因为 θ\theta 一定落在 [θ′−ULP(θ′)/2,θ′+ULP(θ′)/2][\theta' - \mathrm{ULP}(\theta')/2, \theta' + \mathrm{ULP}(\theta')/2] 区间内。作者利用这一性质,将误差按 ULP(θ′)/2\mathrm{ULP}(\theta')/2 归一化后均匀量化为 bb 位整数(式 1、2)。例如,取 θ′\theta' 为 BF16,ρ\rho 为 INT8,则相当于 24 位有效精度(PXR24 格式)。算法 1 给出了压缩与解压的细节,包括数值稳定性的处理。

该方法与现有 BF16+BF16 方案相比,重建误差大幅降低:在 BF16 基础权重的场景下,用 16 位误差修正可实现 99.92%99.92\% 的比特级完美重建(见第 4.4 节,图 3),且在 FP16 的数值范围内误差恒定且远小于现有方案。

带压缩扩展的优化器状态量化

优化器中的动量和方差通常遵循重尾分布,直接做组内线性量化会产生较大误差。FlashOptim 引入了可逆、无超参的压缩扩展函数,重塑分布使其更接近均匀,然后再进行组内缩放和整数量化。

对动量 mm 使用类似 softsign 的变换:

ϕm(x)=2x1+∣x∣,ϕm−1(z)=z2−∣z∣\phi_m(x) = \frac{2x}{1+|x|}, \quad \phi_m^{-1}(z) = \frac{z}{2-|z|}

对 Adam 的二阶矩 vv,先取平方根再量化:

ϕv(x)=x,ϕv−1(z)=z2\phi_v(x) = \sqrt{x}, \quad \phi_v^{-1}(z) = z^2

量化时采用组大小 G=32G=32,每组分别存储一个 FP16 缩放因子。动量存储为 INT8,方差存储为 UINT8,分别见算法 2 和算法 3。这种简单的变换不仅显著降低了量化误差(图 4),还能防止训练发散(图 5)。

融合优化器步骤

完整的 FlashAdamW 流程(算法 4)在每次更新前先反量化动量/方差、重建主权重,执行标准的 AdamW 更新,再将新状态量化并重新分裂权重。所有压缩、解压和更新操作都融合在一个 Triton 内核中执行,避免了额外的内存访问,保持与标准优化器相当的吞吐。FlashOptim 还支持梯度释放,在不需要梯度累积时提前释放梯度内存,进一步节省 2 字节/参数。

创新点和贡献

论文的创新之处体现在以下几点:

  1. 改进的浮点误差校正:利用 ULP 归一化,将误差限制在半个末位单位内,用整数精确编解码,以极低成本实现接近 FP32 的精度。这比之前直接存储 BF16 误差的方案有数量级更小的重建误差(见图 3),且无需特殊数据格式。

  2. 精简且有效的压缩扩展函数:仅通过一行的非线性变换(动量的 softsign 式函数和方差的平方根),就使普通的组内线性量化在优化器状态上达到与复杂量化方案相当的质量,同时避免了训练发散。这为其他类型张量的压缩提供了设计思路。

  3. 系统级设计:将两项技术无缝嵌入标准优化器,提供 drop-in 替换,无需超参调优。FlashOptim 可与 FSDP、激活检查点等正交组合,适合多种分布式训练场景。开源实现便于社区直接使用。

  4. 全面的验证:在图像分类(ResNet-50)、语言模型预训练(GPT-2)、大语言模型微调(Llama-3.1-8B)任务上,搭配 SGD、AdamW、Lion 三种优化器,均展现了与参考实现无统计差异的收敛性和最终精度(表 2、表 3、图 2),而内存收益显著。

实验结果分析

实验部分从收敛性、内存/速度、权重复重建误差和优化器状态量化误差四个方面验证了 FlashOptim 的有效性。

首先,收敛与精度:GPT-2 预训练损失曲线(图 2a)显示 FlashAdamW 与参考 AdamW 几乎重合,直至 20,000 步更新没有偏移;ResNet-50 训练也类似(图 2b)。在 ImageNet Top-1 准确率上,FlashSGD(77.16%) 与 FlashAdamW(75.67%) 均不低于参考优化器(表 2)。对于 LLM 微调的 GSM8k 准确率,FlashAdamW(74.98%)也与参考(75.09%)处于同样方差范围内。预训练后的常识推理平均分(Mean ICL)也无明显差异(表 3)。这些结果表明压缩对模型学习动态无负面影响。

其次,内存与速度:在 Llama-3.1-8B 微调分析中(表 4),FlashOptim 将参数内存从 29.9 GiB 降至 15.0 GiB(-50%),优化器状态从 59.8 GiB 降至 23.4 GiB(-61%),加上激活后的峰值内存从 175.2 GiB 降至 112.9 GiB,下降 36%。优化器步时间仅从 12.5 ms 降至 11.5 ms(实际略有加快,因为内核融合),没有引入明显开销。消融实验表明权重复分裂单独可使参数内存减半,但会略微增加优化器内存(误差修正项),而优化器状态量化可带来约 73% 的状态压缩。

权重复重建误差:图 3 的穷举测试显示,即使使用 8-bit 误差修正(24 位格式),对于 BF16 基础权重,相对误差已控制在实际应用可接受的范围;而 16-bit 修正(BF16+FP32 重建)几乎完美,误差 <10−9<10^{-9}。相比 BF16+BF16 方案,ULP 方法平均和最大误差均有数量级改善。

优化器状态量化误差与稳定性:图 4 的 NMSE 分布表明,companding 对动量有一定的改善,而对方差的改善极为明显,可将大规模误差尾部大幅压窄。GPT-2 训练的消融实验(图 5)证实,若直接对方差做线性量化(不 companding),训练会迅速发散,而 companding 确保训练稳定收敛。这从实践角度强调了压缩扩展函数的不可或缺性。

实践建议

FlashOptim 最适合参数规模大、优化器内存占主导的场景,如大语言模型的预训练和全参数微调。实际部署时,可以参考以下几点:

集成方式:使用论文开源的 PyTorch 库,仅需将标准优化器替换为 FlashSGD、FlashAdamW 或 FlashLion,无需修改训练超参。由于内核融合,通常在 H100 等 GPU 上不会带来额外时间开销,甚至可能因为内存带宽压力减小而略有加速。

与其他优化组合:FlashOptim 与 FSDP、激活重计算、CPU offload 等完全正交。可以优先启用 FSDP 将模型分片,再在每个设备上应用 FlashOptim,进一步降低单卡内存。如果激活性内存显著(如大 batch 高分辨率图像),建议仍保留激活检查点。

选择性使用:若某些层对量化特别敏感,实现支持按层禁用压缩或阻止量化。可先全量启用,观察损失和评估指标,若出现不稳定再局部回退。

梯度释放:当不需要累积梯度(单一 batch 更新)时,强烈建议打开梯度释放选项,配合 FlashOptim 可将每参数额外减少 2 字节,对超大模型收益显著(如 7B 模型节省约 14 GB)。

检查点优化:训练过程中的模型保存也受益:Standard Adam 检查点每参数需要 12 字节,FlashAdamW 仅需 5 字节,对于 70B 模型,单个检查点文件从 840 GB 降至 350 GB。

监控与调试:初次使用时,可在训练开始时记录优化器步时间和峰值内存,与参考实现对比。若发现精度下降,首先检查方差量化是否引发不稳定,可临时对方差使用更高精度(或不量化),并审视线学习率、权重衰减在低精度下的等效行为。

适应性评估:虽然论文在常见视觉和语言基准上验证了安全性,但面对全新的数据分布或极端训练设置(如超大 tensor 尺寸、非常小的 batch),宜进行小规模前测,确保压缩不影响最终指标。

FlashOptim 作为一款低风险、高收益的内存优化方案,适合希望在不增加硬件投入的情况下训练更大模型的研究团队,也为其他内存紧凑的训练工具体系提供了简洁的设计参考。