[Databricks] FlashOptim:重新定义显存效率,让 8B 模型训练内存直降 50%
Spa3R: Predictive Spatial Field Modeling for 3D Visual Reasoning
本文推出了 FlashOptim,一套针对深度学习优化器的内存优化技术方案。通过结合改进的权重拆分(Weight Splitting)和基于压扩函数(Companding)的 8-bit 状态量化,FlashOptim 在保持模型质量和 API 兼容性的前提下,将 AdamW 的每参数内存开销从 16 字节降低至 7 字节(结合梯度释放可达 5 字节),并在 Llama-3.1-8B 等大规模任务上实现了 SOTA 性能的无损保持。
TL;DR
在大模型(LLM)竞赛中,显存(VRAM)往往比算力更稀缺。近日,Databricks AI 研究团队发布了 FlashOptim,一个能够通过底层数值优化将优化器内存占用削减一半以上的工具包。它通过 24-bit 权重拆分和压扩 8-bit 量化技术,让原本需要 175 GiB 显存的 Llama-3.1-8B 微调任务,在 113 GiB 内即可跑出与 FP32 完全一致的精度。
痛点深挖:消失的显存去哪了?
标准的混合精度训练(Mixed-precision Training)表面上使用 FP16/BF16 运算,但在后台,为了保证梯度更新不“迷失”,优化器必须维护一份 FP32 的 Master Weights,加上 AdamW 存储的一阶动量(Momentum)和二阶动量(Variance),每个参数至少需要额外消耗 12 字节显存。
对于传统的量化方案(如直接线性量化到 8-bit),面临两个死穴:
- 分布不均:优化器状态(尤其是 Variance)呈现长尾分布,线性量化会导致精度信息大量丢失,训练极易崩溃。
- 冗余存储:Master Weights 包含的信息与前向传播用的低精度权重高度重叠,产生了巨大的内存带宽浪费。
FlashOptim 核心逻辑:数值直觉的胜利
1. 基于 ULP 的权重拆分 (Weight Splitting)
FlashOptim 并不是简单地存储 FP32。作者观察到,前向传播用的 (BF16)其实已经是全精度 的一个“近似”。根据 IEEE 754 标准, 一定落在 的 ULP(Unit in the Last Place,最后位单位) 半径内。
- 创新点:FlashOptim 只存储这个微小的“修正值” 。通过将这部分差异映射到 8-bit 整数空间,它能以 24-bit 的总位深实现等效于 FP32 的更新精度,相比直接存 FP32 节省了 25% 的参数空间。
2. 压扩量化 (Companded Quantization)
这是 FlashOptim 保持训练稳定的秘密武器。为了解决 8-bit 量化导致的梯度崩坏,作者引入了信号处理中的**压扩(Companding)**概念:
- 针对 Momentum:使用类似
softsign的函数,将极端值向中心挤压,使数值在量化箱(Bins)中分布更均匀。 - 针对 Variance:先开平方(Square Root)再量化。因为 Variance 是梯度的平方,开方能将其极大的动态范围拉回到线性量化可控的区间。
图 1:Llama-3.1-8B 微调显存对比,FlashOptim 显著降低了 Parameter 和 Optimizer 状态占比。
实验结果:无损的降本增效
在严苛的实验设置下(包括 ImageNet 视觉分类、10B Token 的 GPT-2 预训练、Llama-3.1-8B 数学微调),FlashOptim 展现了惊人的稳定性:
- 收敛一致性:无论是 SGD 还是 AdamW,FlashOptim 的 Loss 曲线与 FP32 几乎完全重合(见下方实验图)。
- 内存战绩:
- Master Weights: 4 Bytes -> 2 Bytes (-50%)
- Optimizer States: 8 Bytes -> 2.1 Bytes (-73%)
- 总体峰值显存:在 LLM 微调中降低 36% 左右。
图 2:消融实验显示,如果不使用压扩函数(紫色曲线),量化会导致训练迅速发散(Divergence);而 FlashOptim(蓝色)与基准对齐。
深度洞察:为什么它值得关注?
- Drop-in Replacement:FlashOptim 基于 Triton 实现了融合算子,其 API 与 PyTorch 原生优化器完全兼容,开发者只需更改一行代码。
- 正交性:它并不与 ZeRO (FSDP) 或梯度检查点(Activation Checkpointing)冲突。相反,这些技术可以叠加使用,产生乘数效应。
- 未来的路:目前 FlashOptim 主要优化的是参数相关的显存。对于卷积网络等“激活内存(Activation-heavy)”主导的模型,其收益相对较小。未来的研究方向可能会转向将类似的压扩逻辑应用到梯度压缩和激活值量化上。
总结 (Takeaway)
FlashOptim 证明了深度学习优化器并不需要始终如一的 FP32 精度。通过对数值分布的深刻理解,我们可以在 8-bit 的精度预算下,完成原本需要更高昂代价才能支撑的大规模模型训练任务。
代码现已开源:https://github.com/databricks/flashoptim
