DataStates-LLM:使用可组合状态提供程序实现Transformer模型的可扩展检查点

arXiv cs.AI 论文

摘要

DataStates-LLM提出了一种可扩展的检查点架构,利用可组合的状态提供程序,相比于现有解决方案,吞吐量提升高达4倍,训练时间减少2.2倍。

arXiv:2601.16956v1 Announce Type: cross 摘要:基于Transformer的大型模型(特别是大型语言模型(LLM))的快速增长,现已扩展到数万亿参数,需要跨数千个GPU使用复杂的混合并行策略(例如数据、张量和流水线并行)进行训练。对这些大规模分布式状态进行检查点对于多种用例至关重要,例如弹性、挂起恢复、调查不良训练轨迹以及解释模型演化。然而,现有的检查点解决方案通常将模型状态视为不透明的二进制块,忽略了底层数据结构的“3D异质性”——这些结构因内存位置(GPU与主机)、分布在多个文件中的“逻辑”对象分片数量、数据类型(张量与Python对象)及其序列化要求而异。这导致了由于阻塞的设备到主机传输、数据无关的序列化和存储I/O争用而带来的显著运行时开销。在本文中,我们介绍了DataStates-LLM,一种新颖的检查点架构,利用状态提供程序将状态抽象与数据移动解耦。DataStates-LLM利用模型参数在前向和后向传播中的不变性来执行“惰性”、非阻塞的异步快照。通过引入状态提供程序,我们有效地合并了碎片化的异质分片,并将元数据的序列化与批量张量I/O重叠。我们在256个A100-40GB GPU上对高达70B参数的模型进行了DataStates-LLM评估。结果表明,与最先进的解决方案相比,DataStates-LLM实现了高达4$\times$倍的检查点吞吐量提升,并将端到端训练时间减少了高达2.2$\times$倍,有效缓解了超大规模LLM训练中的序列化和异质性瓶颈。
查看原文
查看缓存全文

缓存时间: 2026/06/29 05:28

# DataStates-LLM:使用可组合状态提供者为Transformer模型实现可扩展检查点

来源:https://arxiv.org/html/2601.16956

###### 摘要

大型Transformer模型,特别是大型语言模型(LLM),其规模已扩展至数万亿参数,必须利用复杂混合并行策略(例如数据并行、张量并行和流水线并行)在数千GPU上进行训练。对这些大规模分布式状态进行检查点操作,对于弹性恢复、暂停恢复、研究不良训练轨迹以及解释模型演化等广泛用例至关重要。然而,现有检查点解决方案通常将模型状态视为不透明的二进制数据块,忽略了底层数据结构的“3D异构性”——这些结构因内存位置(GPU vs. 主机)、跨多个文件分片和拆分的“逻辑”对象数量、数据类型(张量 vs. Python对象)及其序列化要求而异。这导致了由于阻塞的设备到主机传输、数据无关的序列化和存储I/O争用而产生的显著运行时开销。

在本文中,我们提出**DataStates-LLM**,一种新颖的检查点架构,它利用**状态提供者**将状态抽象与数据移动解耦。**DataStates-LLM**利用前向和反向传播过程中模型参数的不可变性,执行“惰性”、非阻塞的异步快照。通过引入状态提供者,我们高效地合并碎片化的异构分片,并将元数据的序列化与批量张量I/O重叠。我们在256个A100-40GB GPU上对高达700亿参数的模型评估了**DataStates-LLM**。结果表明,与现有最先进的解决方案相比,**DataStates-LLM**的检查点吞吐量提高了最多4倍,端到端训练时间减少了最多2.2倍,有效缓解了超大规模LLM训练中的序列化和异构性瓶颈。

## I. 引言

大型语言模型(LLM)已成为现代人工智能的基石,推动了文本生成、理解以及特定领域科学发现方面前所未有的能力[6]。为了实现这些能力,模型规模急剧膨胀,通常超过数千亿到数万亿参数[3]。训练如此庞大的模型需要包含数千个GPU的高性能计算(HPC)基础设施,并需要运行数周或数月[34, 6]。为管理这些模型的内存占用,训练运行时采用了复杂的并行策略:它们结合了数据并行(DP)、张量并行(TP)和流水线并行(PP),以及诸如ZeRO[23]和FSDP[36]的优化器状态分片技术(例如BLOOM[34]和Llama 3[3])。

##### 动机:对可扩展检查点的需求

鉴于LLM训练的规模之大和持续时间之长,检查点是确保弹性和生产力的基本原语。在规模化运行中,硬件故障、软件错误和超时在统计上不可避免[30]。如果没有频繁的检查点,需要重新计算的丢失计算无论在时间还是资源方面都代价过高。例如,Llama 3 405B模型训练涉及16K GPU,耗时54天,每2.8小时就会遇到一次故障[3, 30]。另一个例子是阿里巴巴的Unicron训练,报告故障率为43.4%[8]。除了弹性,检查点对于解决训练不稳定(例如PaLM和GLM-130B模型中的损失尖峰[27])也至关重要,这些不稳定难以预测和防御,使得回滚后调整成为最可行的修正策略。其他几个重要场景的生产力也依赖于可扩展的检查点:基于人类反馈的强化学习(RLHF)、迁移学习、通过沿训练轨迹合并不同模型状态来加速收敛。在这些情况下,检查点频率可以高达每次迭代。无论什么场景,检查点都涉及基本的能力:频繁地捕获分布式AI模型状态,而不引起显著的开销(即阻塞),从而中断应用程序的进展。

##### 最先进技术的挑战和局限性

与传统深度学习模型(ResNet、VGG等,通常大小为几百MB,适合单个GPU内存)相比,LLM通常由数十亿参数组成,导致**巨大的检查点体积**。大型模型和优化器状态需要一致地组合多种不同数据类型和大小的数据结构,如张量、数组和自定义Python对象(例如字典、伪随机数生成器的种子等)。这些数据结构可能在不同运行时和语言中生成,并且可能驻留在不同内存层级(例如,主机内存中的Python字典、GPU内存中的C++/CUDA张量等)。此外,数据并行、流水线并行和张量并行以及冗余消除方法(例如ZeRO[23]、FSDP[36])的组合,导致这些数据结构在大量计算节点上细粒度分布。我们将这些在数据类型、数据大小以及不同存储层级和计算节点上的不同数据分片/分布策略方面新兴的多方面变化称为**LLM检查点的3D异构性**。这一方面未被最先进的检查点方法[31, 12, 20, 15, 33]充分解决,导致检查点期间产生显著开销。具体来说,最先进LLM训练运行时(例如DeepSpeed[23])中实现的许多检查点方法通过从所有GPU并行启动检查点捕获以饱和I/O带宽来处理数据分片。然而,它们以阻塞方式进行,中断了关键训练路径。替代方案如多级异步检查点技术旨在通过在快速层级上捕获检查点,然后在后台将检查点从快速层级刷新到较慢层级[16, 15, 31]来掩盖这些开销。然而,由于以下几个原因,将这些技术直接应用于LLM运行时并不简单。首先,GPU通常没有足够的备用内存容量来捕获完整检查点(CheckFreq[15]、GEMIMI[33])。其次,虽然可以直接在主机内存上捕获检查点(例如TorchSnapshot[20]、CheckFreq[15]),但有限的GPU到主机PCIe传输带宽比GPU内存带宽慢几个数量级,并且还需要由通常位于同一计算节点上的多个GPU共享,从而导致较高的I/O开销。没有更好的技术在初始快速层级上捕获检查点,多级异步检查点的好处会大大降低,以至于不比同步检查点快多少。例如,尽管有高速链路(50+ GB/s网络和25+ GB/s PCIe),LLM检查点吞吐量远未达到链路容量饱和(例如,REFT[32]报告饱和度为38%),并且常常低至几GB/s(例如,如Nebula[14]所报告)。由不同运行时和编程语言在不同内存层级生成的数据的异构性也是一个问题。上述大多数检查点方法高度优化了张量捕获和序列化,但忽略了其他数据结构。因此,训练运行时(如DeepSpeed[23])通常在并行捕获分布式张量内容之前,以集中方式单独收集和捕获高层元数据和其他Python数据结构。在大规模下,即使高层元数据的大小远小于张量的大小,这个阻塞步骤也成为一个显著瓶颈。

##### 关键见解与贡献

在本文中,我们提出**DataStates-LLM**,一个专门为**3D检查点异构性**设计的高性能异步检查点系统。我们的方法建立在两个关键见解之上。首先,模型参数和优化器状态在训练迭代的计算密集型前向和反向传播过程中保持**不可变**。这为在不使用中间暂存区域或阻塞I/O传输和计算的情况下执行“惰性”设备到主机(D2H)拷贝创造了机会窗口,解决了最先进LLM检查点方法面临的多级异步技术的局限性。其次,让检查点运行时能够感知所有数据结构(无论其类型或位于何处)为更好地并行化和重排序操作提供了机会(例如,在将GPU驻留的数据结构刷新到主机的同时,将主机驻留的数据结构刷新到磁盘)。本文利用这两个关键见解,并扩展了我们之前的工作[12],设计和实现了一个专门优化的检查点运行时,以高效处理大规模3D检查点异构性。通过引入**可组合的状态提供者**,这些是针对捕获异构性特定方面而优化的互补抽象,我们能够高效地异步流式传输到一个全局一致的状态,该状态代表一个完整的检查点。我们将贡献总结如下:

1. 1.**LLM检查点的差距分析:** 我们量化了3D并行性对检查点组成的影响,突出了数据结构的“3D检查点异构性”(GPU vs. 主机,张量 vs. 对象),并将序列化确定为关键阻塞瓶颈(§IV)。
2. 2.**状态提供者的设计:** 我们引入**状态提供者**,一种封装异构数据结构语义的中间件抽象。这允许**DataStates-LLM**对张量执行零拷贝序列化,同时高效处理复杂的Python对象,解决了状态异构性挑战(§V-A)。
3. 3.**惰性、非阻塞异步性:** 我们实现了一种惰性状态捕获机制,将D2H传输与训练的不可变阶段(前向/反向传播)重叠,从而隐藏了状态捕获的成本(§V-A2)。
4. 4.**精简的多层级内核加速I/O引擎:** 我们设计了一个流水线I/O引擎,它管理固定主机内存池,并使用低延迟I/O库(如liburing)执行内核加速、多线程刷新到持久存储,从而最大化I/O带宽利用率(§V-A4)。
5. 5.**可扩展性评估和I/O重叠分析:** 我们在Polaris超级计算机上评估**DataStates-LLM**,在256个A100-40GB GPU上训练高达700亿参数的Llama 2模型,并进行消融研究以调查每个设计提案的收益。我们展示了与最先进的方法TorchSnapshot以及我们之前的工作相比,检查点吞吐量提高了3倍–4.2倍,端到端训练时间减少了1.3倍–2.2倍(§VI)。

## II. 背景

##### 3D并行:数据、流水线和张量策略

数据并行(DP)仍然是通过跨多个工作器复制模型以加速深度学习的基础技术,这些工作器处理独立的微批次[22, 34, 3]。为保持副本一致性,在模型更新阶段之前通过all-reduce集合通信同步梯度。然而,随着LLM规模超过单个GPU容量,需要混合3D并行。流水线并行(PP)将模型垂直划分为由顺序Transformer层组成的阶段。为最大化吞吐量和GPU占用率,将微批次划分为微批次,允许前向和反向传播以流水线方式跨阶段重叠。张量并行(TP)通过将单个Transformer块及其关联内存分布到多个GPU上提供水平分片[26]。由于层内通信开销较高,TP通常限于通过高速互联的节点本地GPU。在极端规模下,这三种策略被组合成**3D并行**,以平衡内存效率和通信延迟。

参见图注
图1:AI模型训练期间针对流水线并行(PP)、张量并行(TP)和数据并行(DP)的检查点分片。

##### 通过状态分片减少冗余

DP副本固有的冗余可以通过跨工作器分片模型状态来缓解,这样每个工作器负责管理唯一的分片。像DeepSpeed[21]和FSDP[36]这样的运行时实现优化阶段(例如ZeRO阶段-1/2/3),分别分片优化器状态、梯度和模型参数。虽然分片通过减少容纳模型所需的GPU内存总量来改善内存效率,但它在前向/反向传播期间需要集合通信来重建完整状态,并在检查点期间使状态管理复杂化。

##### 状态分片对检查点的影响

在单个GPU上训练的情况下,模型和优化器状态被序列化为单个文件(图1(a))。DP允许通过让独立工作器检查点不同分片来实现I/O并行化(图1(b)),这是像DeepFreeze[16]和TorchSnapshot[20]等系统使用的方法。对于LLM,分片不仅延伸到DP,还延伸到单个模型层,使...(注:原文在此处中断,但根据上下文,应该是继续描述分片对检查点的影响。由于只给了这部分内容,我们按原文翻译。)

(注:原文在“使...”处截断,后续内容未提供。我们完成现有部分的翻译。)

此外,由于存在冗余消除(如ZeRO),分片策略变得更加复杂。例如,在ZeRO阶段3中,所有模型状态(参数、梯度和优化器状态)都在数据并行工作器之间分片。这意味着每个工作器只持有其分片部分的完整副本。在检查点时,需要捕获每个工作器的独立分片,而不是捕获完整的模型副本。这种分片不仅增加了检查点的数量(更多文件),还引入了协调挑战,因为需要确保所有工作器同时捕获状态而不产生不一致。此外,元数据(如学习率、训练步数、随机状态)通常较小但必须全局一致,这进一步增加了复杂性。

相似文章

@askalphaxiv: 另一项关于循环Transformer的酷研究。他们提出一个问题:“我们能否直接在推理时循环一个冻结的、现成的检查点…

X AI KOLs Timeline

本研究介绍了一种技术,通过使用阻尼Runge-Kutta子步骤,在推理时循环冻结的、现成的Transformer检查点,将Transformer层视为残差ODE中的欧拉步骤。这无需微调、架构更改或新权重即可增加额外的潜在计算,在MMLU-Pro、GPQA和ARC等知识任务上显示出收益。

语言模型需要睡眠

Hacker News Top

本文提出了一种类似睡眠的巩固机制,适用于基于Transformer的大语言模型,该机制定期将最近上下文转换为SSM块中的持久快速权重,清除KV缓存,从而在不增加推理延迟的情况下提升长期推理能力。

语言模型需要睡眠

Hugging Face Daily Papers

本文提出了一种针对Transformer模型的类睡眠巩固机制,该机制利用快速权重和递归传递来改进长上下文处理,同时保持推理速度。