极长上下文状态空间模型训练的 adjoint 分片技术
机器学习
2025-01-03 v1 人工智能
计算与语言
摘要
尽管取得了非常快的进展,但高效训练大型语言模型 (LLMs) 在极长上下文中仍然具有挑战性。现有方法会回退到使用短上下文(训练中最大为几千个标记)进行训练,并在评估时使用推理时间技术处理长上下文(超过100万个标记的上下文窗口)。与长上下文推理不同,训练极长上下文输入提示受限于 GPU 内存可用性以及在最先进硬件上所需的惩人训练时间。与此同时,许多实际应用需要不仅推理,还需要在特定任务上进行长上下文训练/微调。例如, augmenting the context with various sources of raw reference information for fact extraction, fact summarization, or fact reconciliation tasks。我们提出了 adjoint 分片 (adjoint sharding),一种新技术,包含在训练期间对梯度计算进行分片,以将内存要求降低数个数量级,使训练极长上下文计算上变得可行。adjoint 分片基于 adjoint 方法,并计算等效于反向传播的梯度。我们还提出了截断 adjoint 分片以加快算法速度同时保持性能。我们提供了一个分布式版本和一个并行版本的 adjoint 分片以进一步加快训练速度。实证结果表明,所提出的 adjoint 分片算法在127万参数的大型语言模型上进行100万长度上下文训练时,将内存使用降低最高可达3倍。这使得在由五个 AWS P4 实例组成的训练基础设施上,将127万参数模型的训练/微调最大上下文长度从35K 标记提升到超过100K 标记。
引用
@article{arxiv.2501.00692,
title = {Adjoint sharding for very long context training of state space models},
author = {Xingzi Xu and Amir Tavanaei and Kavosh Asadi and Karim Bouyarmane},
journal= {arXiv preprint arXiv:2501.00692},
year = {2025}
}