论文

FORGE:融合寄存器上梯度消除以实现内存高效的 LLM 训练

FORGE: Fused On-Register Gradient Elimination for Memory-Efficient LLM Training

模型训练参数高效训练

摘要

反向模式微分计算每个权重梯度,将其写入内存,然后才让优化器将其读回。这种两阶段的时间表设定了现代训练的内存上限:在阶段之间的接缝处,每一层的梯度都是实时的。我们认为,这种物化梯度是分化如何进行的产物,而不是学习所需的数量——我们消除了它。 FORGE 将优化器步骤折叠到后向传递中,并一次将其应用到一个图块中,完全在寄存器中,因此每个梯度图块在生成时立即被消耗,并且永远不会成为张量。融合仅在更新发生时发生变化,而不是它计算的内容:在完全精度下,融合步骤可证明是精确的——对于每个元素规则,相同的优化器更新——并且这种精确性在张量和序列并行分片中仍然存在;在实践中使用的 bf16 和 8 位机制中,它是忠实的而不是位相同的,其偏差是有限的,并且对于权重存储,通过随机舍入呈现无偏差。因为每个梯度图块都是在相同的寄存器中生成和消耗的,所以它永远不会转换为 bf16 来存储和读回;因此,FORGE 保留了 bf16 和 8 位优化器因该转换而失去的全精度保真度。该方法也不依赖于一种架构或一种优化器:线性层无处不在,并且 FORGE 在任何元素规则下回收其中任何一个的梯度内存。根据经验,FORGE 可以将优化器步骤的内存减少一半以上,并且在 微调 典型的小批量大小和持续预训练下,运行速度提高约 1.5 倍;集成到张量并行 Megatron-LM 中,它适合 8B 训练,其微批次是标准优化器在相同 GPU 上允许的微批次的四倍。