ByteScale: Efficient Scaling of LLM Training with a 2048K Context Length on More Than 12,000 GPUs

发表时间: 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)

A1 主要贡献

为了满足现代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,一个为大规模长短序列混合训练设计的高效、灵活且可扩展的训练框架。其主要贡献如下:

A3 背景知识与关键观察

2. 背景知识

2.1 Transformer与大语言模型

Transformer架构【40,Attention is All you Need,2017,NeurIPS 2017】已成为当今大语言模型(LLMs)【5, 14, 32, 39】最流行和广泛使用的基础架构。它通常由一系列Transformer层组成,每层包含一个注意力模块和一个前馈网络(FFN)模块。如图1所示,自注意力机制需要序列中的所有令牌参与计算以捕获整个文本的上下文信息。相比之下,其他操作如归一化、线性投影和激活函数则执行令牌级计算,允许每个令牌独立处理。

图1. Transformer层的架构
图1. Transformer层的架构

2.2 分布式LLM训练

随着模型大小和训练数据的持续扩展,分布式训练技术在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令牌是常见做法)。受硬件内存限制,一次性处理整个大批量是不可行的。梯度累积将每个全局批次(即每个训练步骤中采样的数据)分成多个微批次。这些微批次的梯度被累积起来,等同于一次性处理整个全局批次所产生的梯度。

2.3 填充与打包

为了在当前静态并行策略中支持可变长度的序列,需要使用填充(padding)和打包(packing)等技术。如图2所示,填充将同一批次中的序列填充到相同长度,但这会导致计算浪费。打包【22,Efficient sequence packing without cross-contamination: Accelerating large language models without impacting performance,2021,CoRR abs/2107.02027】将多个序列连接成一个单一序列,不含填充令牌。它采用一种特殊的分段注意力掩码,以确保每个序列在自注意力中被独立处理。

图2. 序列填充与打包
图2. 序列填充与打包

2.4 长上下文训练

由于自注意力的时间和内存复杂度均为$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)所示。

图3. 带打包的上下文并行
图3. 带打包的上下文并行

3. 观察与动机

3.1 数据异质性

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过程中训练集上的平均响应长度,表明逐渐增加和多样化的响应长度有助于提高模型性能。

图4. 两个数据集中的样本和令牌分布
图4. 两个数据集中的样本和令牌分布

3.2 冗余通信

现有系统在整个训练过程中应用静态并行策略。通常,它们假设所有(打包的)序列长度相同,并设置一个固定的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等级,会导致过度的冗余通信。当序列长度高度偏斜时,这个问题会更加严重。

3.3 计算不均衡

图6. 计算不均衡
图6. 计算不均衡
图5. 不均衡的数据和流水线并行
图5. 不均衡的数据和流水线并行

A2 方法细节

4. ByteScale 概览

我们提出ByteScale来解决这些挑战。如图7所示,它由三个主要组件构成。Profiler用于分析环境、模型配置、数据分布,并为其他组件构建成本模型。Communication Optimizer通过数据感知分片、动态通信和选择性卸载,为长短序列提升通信效率。Balance Scheduler通过并行感知的数据分配来解决计算不均衡问题。

图7. ByteScale概览
图7. ByteScale概览

5. 通信优化器

本节描述ByteScale如何优化通信开销。首先,它通过动态序列分片和通信减少短序列的冗余通信。其次,它通过选择性卸载进一步压缩长序列的通信成本。

5.1 数据感知分片与通信

图8. HDP示意图
图8. HDP示意图
图9. 令牌级梯度
图9. 令牌级梯度

5.2 数据感知的选择性卸载

图10. 逐层激活值卸载
图10. 逐层激活值卸载
# 列表1. act_ctx的用法
with act_ctx(offload_ratio=0.5):
    # 前向传播
    hidden_states = model_layer(hidden_states)
# 后向传播
loss.backward()
图11. 数据感知的选择性卸载
图11. 数据感知的选择性卸载

5.3 整体流程

ByteScale的整体流程在算法1中概述。简而言之,该算法遍历全局批次中的每个序列$S_i$。对于长序列,它推导出卸载比率$\rho$并确定所需的ranks数量$n(S_i)$(第1-6行)。对于短序列,它将它们打包以填满每个rank的容量$C$(第7-9行)。处理后的序列随后被分配给$P_{hdp}$个ranks,算法返回微批次和offload_ctx以供执行(第10-12行)。

算法1:朴素HDP解决方案
算法1:朴素HDP解决方案

6. 平衡调度器

本节我们介绍平衡调度器,以解决DP和PP的不平衡问题。通过精心编排数据分配(替代算法1中的第10行),它在保持§5中实现的最小通信的同时,缓解了这些不平衡。我们将首先概述几个关键见解,然后提出我们的启发式解决方案。

6.1 重新定义微批次

梯度累积要求不同的DP ranks执行相同数量的微批次,这是基于所有微批次具有相同计算负载的假设。然而,如§3.3所述,不同微批次的执行时间可能显著不同。在ByteScale中,我们重新定义了一个更灵活的策略,允许不同的HDP ranks处理不同数量的微批次(大小相同但工作负载不同),以缓解不平衡问题。如图13所示,这使得所有ranks几乎同时完成计算。更重要的是,这个策略不影响模型收敛。无论序列如何分配给HDP ranks,我们最终都计算全局批次中所有令牌的梯度总和,如§5.1所讨论的,这确保了数学上的等价性。

图13. 平衡策略
图13. 平衡策略

6.2 解决PP不平衡

图12. 平衡的数据和流水线并行
图12. 平衡的数据和流水线并行

6.3 解决DP不平衡

6.4 平衡策略

算法2描述了平衡策略。首先,我们按长度降序对全局批次B中的序列进行排序。然后将这些有序序列划分为FLOPs总和近似相等的桶,因此平均长度较长的桶包含的序列较少(第3-5行)。其次,我们确定那些执行时间较短的ranks,以便后续分配(第7-9行)。第三,如果使用DP-Balance策略,我们从同一个桶中选择序列。否则,如果使用PP-Balance策略,我们从所有桶中顺序选择序列。实践中,执行时间较短的ranks会被分配更多的序列(第12-15行)。最后,我们重复第二和第三步,直到所有桶都为空。

算法2:HDP的平衡策略
算法2:HDP的平衡策略

7. 实现细节

ByteScale基于Python、C++和CUDA实现了约16000行代码,并已与MegaScale【18,MegaScale: scaling large language model training to more than 10,000 GPUs,2024,NSDI’24】集成,后者是一个用于LLM训练的高性能框架。为了支持大规模训练和通信,我们还应用了以下优化。

图14. 为打包序列优化的Dist-attn
图14. 为打包序列优化的Dist-attn
图15. 远程数据加载器
图15. 远程数据加载器
图16. 融合的SoftmaxCrossEntropy
图16. 融合的SoftmaxCrossEntropy

A4 实验环境与结果

8.1 实验设置

Table 1. 用于评估的模型

8.2 端到端评估

我们首先通过测量每个训练步骤的平均吞吐量来评估三种方法的端到端性能,总体结果如图17所示。结果表明,朴素HDP和平衡HDP解决方案都优于基线,最大实现了7.89倍的加速。

图17. 端到端评估(单位:每秒令牌数)
图17. 端到端评估(单位:每秒令牌数)

8.3 案例研究

为了更深入地剖析ByteScale的卓越性能,我们选择Byted数据集,并在一个有1024个GPU的集群上训练LLaMA-7B,上下文长度为2M。图18展示了单个训练步骤中不同ranks的详细运行时状态。

图18. 案例研究
图18. 案例研究

8.4 消融研究

为了探究ByteScale中每个组件的有效性,我们使用与§8.3相同的配置进行了消融实验,如图20所示。

图19. 网络流量和张量核心利用率
图19. 网络流量和张量核心利用率
图21. 激活值卸载的有效性
图21. 激活值卸载的有效性