第18章 大模型训练中的 Bug 修复:梯度累积
梯度累积旨在以较低的显存占用模拟全批次训练。由于梯度累积也常用于 DDP 和多 GPU 配置,因此该问题同样影响大规模训练运行。附注:如果您喜欢我们的作品,请别忘了在 GitHub 上 ⭐ 给我们星标,并加入我们 Discord 社区——你们的支持对我们意义重大!🦥
早在 2021 年,Zhaofeng 首次发现了该问题;上周,Benjamin Marie 也重新发现了它。他们表明,如果使用梯度累积,损失值会比使用全批次训练时更高:我们制定了一种新的方法论来解决这个问题——它现已集成到 Unsloth 中!请更新 Unsloth(pip install --upgrade unsloth)并使用 unsloth_train!我们提供了一个免费的 Colab 笔记本,使用修复后的训练器微调 Llama 3.2 1/3B——Colab 笔记本。还有一个免费的 Kaggle 笔记本。from unsloth import unsloth_train
# trainer_stats = trainer.train() << 带有缺陷的梯度累积
trainer_stats = unsloth_train(trainer)
💡 我们的发现
复现该问题
在尝试修复之前,我们首先需要复现这个错误。理论上,梯度累积在数学上应当等效于全批次训练。我们使用有效全批次大小(effective full batch size)为 16 进行训练,因此 bsz * ga(批次大小乘以梯度累积步数)应保持恒定。我们测试了 bsz 为 1、2、4、8 和 16 的情况,不幸的是,我们成功复现了该问题。对于较大的梯度累积步数,训练损失始终偏高。
什么是梯度累积?
在训练或微调过程中,我们需要从训练数据集中选取一定数量的随机行(样本)来更新模型权重。那么应该选多少行呢?在 Llama 3.1 等超大规模预训练任务中,为了减少过拟合并提升泛化能力,批次大小可能高达数百万。而在 Unsloth 的 Llama 3.2 等微调任务中,批次大小可能仅为 32 这样的小数值。
主要问题在于,大批次会占用大量内存。如果 1 个批次使用 1 个单位的内存,那么 100 万大小的批次就需要 100 万个单位的内存。我们如何既能模拟大批次训练,又不过度占用内存?
梯度累积应运而生!我们不再预先计算梯度,而是每当新的迷你批次(mini batch)输入时即时生成梯度。随后,我们将所有迷你梯度相加并进行缩放,从而得到最终的大批次梯度。
可能的原因分析
一种流行的理论认为,梯度累积在累加步骤中存在数值误差。但研究人员发现,即使使用 float32 进行累加,也会出现同样的效果。我们的发现表明,确实存在极其微小的累积误差。
但我们在下文可以看到,最终的求和不等于原始完整批次的损失——实际上它大了 G 倍(其中 G 是梯度累积步数)。 L = ∑ [ L 1 m 1 | L 2 m 2 | L 3 m 3 | L 4 m 4 ] L = ∑ [ L ¯ m ¯ | L ¯ m ¯ | L ¯ m ¯ | L ¯ m ¯ ] L = G L ¯ m ¯ ≠ L ¯ m ¯ 因此在梯度累积中,我们必须用梯度累积步数来缩放每个迷你梯度累积器,才能得到预期结果。 L = ∑ [ 1 G L 1 m 1 | 1 G L 2 m 2 | 1 G L 3 m 3 | 1 G L 4 m 4 ] 这在大批量下通常表现良好。但如果序列长度不同会怎样?这不会引发问题吗?我们通过完全移除分母来测试,即不使用归一化的交叉熵损失,而是直接使用未归一化的损失,以确认梯度累积是否仍然有效。以下是修改后的 Unsloth 训练运行中的训练损失: 奇迹般地,我们观察到所有训练损失曲线都吻合!这意味着分母绝对是罪魁祸首!这意味着简单地对每个梯度累积步进行平均是错误的,我们必须事先推导分母。
我们已经在 Unsloth 中实现了这个修复,现在所有 loss 曲线都能对齐,证明梯度累积确实等价于全量 batch 训练。数值差异另一个需要考虑的问题是:这个误差是否真的会影响最终权重的差异。为此,我们分别用 Unsloth 训练了一个 LoRA adapter:一次用全量 batch(bsz=16, ga=1),一次用梯度累积(bsz=1, ga=16)。
我们跑了所有组合(从 bsz=1,ga=16 一直到 bsz=16,ga=1),并将得到的 LoRA 权重与全量 batch 版本(bsz=16,ga=1)对比,计算 L2 范数差异。
结果表明:(1) 由于浮点运算,梯度累积本身存在固有误差(0.0068 L2 范数);(2) 梯度累积步数越多,L2 范数差异越大(从 0.0196 增长到 0.0286)。这实质上意味着梯度累积天然带有微小的浮点加法损失,而且直观上看,累积步数越大,偏差就越高。使用我们修复后的 Unsloth 梯度累积版本,L2 范数误差可以降低一个数量级以上。所以,请更新 Unsloth(pip install --upgrade unsloth)并使用 unsloth_train!我们还提供了免费的 Colab notebook,可以用修复后的 trainer 微调 Llama 3.2 1/3B,地址在这里 - Colab notebook.from unsloth import unsloth_train
# trainer_stats = trainer.train() << 有 bug 的梯度累积 ```html
trainer_stats = unsloth_train(trainer)
补充 - 数学证明
假设批大小(batch size)为 2,梯度累积步数(gradient accumulation steps)为 2,则在不使用梯度累积(即全批次训练)和使用梯度累积两种情况下的最终损失如下:
$$L = \frac{L_1 + L_2 + L_3 + L_4}{m_1 + m_2 + m_3 + m_4}$$
$$L = \frac{\frac{1}{2} \cdot \frac{L_1 + L_2}{m_1 + m_2} + \frac{1}{2} \cdot \frac{L_3 + L_4}{m_3 + m_4}}{\frac{1}{2} \cdot \frac{L_1 + L_2}{m_1 + m_2} + \frac{1}{2} \cdot \frac{L_3 + L_4}{m_3 + m_4}} \text{ (注意:原文第二个公式在排版上可能存在缺失分母的情况,根据上下文逻辑,这里展示的是累积梯度的平均效果对比)}$$
*注:根据原文排版,第二个公式似乎是想表达梯度累积后的损失计算方式。通常梯度累积的损失是各个 mini-batch 损失的加权平均。原文公式 $$L = \frac{1}{2} \frac{L_1+L_2}{m_1+m_2} + \frac{1}{2} \frac{L_3+L_4}{m_3+m_4}$$ 可能是指两个 mini-batch 平均后的损失值,而全批次是整体平均。*目标是证明标准的(或朴素的)梯度累积产生的损失始终不同于全批次训练(更高或更低)。我们先针对批大小为 2、累积步数为 2 的特定情况进行证明。其他长度情况类似。因此,我们需要证明:
$$\frac{\frac{1}{2} \cdot \frac{L_1 + L_2}{m_1 + m_2} + \frac{1}{2} \cdot \frac{L_3 + L_4}{m_3 + m_4}}{?} \ge \frac{L_1 + L_2 + L_3 + L_4}{m_1 + m_2 + m_3 + m_4}$$
*注:原文此处公式排版较为复杂,核心逻辑是比较“mini-batch 损失的平均值”与“整体序列的加权平均损失”。*随后,我们以类似的方式证明反向不等式。我们注意到,与其直接证明,不如证明该比值大于 1。既然已知序列长度 $m$ 始终大于 0,这是可行的。我们还使用了平均损失。
$$\frac{\frac{1}{2} \cdot \frac{L_1 + L_2}{m_1 + m_2} + \frac{1}{2} \cdot \frac{L_3 + L_4}{m_3 + m_4}}{\frac{L_1 + L_2 + L_3 + L_4}{m_1 + m_2 + m_3 + m_4}} \ge 1$$
$$\frac{\frac{1}{2} \cdot \frac{2\bar{L}}{m_1 + m_2} + \frac{1}{2} \cdot \frac{2\bar{L}}{m_3 + m_4}}{\frac{4\bar{L}}{m_1 + m_2 + m_3 + m_4}} \ge 1$$
通过简化和代数运算,我们得到:
$$\frac{\left( \frac{1}{2} \cdot \frac{2\bar{L}}{m_1 + m_2} + \frac{1}{2} \cdot \frac{2\bar{L}}{m_3 + m_4} \right) \cdot (m_1 + m_2 + m_3 + m_4)}{4\bar{L}} \ge 1$$
$$\frac{\left( \frac{\bar{L}}{m_1 + m_2} + \frac{\bar{L}}{m_3 + m_4} \right) \cdot (m_1 + m_2 + m_3 + m_4)}{4\bar{L}} \ge 1$$
$$\frac{\left( \frac{(m_3 + m_4)\bar{L}}{m_1 + m_2} + \frac{(m_1 + m_2)\bar{L}}{m_3 + m_4} \right) \cdot (m_1 + m_2 + m_3 + m_4)}{4(m_1 + m_2)(m_3 + m_4)\bar{L}} \ge 1$$
*注:分母通分后的推导过程*$$\frac{\left( \frac{(m_3 + m_4)}{m_1 + m_2} + \frac{(m_1 + m_2)}{m_3 + m_4} \right) \cdot (m_1 + m_2 + m_3 + m_4)}{4(m_1 + m_2)(m_3 + m_4)} \ge 1$$
$$\frac{\left( \frac{(m_1 + m_2 + m_3 + m_4)^2}{(m_1 + m_2)(m_3 + m_4)} \right) \cdot \frac{1}{4} \ge 1$$
$$\left( m_1 + m_2 + m_3 + m_4 \right)^2 \ge 4 (m_1 + m_2)(m_3 + m_4)$$
现在假设所有序列长度相同。我们预期全批次训练应与梯度累积相同。
$$\left( m + m + m + m \right)^2 \ge 4 (m + m)(m + m)$$
$$\left( 4m \right)^2 \ge 4 (2m)(2m)$$
$$16m^2 \ge 16m^2$$
我们可以看到结果符合预期——全批次训练和梯度累积是相同的!
```但如果某一个序列长度(只有 1 个)比其他序列长了一个微小的 ϵ,会发生什么?这会带来什么影响? $$\left( 4m \right)^2 \ge 4\left(2m\right)\left(2m\right)$$ $$\left( 4m + \epsilon \right)^2 \ge 4\left( 2m + \epsilon \right)\left( 2m \right)$$ $$16m^2 + 8m\epsilon + \epsilon^2 \ge 16m^2 + 8m\epsilon$$ 可以看到多出了一项 ϵ²,而它永远大于 0!不过我们还需要证明,当某一个序列长度比其他序列略短时,结论同样成立: $$\left( 4m \right)^2 \ge 4\left(2m\right)\left(2m\right)$$ $$\left( 4m - \epsilon \right)^2 \ge 4\left( 2m - \epsilon \right)\left( 2m \right)$$ $$16m^2 - 8m\epsilon + \epsilon^2 \ge 16m^2 - 8m\epsilon$$ 两种情况下不等式都成立,因为 ϵ² 永远大于等于 0。这基本上就证明了:在 bsz=2、ga=2 时,朴素的(也就是标准的)梯度累积得到的 loss 总是高于 full batch。随后我们把这个证明推广到其他 bsz 和 ga 的组合上,那部分推导会更复杂一些。
我们还必须证明该不等式在反方向上也成立——也就是说,目标是证明通用的朴素梯度累积并不等价于全批次训练。10月17日更新:我们与 Hugging Face 的朋友们合作,修复了其训练器中的该问题,查看此处的 PR。我们还了解到其他训练框架也在积极解决此问题,并正在与其中一些框架合作进行修复。💕 感谢大家!像往常一样,非常感谢所有使用和分享 Unsloth 的用户,我们由衷感激。同时特别鸣谢新加入的支持者:Dario、Bronson、Jun、John、Steven 和 Aaron!🙏
顺便说一下,我们正在招聘,欢迎通过 [email protected] 与我们联系!如常,请务必加入我们的 Reddit 页面或 Discord 服务器获取帮助或表达支持!你还可以在 Twitter 和 Substack 上关注我们。感谢阅读!Daniel 和 Michael Han 🦥
2024年10月15日下一步是视觉微调!免费开始使用 加入我们的 Discord