发表时间: 2025-02 · arXiv:2502.21231 (ByteDance)
文章标题:ByteScale: 在超过12,000个GPU上以2048K上下文长度高效扩展LLM训练
作者/机构:Hao Ge (Peking University), Junda Feng (ByteDance Seed), Qi Huang (ByteDance Seed), Fangcheng Fu (Peking University), Xiaonan Nie (ByteDance Seed), Lei Zuo (ByteDance Seed), Haibin Lin (ByteDance Seed), Bin Cui (Peking University), Xin Liu (ByteDance Seed)
为了满足现代LLM应用(如文档摘要、视频理解、智能体交互和代码补全)对长程依赖理解的需求,扩展模型的上下文长度至关重要。然而,扩展至长上下文面临着自注意力机制带来的内存和计算复杂度呈二次方增长的根本挑战。现有方法如Flash Attention将内存复杂度从$O(N^2)$降至$O(N)$,但要进一步扩展,则需跨设备分区序列。现有框架通常采用数据并行(DP,分发不同序列)和上下文并行(CP,切分单个序列)两种正交技术,并将设备组织成静态的2D网格。这种静态设计依赖于一个假设:所有序列长度相同,以保证负载均衡。
然而,在真实世界的训练场景中,无论是文本还是多模态数据,序列长度通常是可变的且分布倾斜。此外,强化学习中思维链推理过程的长度增加也加剧了长度异质性。现有框架为了处理最长序列,必须配置足够大的CP组,导致所有短序列也必须在整个CP组上进行不必要的分区和通信。这种数据异质性与静态系统设计之间的不匹配导致了两个核心问题:
1. 冗余通信:短序列被迫参与为长序列设计的复杂通信过程,即使它们本可以在单个或少数设备上处理。此外,对于短序列,用$O(N^2)$的计算来掩盖$O(N)$的通信非常困难。
2. 计算不均衡:尽管通过CP可以均匀分配令牌以平衡内存,但每个令牌的计算复杂度与原始序列长度相关($O(N^2)$),导致不同设备上的执行时间不同,从而产生同步等待的空闲时间。
为解决上述挑战,本文提出了ByteScale,一个为大规模长短序列混合训练设计的高效、灵活且可扩展的训练框架。其主要贡献如下:
Transformer架构【40,Attention is All you Need,2017,NeurIPS 2017】已成为当今大语言模型(LLMs)【5, 14, 32, 39】最流行和广泛使用的基础架构。它通常由一系列Transformer层组成,每层包含一个注意力模块和一个前馈网络(FFN)模块。如图1所示,自注意力机制需要序列中的所有令牌参与计算以捕获整个文本的上下文信息。相比之下,其他操作如归一化、线性投影和激活函数则执行令牌级计算,允许每个令牌独立处理。
随着模型大小和训练数据的持续扩展,分布式训练技术在LLM训练中不可或缺。
- 数据并行(Data Parallelism, DP):DP【9, 24, 37】将训练数据均匀分布在各个设备上,每个设备持有一个模型副本。在每个训练步骤中,设备独立处理其本地数据,然后全局同步梯度以更新模型。ZeRO系列方法【35,ZeRO: memory optimizations toward training trillion parameter models,2020,SC 2020】进一步增强了DP的可扩展性。
- 模型并行(Model Parallelism):模型并行将模型分布在设备上,包括张量并行(TP)【38,Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism,2019,CoRR abs/1909.08053】和流水线并行(PP)【16, 28, 29】。TP执行操作内分区,将层内的操作和参数划分到不同设备(例如Megatron-LM中的行并行和列并行线性层),这需要通信中间结果(激活值),通常在单个节点内使用。PP采用操作间分区,将模型层分段为不同阶段,仅需在连续阶段间通过点对点(P2P)通信交换激活值,从而实现跨多节点的模型分区。
- 混合并行(Hybrid Parallelism):混合并行结合多种并行策略以提高训练效率。特别是,Megatron-LM通过集成DP、TP和PP,采用了3D并行策略【21, 30, 38】,使其成为当今大规模模型训练的主流方法。
- 梯度累积(Gradient Accumulation):为提高效率和收敛性,LLM通常需要大批量大小【6, 15, 39】(例如,在拥有10K GPU的集群中,每个批次处理近30-80M令牌是常见做法)。受硬件内存限制,一次性处理整个大批量是不可行的。梯度累积将每个全局批次(即每个训练步骤中采样的数据)分成多个微批次。这些微批次的梯度被累积起来,等同于一次性处理整个全局批次所产生的梯度。
为了在当前静态并行策略中支持可变长度的序列,需要使用填充(padding)和打包(packing)等技术。如图2所示,填充将同一批次中的序列填充到相同长度,但这会导致计算浪费。打包【22,Efficient sequence packing without cross-contamination: Accelerating large language models without impacting performance,2021,CoRR abs/2107.02027】将多个序列连接成一个单一序列,不含填充令牌。它采用一种特殊的分段注意力掩码,以确保每个序列在自注意力中被独立处理。
由于自注意力的时间和内存复杂度均为$O(N^2)$,当上下文长度扩展时,这种二次复杂度成为瓶颈。Flash Attention【7, 8】通过优化内存I/O和采用分块(tiling)技术,将内存复杂度从$O(N^2)$降低到$O(N)$,但时间复杂度仍为$O(N^2)$。上下文并行(CP)【4, 23, 25, 31】进一步将序列划分到$P$个设备上,将每个设备的内存从$O(N)$减少到$O(N/P)$。根据图1,CP沿序列维度对QKV进行分片,跨令牌操作需要使用环形点对点(ring-style P2P)通信在设备间交换KV切片,该通信与计算重叠。该技术也适用于打包序列,其实现细节将在第7节详述。值得注意的是,每个子序列也必须在所有CP ranks上进行分片,如图2(c)和图3(a)所示。
LLM在序列数据上进行训练。如第1节所述,训练数据通常包含可变长度的序列。存在两个观察和一个重大挑战:
- 观察1:真实世界数据集中序列长度呈偏斜分布。 如图4所示,我们分析了两个用于长上下文训练的数据集:开源的GitHub数据集和生产环境的Byted数据集。我们观察到两者在序列长度上都呈现偏斜分布。例如,在Byted数据集中,如果我们随机采样一个全局批次,近80%的样本是4K令牌或更短,而只有0.05%的样本能达到2M令牌。然而,从令牌分布的角度看,这0.05%的样本(>=2M)贡献了全局批次中12.1%的令牌,而1%的样本(>=128K)贡献了44.3%的令牌。尽管GitHub数据集中长序列的比例较低,但其16.2%的令牌来自超过128K的序列,显示出显著的数据异质性。
- 观察2:混合长短序列能提升模型性能。 现有工作【12,How to Train Long-Context Language Models (Effectively),2024,CoRR】已证明,仅在长上下文数据上训练会导致短上下文性能下降。LLaMA3报告【11,The Llama 3 Herd of Models,2024,CoRR】指出,在训练一个128K上下文的模型时,将0.1%的长数据与原始短数据混合,可以优化短上下文和长上下文基准测试的性能。DeepSeek-R1【10,DeepSeek-R1: Incentivizing Reasoning Capability in LLMs via Reinforcement Learning,2025,CoRR abs/2501.12948】展示了RL过程中训练集上的平均响应长度,表明逐渐增加和多样化的响应长度有助于提高模型性能。
现有系统在整个训练过程中应用静态并行策略。通常,它们假设所有(打包的)序列长度相同,并设置一个固定的CP等级,以将它们分摊到足够多的设备上,从而避免内存溢出(OOM)错误。如§2.3所述,为了处理可变长度序列,通常会将序列打包至上下文长度。然而,如图3(a)-(b)所示,所有序列都必须在整个CP组中进行分区,即使对于较短的序列来说这是不必要的。
例如,假设每个设备能处理8K令牌,要训练一个上下文长度为1M令牌的LLM,需要128的CP等级。此配置需要128个独立设备来处理一个1M令牌的序列。同时,大量较短的序列,如长度为4K、8K和16K的序列,被打包至1M令牌,并在一个有128个设备的CP组中处理。如图14所示,打包序列中的每个子序列都需要在CP ranks上被分区成128个块,并执行环形P2P通信。实际上,对于长度低于8K的序列,执行跨设备分区和通信是不必要的。对于16K令牌的序列,只需要两个CP ranks。对这些较短序列使用与最大序列长度相同的CP等级,会导致过度的冗余通信。当序列长度高度偏斜时,这个问题会更加严重。
我们提出ByteScale来解决这些挑战。如图7所示,它由三个主要组件构成。Profiler用于分析环境、模型配置、数据分布,并为其他组件构建成本模型。Communication Optimizer通过数据感知分片、动态通信和选择性卸载,为长短序列提升通信效率。Balance Scheduler通过并行感知的数据分配来解决计算不均衡问题。
本节描述ByteScale如何优化通信开销。首先,它通过动态序列分片和通信减少短序列的冗余通信。其次,它通过选择性卸载进一步压缩长序列的通信成本。
混合数据并行(Hybrid Data Parallelism, HDP):我们首先引入一种新的并行策略,即混合数据并行(HDP),以实现对不同序列长度的高效训练。DP和CP都将训练数据在设备间进行分区。DP通过将不同样本均匀分布到设备上执行数据间分区,而CP通过将单个样本分片到设备上执行数据内分区。HDP统一了数据间和数据内分区,其定义为将令牌均匀分布在设备上。它可以替代传统的DP和CP,HDP的并行度等于DP和CP并行度的乘积(即$P_{hdp} = P_{dp} \times P_{cp}$)。
HDP的异构行为:与DP和CP要求所有DP/CP ranks执行一致的计算或通信行为(例如,CP要求所有CP ranks参与同构的环形P2P通信)不同,HDP允许HDP ranks之间存在异构行为。它有两个关键特性:
NCCL缓冲优化:创建NCCL通信组会产生额外开销。首先,建立通信组的过程本身很慢,为每个序列动态创建新组会显著降低训练效率。其次,创建过多的通信组会为每个GPU消耗额外的5-10GB内存用于NCCL缓冲区,进一步减少可用内存。幸运的是,分布式注意力利用P2P通信。通过在所有HDP ranks上建立一个全局通信组,任意两个设备之间的P2P通信可以直接重用现有组,从而减轻了创建临时通信组带来的时间和内存压力。
优化器状态分片:HDP将令牌均匀地划分到设备上,并且既不分片模型参数也不分片梯度。这意味着HDP ranks像DP一样复制模型状态。因此,ZeRO系列技术也适用于HDP,如图8(a)所示,HDP在所有HDP ranks上利用ZeRO-1来最大程度地分片优化器状态,从而最小化内存使用。
损失与模型更新:尽管HDP ranks可能在不同的微批次中执行不同的异构通信,但参数的最终梯度与标准DP中获得的梯度是等效的。如图9所示,每个令牌都对参数$W$贡献一个梯度,最终梯度$grad_W$是全局批次(表示为B)中所有令牌梯度的总和。设$grad(t, W)$表示令牌$t$对参数$W$的梯度。那么$grad_W$可以表示为:
由于参数是复制的,并且令牌在HDP ranks(表示为R)之间均匀分布,本地累积的梯度对应于分配给每个rank(表示为$B_r$,即rank $r$中的微批次)的令牌梯度的部分和。因此,与DP类似,将在所有HDP ranks上执行全局集体通信(如All-Reduce或Reduce-Scatter)来聚合部分梯度。这也得到了来自所有令牌的梯度$grad_W$:
公式(2)等价于公式(1),并确保HDP中梯度累积的结果与标准DP中的结果等效。此外,由于我们计算了全局批次中所有令牌的梯度$grad_W$,它也需要通过令牌总数进行缩放,我们通过令牌级损失来实现这一点,该损失按令牌数量而不是样本数量来缩放损失。
act_ctx的通用组件(列表1)来支持激活值卸载。该组件分别为D2H(设备到主机)和H2D(主机到设备)维护两个cuda流。它自动从计算图中捕获激活张量,并在前向传播的适当时间将它们卸载到CPU(使用asyncCudaMemcpy API),并在D2H流和计算流之间建立异步依赖关系。计算图中的原始张量被元数据{层id, 激活id}替换。同样,在后向传播期间,存储在计算图中的元数据用于索引和在H2D流中重新加载相应的激活值。图10展示了整个过程。act_ctx还支持一个名为offload_ratio的参数,提供令牌级的细粒度控制,以控制卸载到CPU的激活值比例。此功能在节省GPU内存与实现最佳计算重叠之间取得平衡。# 列表1. act_ctx的用法
with act_ctx(offload_ratio=0.5):
# 前向传播
hidden_states = model_layer(hidden_states)
# 后向传播
loss.backward()
offload_ratio。这种方法有效地将长序列所需的ranks数量从$S_i/C$压缩到$n(S_i)$,如图11(a)所示。它不仅显著减少了通信开销,还使更多可用的HDP ranks能够处理数据,从而提高效率。ByteScale的整体流程在算法1中概述。简而言之,该算法遍历全局批次中的每个序列$S_i$。对于长序列,它推导出卸载比率$\rho$并确定所需的ranks数量$n(S_i)$(第1-6行)。对于短序列,它将它们打包以填满每个rank的容量$C$(第7-9行)。处理后的序列随后被分配给$P_{hdp}$个ranks,算法返回微批次和offload_ctx以供执行(第10-12行)。
本节我们介绍平衡调度器,以解决DP和PP的不平衡问题。通过精心编排数据分配(替代算法1中的第10行),它在保持§5中实现的最小通信的同时,缓解了这些不平衡。我们将首先概述几个关键见解,然后提出我们的启发式解决方案。
梯度累积要求不同的DP ranks执行相同数量的微批次,这是基于所有微批次具有相同计算负载的假设。然而,如§3.3所述,不同微批次的执行时间可能显著不同。在ByteScale中,我们重新定义了一个更灵活的策略,允许不同的HDP ranks处理不同数量的微批次(大小相同但工作负载不同),以缓解不平衡问题。如图13所示,这使得所有ranks几乎同时完成计算。更重要的是,这个策略不影响模型收敛。无论序列如何分配给HDP ranks,我们最终都计算全局批次中所有令牌的梯度总和,如§5.1所讨论的,这确保了数学上的等价性。
见解1:当不同长度级别的序列被分配到不同的流水线时,PP气泡较少。
确保流水线处理的微批次具有相似的执行时间至关重要。如图13(b)所示,当$P_{pp}=4$时,时间轴上任意4个连续的微批次将由4个PP阶段同时执行。如果它们的执行时间差异很大,就会出现额外的PP气泡。由于全局批次中长序列数量有限,一些流水线不得不被分配多种长度级别的序列。幸运的是,只有在过渡阶段(例如,当4个连续的微批次属于不同长度级别时)才会导致额外的PP气泡。
策略:我们为平均执行时间较短的流水线分配更多的微批次。如图12(a)-(b)所示,pipeline-0处理平均执行时间较长的微批次,因此只分配了8个微批次。相比之下,pipeline-1被分配了18个微批次,以与pipeline-0同步。此外,由于微批次更多,气泡率进一步降低。
见解2:当不应用流水线并行时,只需在每个时间步保持负载均衡。
如果只应用DP而不应用PP,实现负载均衡只需要在任何给定时间,由不同HDP ranks执行的微批次具有相似的执行时间。无需考虑时间轴上不同时间步之间微批次的工作负载不平衡。
策略:一个直接的方法是在同一时间将相同长度级别的序列分配给不同的HDP ranks,如图13(a)所示。此外,我们仍然为处理较短序列的ranks分配比其他ranks更多的微批次。最终,这确保了所有HDP ranks几乎同时同步梯度。
算法2描述了平衡策略。首先,我们按长度降序对全局批次B中的序列进行排序。然后将这些有序序列划分为FLOPs总和近似相等的桶,因此平均长度较长的桶包含的序列较少(第3-5行)。其次,我们确定那些执行时间较短的ranks,以便后续分配(第7-9行)。第三,如果使用DP-Balance策略,我们从同一个桶中选择序列。否则,如果使用PP-Balance策略,我们从所有桶中顺序选择序列。实践中,执行时间较短的ranks会被分配更多的序列(第12-15行)。最后,我们重复第二和第三步,直到所有桶都为空。
ByteScale基于Python、C++和CUDA实现了约16000行代码,并已与MegaScale【18,MegaScale: scaling large language model training to more than 10,000 GPUs,2024,NSDI’24】集成,后者是一个用于LLM训练的高性能框架。为了支持大规模训练和通信,我们还应用了以下优化。
GQA(Group Query Attention):GQA已成为现代LLM(如LLaMA3和Mistral)中不可或缺的特性,它有助于减少KV头的数量,从而降低分布式注意力(dist-attn)的通信量。本文提到的所有系统都应用了GQA技术。
带打包的Dist-attn:由于工作负载与注意力掩码的面积成正比,按顺序将序列划分到设备上会导致工作负载不平衡。已有几种技术【4, 23, 31】被提出来解决这个问题。然而,它们不适用于打包序列的特殊分段因果注意力掩码。如图14所示,为了避免CP组内的异构计算和通信,我们优化了当前的dist-attn。打包序列的每个子序列被均匀地分成$2P$部分,并对称地分配给$P$个设备。这确保了每个设备持有所有子序列的$1/P$,并覆盖了注意力掩码面积的$1/P$。所有设备参与相同的环形P2P通信,数据交换量相同。
基线:我们的系统构建在MegaScale之上,这是一个用于大规模GPU集群的生产级LLM训练框架,已证明其性能优于DeepSpeed和Megatron-LM。因此,我们通过在三种情况下进行比较来展示ByteScale的优势:
模型和数据集:我们使用密集和稀疏LLM评估我们的工作,如表1所示。对于密集模型,我们选择了四种不同大小的LLaMA系列LLM:LLaMA-7B、LLaMA-13B、LLaMA-30B和LLaMA-70B。对于稀疏模型,我们选择了两种不同大小的Mistral系列LLM(MoE):Mistral-8x7B(激活参数=13B/47B)和Mistral-8x22B(激活参数=39B/141B)。实验中使用了两个数据集,即GitHub和Byted,我们已在§3.1中介绍过。图4展示了这两个数据集的数据分布。
Table 1. 用于评估的模型
我们首先通过测量每个训练步骤的平均吞吐量来评估三种方法的端到端性能,总体结果如图17所示。结果表明,朴素HDP和平衡HDP解决方案都优于基线,最大实现了7.89倍的加速。
数据集差异:Byted数据集比GitHub数据集包含更多的长序列,全局批次中有37%的令牌长于256K。因此,平均吞吐量和加速比低于GitHub数据集。然而,由于ByteScale为长短序列都提供了通信优化,加速比仍可达到4.26倍。
并行策略差异:像LLaMA-7B、13B和30B这样的模型使用包括HDP和TP在内的并行策略,因此应用DP-Balance策略。相比之下,像LLaMA-70B、Mistral-8x7B和Mistral-8x22B这样的模型采用HDP、TP和PP,我们应用PP-Balance策略。可以观察到,与PP-Balance相比,带DP-Balance的HDP实现了更高的加速比。例如,在GitHub数据集和2M上下文长度下,带DP-Balance的HDP的加速比在6.21倍-7.89倍之间,而带PP-Balance的HDP的加速比仅在3.42倍-4.28倍之间。如图13所示,DP-Balance策略只需要在每个时间步平衡计算,这比PP-Balance策略要求的在所有时间步平衡计算更容易实现。
为了更深入地剖析ByteScale的卓越性能,我们选择Byted数据集,并在一个有1024个GPU的集群上训练LLaMA-7B,上下文长度为2M。图18展示了单个训练步骤中不同ranks的详细运行时状态。
通信密集型案例:首先,我们从集群中随机选择4个ranks,并记录它们在每个方法的训练步骤中的前向和后向时间。如图18(b)所示,对于基线,微批次数量设为8,我们必须设置$P_{cp}=256$以支持2M的序列长度。可以观察到这4个ranks表现出相似的执行时间。这是因为大多数微批次(除了第三个)不具有$O((2M)^2)$的计算复杂度,但必须处理2M的通信量。如图18(a)所示,P2P通信时间远超计算时间,导致一个微批次的执行时间几乎由通信决定(占总时间的97.6%)。
计算不平衡案例:在朴素HDP解决方案下,全局批次内的序列按所需的最小ranks数量进行分片。如图18(a)所示,一个312K的序列仅由39个HDP ranks分片作为微批次,因此计算时间可以与通信开销重叠。然而,由于ranks之间的不平衡,训练效率仍然存在问题。如图18(b)所示,尽管第三个rank在1分41秒内完成了它的8个微批次,但它必须等待第一个rank在4分32秒时完成,导致171秒的空闲时间。即便如此,朴素HDP解决方案相比基线节省了4分8秒。
平衡案例:在平衡HDP解决方案下,所有ranks几乎同时完成执行。如图18(b)所示,在任何时间步,每个rank都被分配了具有相似FLOPs的微批次,而执行时间较短的ranks(如第三和第四个rank)将被分配更多的批次。因此,该步骤的总时间进一步减少到2分37秒,相比基线节省了6分3秒。
总体比较:如图18(c)所示,我们记录了所有1024个GPU在单个步骤中的有效计算时间。可以发现,朴素HDP解决方案相比基线将峰值执行时间减少了1.7倍,但ranks之间存在显著的时间差异,最大值和最小值之间相差4.7倍(min=60s, max=279s, std=68s)。平衡HDP解决方案消除了时间差异,从而相比朴素解决方案进一步将执行时间减少了2.3倍。
为了探究ByteScale中每个组件的有效性,我们使用与§8.3相同的配置进行了消融实验,如图20所示。
平衡策略的有效性:如图19所示,平衡策略稳定了RDMA流量,并使张量核心利用率持续保持在40%左右。这表明计算和通信的硬件单元在超过两个小时内持续满负荷工作而没有空闲。因此,平衡HDP解决方案将加速比从2.01倍提高到3.69倍,超过了任何其他策略带来的改进。
远程数据加载器的有效性:我们采用图15所示的远程加载器,并使用CPU预取来将数据读取与计算重叠。这种方法进一步将加速比从3.69倍提高到3.89倍。