EFFICIENT LONG-CONTEXT LANGUAGE MODEL TRAINING BY Core Attention Disaggregation

发表时间: 2025-10 · arXiv:2510.18121

原文: https://arxiv.org/abs/2510.18121

作者/机构: Yonghao Zhuang, Junda Chen, Bo Pang, Yi Gu, Yibo Zhu, Yimin Jiang, Ion Stoica, Eric Xing, Hao Zhang


速读

一句话结论 本文提出了一种名为核心注意力解耦(CAD)的新架构,通过将无参数的注意力计算从模型其他部分分离并独立调度,解决了长上下文大语言模型训练中的负载不均衡问题,在高达 512K 上下文长度的训练中实现了最高 1.35 倍的端到端吞吐量提升。

要解决什么问题 在长上下文大语言模型训练中,为了处理可变长度的文档,通常会采用文档打包的方式,但这会导致严重的工作负载不均衡。根本的机制卡点在于计算复杂度的不匹配:对于长度为 $l$ 的文档,其计算量可以表示为 $FLOPs(l) = \alpha l^2 + \beta l$,其中二次项 $\alpha l^2$ 来自核心注意力,线性项 $\beta l$ 来自上下文无关层(如前馈网络、线性层等);而激活显存的消耗 $M(l) = \gamma l$ 仅呈线性增长。在数据并行或流水线并行中,要让不同设备同时满足显存平衡(即 $\sum l_i = \sum l'_j$)和计算平衡(即 $\sum l_i^2 = \sum l'_j^2$)在数学上极难实现。这使得处理长文档的设备会成为拖慢全局的掉队者。现有的妥协方案都存在致命缺陷:如果通过重新分配文档来强行平衡计算量,会导致各设备总 token 数不同,进而引发显存膨胀甚至溢出;如果采用按文档的上下文并行来均分负载,虽然能平衡显存和计算,但会切碎短文档导致算力利用率低下,且需要全局收集键值状态,其通信延迟在扩展到 32 个节点时会占到总耗时的近 40%,同时给最后一个设备带来巨大的显存压力。

怎么做的 核心思路是提出核心注意力解耦,将无参数的核心注意力计算与模型的上下文无关层彻底分离,并调度到独立的资源池中执行。这一设计之所以能绕开上述卡点,是基于核心注意力的两个关键特性:一是无状态性,现代注意力内核通过重计算避免了物化庞大的注意力分数矩阵,使得核心注意力计算 $O = \text{softmax}(QK^\top)V$ 几乎不产生中间状态,因此对其进行调度纯粹是计算负载均衡问题,无需顾虑显存状态;二是可组合性,核心注意力可以在 token 级别任意切分,不同文档的切片可以重新组合成一个大的融合内核调用,只要切片大于内核的分块大小就能保持极高的硬件利用率。基于此,作者开发了 DistCA 系统,包含三个关键部件。首先是原地注意力服务器,为了避免给无状态的注意力计算分配专用 GPU 导致显存闲置,系统让每个 GPU 在计算上下文无关层和充当注意力服务器之间分时复用,兼顾了高算力和高显存利用率。其次是负载均衡调度器,它在 CPU 上运行感知通信的贪心算法,先计算出理想的单设备目标负载 $\bar{F}$,然后将设备分为盈余和赤字两类。调度器会评估将任务从盈余设备迁移到赤字设备的效率,通过计算优先级分数 $E = \Delta F_{max} / V_{comm}$(即单位通信成本带来的计算量转移)来挑选最高效的切片进行迁移,直到各设备负载差距缩小到容忍度 $\epsilon$ 以内。最后是乒乓执行机制,为了隐藏解耦带来的跨节点通信开销,系统将每个微批次拆分为两个等量的纳米批次交错执行,使得一个批次的张量传输与另一个批次的注意力计算完全重叠。

效果如何 实验在多达 512 张 H200 GPU 上进行,评估了 LLaMA 8B 和 LLaMA 34B 模型,上下文长度最高达 512K,使用了模拟预训练分布和专门的长上下文数据集 ProLong。核心对比基线是 WLB-LLM(文中复现为 WLB-ideal),该基线代表了采用可变长度数据块并试图用不平衡的多层感知机来补偿注意力负载不平衡的路线。量化结果显示,在不开启流水线并行的 3D 并行设置下,DistCA 在所有配置下均超越基线,最高实现 1.20 倍加速;在包含流水线并行的 4D 并行设置下,DistCA 的优势进一步扩大,在 8B 模型上端到端训练吞吐量提升高达 1.35 倍,在 34B 模型上提升达 1.25 倍,彻底消除了并行训练中的掉队者现象,并能有效利用流水线预热和排空阶段的闲置算力。该方法也存在一定的局限性和代价:在 34B 模型的 4D 并行实验中,由于每个微批次处理的张量形状不断变化,导致底层内存分配器频繁产生内存碎片并触发垃圾回收,这反过来延迟了 GPU 核函数的启动,限制了性能上限。此外,当前的调度器策略较为保守,强制每个查询切片必须使用完整的键值上下文,且在估算通信量时悲观地假设所有 token 都需要通过网络传输,忽略了目标设备上可能已经缓存的键值状态,这会导致对通信字节数的高估和非最优的传输调度。

A1 主要贡献

核心问题:在长上下文大语言模型(LLM)的训练中,由于文档长度可变,通过文档打包(document packing)方式处理数据会导致严重的负载不均衡。其根本原因在于,Transformer模型中的自注意力(self-attention)计算量随序列长度呈二次方增长,而模型其余部分的计算量仅呈近似线性增长。这种计算复杂度的不匹配导致在数据并行(DP)和流水线并行(PP)中出现“掉队者”(stragglers),即处理较长文档(或注意力计算量大的数据块)的设备会拖慢整个训练过程,降低系统吞吐量。

研究目标:本文旨在通过一种新的系统架构来解决长上下文LLM训练中的负载不均衡问题,从而提升端到端的训练吞吐量。

创新点与核心思想
本文提出核心注意力解耦(Core Attention Disaggregation, CAD),其核心思想是将无参数的核心注意力(Core Attention, CA)计算(即softmax(QK⊤)V部分,如图1所示)从模型的其他部分(如线性层、前馈网络等)中分离出来,并将其调度到一个独立的资源池(称为注意力服务器 Attention Servers)上执行。


图 1. Transformer及其由核心注意力引起的工作负载不平衡。

这一方法的可行性基于两个关键观察:
1. 无状态性(Statelessness): 核心注意力(CA)没有可训练的参数,且只存储极少的瞬时状态(如每行的softmax统计数据)。因此,对其进行负载均衡简化为了一个计算密集型任务的调度问题,而无需考虑内存状态的平衡。
2. 可组合性(Composability): 核心注意力的计算可以在token级别上进行任意粒度的切分。不同文档(来自不同DP副本或PP阶段)的token级计算任务可以被重新组合(re-batch)成一个大的、高设备利用率的融合内核调用,而不会损失现代注意力内核(如FlashAttention)的效率。

基于CAD思想,本文实现了一个名为DistCA的系统,并包含三项关键优化:
- 原地注意力服务器(In-place attention server): 通过让GPU在计算上下文无关层和作为注意力服务器之间分时复用,实现了高计算和高内存利用率。
- 乒乓执行机制(Ping-pong execution): 将每个微批次(microbatch)分为两个更小的“Ping”和“Pong”纳米批次(nano-batches),通过交错执行来将数据通信与计算完全重叠,隐藏通信开销。
- 负载均衡调度器(Workload balanced scheduler): 开发了一个感知通信的贪心调度算法,该算法动态地将文档切分为token级任务,并在注意力服务器之间进行调度,以在最小化通信开销的同时实现近乎完美的计算负载均衡。

通过在多达512个H200 GPU上对高达512K上下文长度的负载进行评估,DistCA相比现有系统,端到端训练吞吐量提升高达1.35倍,并消除了DP/PP中的掉队者现象。


A3 背景知识、关键观察与设计原则

2.1 LLM架构

2.2 LLM训练并行化

3.1 计算与内存的负载不均衡

3.2 现有方法的问题

3.3 核心注意力解耦的动机


A2 方法细节

DistCA系统由一个运行时系统和一个工作负载调度器组成。运行时负责交替执行上下文无关层和核心注意力层,并插入必要的通信。调度器则决定如何对每个文档进行分片以及如何将分片放置到不同设备上。


CA:核心注意力层;Linear:FFN,qkvo-proj;MISC:layer norm, dropout, ...
图 2. DistCA架构。

4.1 运行时

4.2 感知通信的贪心调度算法

5 实现


A4 实验环境

A4 实验结果

6.2 端到端实验

6.3 消融研究


A7 补充细节

7 相关工作

8 局限性


A5 结论

本文提出了核心注意力解耦(CAD),这是一种用于大语言模型训练的新架构,它将核心注意力模块与模型的其余部分分离开来,以实现独立的扩展和调度。基于核心注意力是无状态且在token粒度上可组合的观察,本文实现了DistCA系统,该系统具有一个感知工作负载的调度器以平衡计算同时最小化通信,以及一个乒乓执行方案来隐藏分派延迟。端到端评估表明,与最先进的训练系统相比,DistCA的吞吐量提升高达1.35倍,并且随着规模的扩大,其优势愈发明显。


A6 附录

A 核心注意力服务器最大分区大小的上限

B 通信开销函数