中文

StreamBP:面向LLM长序列训练的内存高效精确反向传播

机器学习 2025-06-04 v1 人工智能

摘要

在长序列数据上训练语言模型是增强模型在复杂任务(例如长链推理)上能力的苛刻要求。然而,随着序列长度的扩大,即使在应用梯度检查点技术的情况下,反向传播过程中存储激活值的内存开销也变得巨大。为了应对这一挑战,我们提出了一种内存高效且精确的反向传播方法,称为StreamBP,它沿序列维度以逐层方式对链式法则进行线性分解,显著降低了激活值和logits的内存开销。所提出的方法适用于SFT、GRPO和DPO等常见目标。从实现角度来看,StreamBP通过利用语言模型的因果结构,实现了更少的计算FLOPs和更快的反向传播速度。与梯度检查点相比,StreamBP将反向传播的最大序列长度扩大了2.8到5.5倍,同时使用相当甚至更少的反向传播时间。请注意,StreamBP的序列长度缩放能力可以直接转化为批大小缩放,从而加速训练。我们进一步开发了通信高效的分布式StreamBP,以有效支持多GPU训练并拓宽其适用性。我们的代码可以轻松集成到任何Transformer模型的训练流程中,并可在 https://github.com/Ledzy/StreamBP 获取。

关键词

引用

@article{arxiv.2506.03077,
  title  = {StreamBP: Memory-Efficient Exact Backpropagation for Long Sequence Training of LLMs},
  author = {Qijun Luo and Mengqi Li and Lei Zhao and Xiao Li},
  journal= {arXiv preprint arXiv:2506.03077},
  year   = {2025}
}