使用 PyTorch Fully Sharded Data Parallel 加速大模型训练

Hugging Face 已将 PyTorch Fully Sharded Data Parallel (FSDP) 集成到 Accelerate 库中,允许从业者通过在数据并行工作节点之间分片优化器状态、梯度和参数,从而训练显著更大的模型。通过支持更大的 batch size 以及通过 CPU offloading 训练原本会超出 GPU 显存限制的模型,这一集成使大模型训练变得更加普及。

FSDP vs. Distributed Data Parallel (DDP)

PyTorch FSDP 通过消除每个 GPU 上全量模型副本带来的冗余显存消耗,对 Distributed Data Parallel (DDP) 进行了改进。

DDP 中,每个工作节点都维护模型参数、梯度和优化器状态的完整副本。虽然每个工作节点处理不同的数据 batch,但它们在更新模型之前必须执行 all-reduce 操作以平均所有工作节点之间的梯度。

FSDP 中,优化器状态、梯度和模型参数在工作节点之间进行分片。在正向和反向传播过程中,FSDP 使用 all-gather 操作仅检索特定层或封装模块所需的参数,并在计算完成后立即释放它们。随后,局部梯度通过 reduce-scatter 操作进行平均和分发,从而允许每个工作节点仅更新其局部的参数分片。这种方法极大地减少了每个 GPU 的显存占用。

GPT-2 性能基准测试

Hugging Face 在因果语言建模 (Causal Language Modeling) 任务上,使用两块 NVIDIA Titan RTX GPU (每块 24GB) 对 FSDP 与 DDP 进行了基准测试。

GPT-2 Large (762M Parameters)

FSDP 与 DDP 相比可以实现显著更大的 batch size。在不使用 CPU offload 的情况下,FSDP 支持高达 15 的 batch size (相比之下 DDP 为 7)。启用 CPU offload 后,batch size 进一步增加到 22。虽然在原始训练时间方面,使用混合精度 (FP16) 的 DDP 最快,但 FSDP 为大 batch size 提供了所需的显存效率,这对于具有动态 batching 的应用特别有益。

GPT-2 XL (1.5B Parameters)

对于 GPT-2 XL 模型,DDP 即使在 batch size 为 1 时也会出现 CUDA Out of Memory (OOM) 错误。相比之下,FSDP 实现了成功的训练:

  • FSDP (Zero-Stage 3): 在 2 块 GPU 上支持每块 GPU 的 batch size 为 5。
  • FSDP with CPU Offload: 在单块 GPU 上支持 batch size 为 10 的训练,在 2 块 GPU 上支持每块 GPU batch size 为 14 的训练。

技术实现与配置

通过 Accelerate 集成

用户可以通过 accelerate config CLI 或通过 FullyShardedDataParallelPlugin 进行更细粒度的控制。关键配置选项包括:

  • Sharding Strategy:FULL_SHARDSHARD_GRAD_OP 之间进行选择。
  • Min Num Params: 一个层被默认 auto-wrap 策略封装所需的最小参数数量。
  • Offload Params: 一个布尔值,用于确定是否应将参数和梯度卸载到 CPU。

Auto Wrap Policy 的作用

min_num_params 设置对于显存优化至关重要。当使用 default_auto_wrap_policy 时,如果一个层的参数数量超过指定阈值,FSDP 就会对其进行封装。在 BERT-Large (330M) 上的基准测试显示,使用 auto-wrap 的 FSDP 消耗的显存大约是 DDP 的一半。降低 min_num_params (例如降低到 2k) 与较高的阈值 (例如 1M) 相比,可以略微进一步减少显存使用。

关键注意事项与限制

使用 FSDP 的从业者必须意识到以下技术约束:

  • Optimizer Initialization: FSDP 会将参数展平并进行原地分片。因此,模型必须在创建优化器 之前 通过 accelerator.prepare(model) 进行准备。在封装模型之前创建优化器可能会导致优化器失效或导致显存占用增加。
  • Parameter Groups: 由于 FSDP 将嵌套模块展平为 1D 数组,在封装之前创建的参数组 (例如,对偏置项应用不同的 weight decay) 会被合并为一个单一组并丢失。
  • Multiple Models: 在训练多个模型时,必须在模型创建各自的优化器之前先准备模型,以避免错误。
  • Mixed Precision: 在本文发布时,由于 PyTorch 尚未修复相关问题,FSDP 还不支持混合精度。

分布式训练方法总结

FSDP 是更广泛的分布式训练策略生态系统的一部分,旨在处理海量模型:

  • ZeRO (Zero Redundancy Optimizer): FSDP 的基础,对优化器状态 (Stage 1)、梯度 (Stage 2) 和参数 (Stage 3) 进行分片。
  • Tensor Parallelism: 将单个大层的参数分片到多个 GPU 上。
  • Pipeline Parallelism: 将不同的层分布在不同的 GPU 上并进行微批次 (micro-batches) 流水线化处理。
  • 3D Parallelism: ZeRO、Tensor 和 Pipeline 并行化的结合,用于拥有数百亿参数的模型。

Sources

相关