发表时间: 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 都需要通过网络传输,忽略了目标设备上可能已经缓存的键值状态,这会导致对通信字节数的高估和非最优的传输调度。
核心问题:在长上下文大语言模型(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中的掉队者现象。
O = softmax(QK⊤)V,它没有可训练的参数,也不需要其他token的CA中间输出。现代的IO感知注意力内核【索引2,Flashattention: Fast and memory-efficient exact attention with io-awareness,2022,NeurIPS】通过在反向传播中重新计算来避免在前向传播中物化巨大的注意力分数矩阵P,这使得核心注意力的中间状态量可以忽略不计,因而是无状态的。l的文档的计算量可表示为 $FLOPs(l) = \alpha l^2 + \beta l$,其中$\alpha l^2$来自核心注意力,$\beta l$来自上下文无关层。激活内存为 $M(l) = \gamma l$。要使两个微批次(分别包含长度为 $\{l_i\}_{i=1}^n$ 和 $\{l'_j\}_{j=1}^m$ 的文档)在计算和内存上都达到平衡,必须同时满足 $\sum_{i=1}^n l_i = \sum_{j=1}^m l'_j$ (内存平衡)和 $\sum_{i=1}^n l_i^2 = \sum_{j=1}^m l'_j^2$ (计算平衡)两个条件,这在实践中极难实现。可变长度数据块(Variable-length data chunk):该方法通过重新分配文档来平衡计算量(即 $\sum l_i^2$),但这会导致各微批次的总token数($\sum l_i$)不同,从而在某些设备上造成激活内存膨胀。随着序列长度增长,该方法会达到内存上限,无法再通过移动序列来完全平衡注意力计算。如图4所示,在512K长度的负载下,DP=8时GPU的平均空闲时间比例高达55%。
图 4. 在8B模型上处理512K token数据块时,不同并行策略下可变长度数据块的吞吐量和内存差异。
按文档上下文并行(Per-document CP):该方法为每个文档分配相等份额给每个CP rank,从而同时平衡了计算和内存。但它在规模化时面临三大瓶颈:
图 3. Llama-8B模型下,上下文并行中all-gather的延迟和内存分解。文档长度均为32k。
图 5. 核心注意力的吞吐量。
两种方法的组合:结合使用这些技术会继承各自的缺点。图6显示,在一个64-GPU、512K-token的实验中,增加CP度可以减少不平衡,但会降低吞吐量并有OOM风险;而增加DP度会导致严重的负载不平衡和次优的吞吐量。
图 6. 应用可变长度数据块和按文档CP时的吞吐量。
DistCA系统由一个运行时系统和一个工作负载调度器组成。运行时负责交替执行上下文无关层和核心注意力层,并插入必要的通信。调度器则决定如何对每个文档进行分片以及如何将分片放置到不同设备上。
CA:核心注意力层;Linear:FFN,qkvo-proj;MISC:layer norm, dropout, ...
图 2. DistCA架构。
核心注意力任务(CA-task)的定义:注意力服务器的工作负载被定义为核心注意力任务(CA-task),记为t。一个CA-task是针对一个查询分片q(t)及其上下文的键值分片kv(t)的核心注意力计算。一个文档被分割成多个不重叠的分片q1, q2, ..., qn,其完整的核心注意力结果是对应任务t1, t2, ..., tn结果的集合。
系统工作流程:如图2所示,在一组处理多个文档的GPU中,每个GPU都可以作为一个独立的注意力服务器。对于一批文档,在经过上下文无关层处理后,它们被分割成CA-tasks。每个CA-task被分配给一个注意力服务器。服务器接收到其被分配的所有任务的输入张量后,利用内核的可组合性,将所有CA-task批处理并在一个单独的内核中执行(例如通过一次FlashAttention调用)。计算完成后,每个CA-task的输出被发送回处理后续上下文无关层的源GPU。
中央调度器:一个运行在CPU上的中央调度器负责确定分片策略。它在GPU处理当前批次时,预取下一批次的文档,并使用预先计算的性能分析数据来估计每个文档和潜在CA-task的计算成本,从而生成一个分片和分配计划。
原地注意力服务器(In-place attention server):为了避免为CA计算分配专用GPU而导致的内存严重未充分利用问题(因为CA是无状态的,而FFN等层内存消耗巨大),DistCA采用原地注意力服务器设计。每个GPU周期性地在计算上下文无关层和充当注意力服务器之间切换角色。这样既能实现高内存利用率,又能平衡GPU间的计算负载。
乒乓执行(Ping-Pong execution):为隐藏通信开销,系统采用乒乓执行调度。每个输入微批次被分为两个等token数的纳米批次——“Ping”和“Pong”。这两个纳米批次的执行是交错的,使得一个的通信可以与另一个的计算重叠。此外,系统还将张量并行所需的节点内通信(通常通过NVLink)与核心注意力解耦引起的节点间通信(通常通过InfiniBand)进行重叠。
流水线并行支持:CAD可以自然地与DP和TP集成,并替代CP。对于PP,CA-task因其无权重的特性,来自不同PP阶段的任务可以与来自不同微批次的任务一同被调度和平衡。对于上下文无关层,由于所有PP阶段的微批次包含相同数量的token,它们的计算负载是相同的,因此是平衡的。为了防止设备在切换角色时空闲,系统调整了调度,使所有阶段在同一个tick内执行相同的阶段(要么全做前向,要么全做后向),这是通过将部分后向微批次逻辑上延迟到调度末尾的流水线气泡中实现的(如图8所示)。此外,在流水线预热和排空阶段,部分空闲的GPU时间被重新用于作为注意力服务器运行CA-task。
图 8. 正常1F1B和解耦注意力下的流水线并行调度。
n,计算出理想的每个服务器负载F¯。然后,将注意力服务器划分为负载有盈余(load > F¯)和有赤字(load < F¯)两类。d,并尝试从盈余源服务器迁移Item来弥补d的负载缺口。∆Fmax = min(F_Item, S_source, D_destination)。∆Fmax且通信成本Vcomm最小的分片。E = ∆Fmax / Vcomm。E越高表示迁移效率越高。∆Fmax,如果∆Fmax = F_Item,则整个Item被迁移;如果∆Fmax < F_Item,则该Item被拆分为两个子Item,新创建的具有∆Fmax FLOPs的子Item被分派到目标服务器。ϵF¯(一个容忍度epsilon)范围内,或者当剩余的迁移无法使E值超过一个很小的阈值时,调度器停止。这样,调度器在确保系统负载平衡的同时,避免了因微不足道的迁移而产生不必要的通信。模型架构:实验使用了LLaMA 8B和LLaMA 34B模型。具体配置如表2所示。
表 2. 实验模型配置。“Hidden”是隐藏维度大小,“#Head”是注意力头数,“Head Size”是每头维度。
硬件配置:所有实验均在NVIDIA DGX H200节点上运行。每个节点包含8块140GB H200 GPU。
3D并行(无PP):
表 3. 3D训练配置。
图 9. 3D并行(无PP)实验。加速比定义为WLB-LLM的平均运行时间除以DistCA的平均运行时间。
4D并行(含PP):
表 4. 4D并行训练配置。
图 10. 4D并行(含PP)实验。加速比定义为WLB-LLM的平均运行时间除以DistCA的平均运行时间。
系统开销:
图 11. 不同通信模式下的吞吐量。
调度器中的超参数:
ϵ如何权衡CA负载平衡与通信量。
图 12. 计算不平衡容忍度因子的影响。
本文提出了核心注意力解耦(CAD),这是一种用于大语言模型训练的新架构,它将核心注意力模块与模型的其余部分分离开来,以实现独立的扩展和调度。基于核心注意力是无状态且在token粒度上可组合的观察,本文实现了DistCA系统,该系统具有一个感知工作负载的调度器以平衡计算同时最小化通信,以及一个乒乓执行方案来隐藏分派延迟。端到端评估表明,与最先进的训练系统相比,DistCA的吞吐量提升高达1.35倍,并且随着规模的扩大,其优势愈发明显。
l的文档被均匀分为s个分片。查询状态的总通信量为$l \cdot h_q$。键值状态的总通信量为$h_{kv} \cdot (l/s \cdot s + l/s \cdot (s - 1) + \dots) = (s + 1)lh_{kv}/2$。t是计算一个token的上下文无关层的时间,B是网络带宽。整理后得到分片数量的上限为:$s \leq 2(tB - h_q)/h_{kv} - 1$。Llama-34B示例:以Llama-34B为例,其配置如表5所示。假设InfiniBand带宽为50GB/s,H200节点的MFU为50%(FP16下990TFLOPs)。
表 5. Llama-34B配置
| hidden (h) | key-value hidden (hkv) | intermediate (i) |
|---|---|---|
| 8192 | 2048 | 22016 |
计算时间t:一个token的上下文无关层总FLOPs为:
计算出的时间t代入上限公式后得到:
t随隐藏大小$h_q$二次方增长,这个上限s甚至会增加。v(·)的构建:对于一个有nq个查询token和nkv个键值token的分片,通信成本为:∆Fmax FLOPs时,选择最优的分片大小nq可以最小化通信量v(·)。最优解的推导如下:i + nq = j = nkv,通信量简化为:nq对应于可能的最小值,具体由以下约束决定: