特色项目

¹ 伊利诺伊大学厄巴纳-香槟分校 SSAIL 实验室,² Anyscale,³ Snowflake

摘要: AutoSP 可以自动将标准 Transformer 训练代码转换为序列并行代码,从而在多个 GPU 上进行长上下文大语言模型(LLM)的训练。它与 DeepSpeed 集成,在几乎不增加运行时开销的情况下,相比手动编写的基准代码,显著提升了可训练的最大上下文长度。

如今,大语言模型(LLM)正越来越多地被用于处理极长上下文的任务,其标记(token)数量可超过 10 万个。在如此大的规模下,即使使用 ZeRO/FSDP 等常规训练技术对设备数量进行扩展,也会出现内存不足(OOM)的问题。为了规避这些问题,序列并行(Sequence Parallelism, SP)——即将输入标记在不同设备间进行分区,从而通过增加 GPU 数量实现长上下文训练——成为了一种常用的并行训练技术。

然而,实现 SP 难度极高,需要对 DeepSpeed 或 HuggingFace 等现有库进行侵入式代码更改。这些更改通常涉及对输入标记上下文(及中间激活值)进行分区、插入通信集合,以及实现通信与计算的重叠,所有这些步骤都需要在正向和反向传播中同时完成。这导致希望探索长上下文能力的科研人员必须投入大量精力去设计系统堆栈,且针对不同的硬件供应商,这项工作往往需要重复进行。

为了避免这种复杂性,我们推出了 AutoSP:一种基于编译器的全自动化解决方案。它能自动将易于编写的训练代码转换为多 GPU 序列并行代码,高效利用 GPU 进行长上下文训练,并可与现有的并行策略(如 ZeRO)兼容。这免去了开发者为长上下文训练而反复修改训练管道的繁琐工作。现在,用户只需导入 AutoSP 并使用 AutoSP 后端编译任意模型,即可让任何人都能拥有长上下文训练的能力。此外,通过将该技术嵌入编译器,我们的方法实现了性能可移植性:可以在多种硬件上实现高性能的 SP。

本文结构如下:(1)AutoSP 及其如何帮助模型科学家实现长上下文训练;(2)AutoSP 的关键设计决策;(3)AutoSP 的核心结果,展示其易用性与影响力;(4)AutoSP 的局限性及功能边界。

AutoSP 使用方法

AutoSP 的核心设计理念是简单易用,为用户屏蔽了多 GPU 编程的大部分复杂性。为此,我们在 DeepCompile 中实现了 AutoSP:这是一个 DeepSpeed 内部的编译器生态系统,旨在通过编程方式为深度神经网络训练启用各种优化。通过它,任何使用 DeepSpeed 的用户几乎都可以毫不费力地启用序列并行。接下来我们看一个示例。

# We instantiate a deepspeed config.
# Assume 8 GPUs with 2 DP ranks and 4 SP ranks.

config = {
    "train_micro_batch_size_per_gpu": 1,
    "train_batch_size": 2,
    "steps_per_print": 1,
    "optimiser": {
        "type": "Adam",
        "params": {
            "lr": 1e-4
        }
    },
    "zero_optimization": {
        "stage": 1, # AutoSP interoperates with ZeRO 0/1.
    },
    # Simply turn on deepcompile and set
    # the AutoSP pass to be triggered on.
    "compile": {
        "deepcompile": True,
        "passes": ["autosp"]
    },
    "sequence_parallel_size": 4,
    "gradient_clipping": 1.0,
}

# Initialise deepspeed with model.
model, _, _ = deepspeed.initialize(config=config,model=model)

# Compiles model and automatically applies AutoSP passes.
model.compile(compile_kwargs={"dynamic": True})

for idx, batch in enumerate(train_loader):
    # Custom function that we expose within:
    #     deepspeed/compile/passes/sp_compile.
    inputs, labels, positions, mask = prepare_auto_sp_inputs(batch)

    loss = model(
        input_ids=inputs,
        labels=labels,
        position_ids=positions,
        attention_mask=mask
    )

    ... # Backwards pass, optimiser step etc...

如上述示例所示,用户只需在现有的单设备训练代码上进行以下操作:(1)使用 prepare_autosp_input 工具函数(在 DeepSpeed 中提供)对输入标记、注意力掩码和位置 ID 进行轻量级标记,供 AutoSP 进行程序分析;(2)调整 DeepSpeed 配置以开启 DeepCompile,并将 “passes” 标志设置为 “autosp”。其余工作均通过编译模型时调用的 AutoSP 编译器传递来处理,它会自动启用序列并行以及其他长上下文训练优化。此外,AutoSP 还可以开箱即用地自动与 ZeRO Stage 1 组合,只需在设置 AutoSP 标志的同时在 DeepSpeed 中设置 ZeRO-1 标志即可结合两种策略。

AutoSP 编译器传递

由于 AutoSP 通过转换用户代码来支持更长的上下文训练,为了保持透明度,我们简要介绍一下 AutoSP 的关键设计点、代码转换方式及其对用户的影响。

序列并行代码转换:AutoSP 自动将单 GPU 代码转换为多 GPU 序列并行(SP)代码。AutoSP 转换的具体 SP 策略是 DeepSpeed-Ulysses。我们特别选择 DeepSpeed-Ulysses 而非其他策略(如 RingAttention),是因为在 NVLink 网络拓扑或胖树(fat-tree)网络中,其通信开销随 GPU 数量的增加保持恒定。不过,DeepSpeed-Ulysses 仅能将 SP 大小扩展至模型头部数量的规模(对于 7-8B 模型,通常为 32)。

用于长上下文训练的激活检查点(Activation Checkpointing):AutoSP 还应用了一种专为长上下文建模定制的自定义激活检查点(AC)策略。AC 会释放计算成本较低的算子的中间激活值,并在反向传播时根据需要重新计算它们以获取相关梯度。PyTorch 2.0 引入了一种基于最大流最小割的自动化 AC 公式,但我们发现对于长上下文建模来说,这种方法过于保守。因此,我们引入了一种针对长上下文训练的新型 AC 策略:序列感知 AC(Sequence-aware AC, SAC),它利用了长上下文特有的 FLOP 动态特性。当开启此功能(AutoSP 中的默认设置)时,训练吞吐量会略有下降。然而,如果不使用它,长上下文训练将无法实现,因此用户可以根据需要选择仅在发生 OOM 的配置中开启此功能。

在真实模型上评估 AutoSP

为了验证 AutoSP 的可行性,我们在不同规模的模型上进行了 NVIDIA GPU 评估,证明其易用性几乎不牺牲任何运行性能。我们在拥有 8 个 A100-80Gb SXM 的节点上对不同的 Llama 3.1 模型进行了基准测试。我们使用 PyTorch 2.7 和 CUDA 12.8,将 AutoSP 与以下手动编写的 torch 编译基准进行了对比:RingFlashAttention、DeepSpeed-Ulysses 和 ZeRO-3。关键结果总结在下图中。

AutoSP 不仅能在同等资源下增加最大可训练序列长度(左图——越高越好),而且这些优势几乎是以不牺牲运行性能为代价的(右图——越低越好)。

局限性

AutoSP 有两个主要局限性。首先,我们要求用户将 Transformer 强制编译为单个可编译工件。有时,PyTorch 用户可能会分别编译许多函数并将它们拼接到一个模型中。这在 AutoSP 中是不被允许的,因为我们需要编译并查看整个模型,以便正确地对输入序列进行分片并将该信息传播到整个图中。其次,我们不允许可编译工件中存在任何图中断(graph breaks)。这会使信息的分析和传播变得复杂,我们将把提高 AutoSP 对图中断的恢复能力作为未来的研究方向。

结论

AutoSP 使用户能够轻松扩展任意 Transformer 训练代码以实现序列并行,并利用自定义的 AC 策略增强长上下文训练。通过与 DeepSpeed 的集成,用户只需更改配置文件,即可轻松使用现有的 DeepSpeed 训练代码进行长上下文训练。我们准备了端到端的示例,供用户在真实模型工作负载(如 Llama 3.1 8B)上进行尝试,地址在此。快来试用一下,看看长上下文训练现在变得多么简单。