Accelerate ND-Parallel:高效多 GPU 训练指南

Accelerate ND-Parallel:高效多 GPU 训练指南

Hugging Face 已将 ND-Parallelism 集成到 Accelerate 和 Axolotl 中,使用户能够组合数据并行、完全分片数据并行、张量并行和上下文并行策略,以优化大规模模型的多 GPU 训练。

Accelerate ND-Parallel:高效多 GPU 训练指南

Hugging Face 在 Accelerate 和 Axolotl 中引入了 ND-Parallelism,提供了一种简化的方式,在单个训练脚本中组合多种并行策略——数据并行(DP)、完全分片数据并行(FSDP)、张量并行(TP)和上下文并行(CP)。这种集成使开发者能够在跨多节点 GPU 集群训练拥有数十或数百亿参数的模型时,优化内存使用与通信开销之间的权衡。

核心并行策略

数据并行(DP)

数据并行在每个设备上复制完整的模型、梯度和优化器状态。每个设备处理不同的子批次数据,梯度在所有设备之间同步后再更新参数。此方式提升吞吐量,但要求整个模型能够放入单个 GPU。

完全分片数据并行(FSDP)

FSDP 将模型权重、梯度和优化器状态在 GPU 之间进行分片,从而降低每个设备的内存占用。进行前向或后向传播时,FSDP 会收集特定层(通常是 transformer 解码器块)所需的权重,随后再重新分片。此方式以增加通信开销为代价,显著降低峰值内存使用。

张量并行(TP)

张量并行将大型线性层(如前馈层或注意力投影)在设备之间拆分。与 FSDP 的动态分片不同,TP 创建静态的内存分区。由于 TP 需要频繁的激活同步,它在使用高带宽链接(例如 NVLink)的单节点内部最为高效,不建议在通过 PCIe 连接的 GPU 上使用。

上下文并行(CP)

上下文并行将输入序列在 GPU 之间进行分片,以处理因注意力的二次增长而导致的极长序列长度,防止超出 GPU 内存。通过 RingAttention,每个 GPU 持有查询、键和值矩阵的一部分,并在 GPU 环形结构中循环传递键值分片,确保每个查询能够对整个序列计算注意力得分,同时分摊计算和内存负载。

ND-Parallelism:组合策略实现多节点扩展

多节点训练常常受到节点间延迟和内存限制的瓶颈。ND-Parallelism 将集群视为多维拓扑结构,以优化通信。

混合分片数据并行(HSDP)

HSDP 是一种 2D 并行方法,在节点内部使用 FSDP(利用快速的节点内链接),在节点之间使用 DP。它将慢速的节点间通信压缩为一次梯度同步步骤,从而提升吞吐量,但相较于纯 FSDP 会增加内存使用。

FSDP + 张量并行

将 FSDP 与 TP 结合,即通过 FSDP 在节点之间对模型进行分片,并在单节点内部通过 TP 拆分层。此方式降低了 FSDP 的延迟,允许训练单个设备无法容纳的超大层,并可降低全局批次大小。

FSDP + 上下文并行

该 2D 策略用于训练极长序列。当 CP 已经与 FSDP 集成时,在 CP 之上再加入 FSDP 可进一步降低模型权重和优化器状态所需的内存预算。

混合分片数据并行 + 张量并行

该 3D 层次结构使用 DP 在节点组之间复制模型,使用 FSDP 在组内对模型进行分片,并使用 TP 在每个节点内部拆分层。此配置为适配特定硬件和扩展约束提供了最大的灵活性。

实现与使用注意事项

在 Accelerate 和 Axolotl 中的配置

用户可以通过 Accelerate 中的 ParallelismConfig 类或 Axolotl 中的特定配置字段来设置这些策略:

  • dp_shard_size:FSDP 的分片度
  • dp_replicate_size:DP 的复制度
  • tp_size / tensor_parallel_size:TP 的并行度
  • cp_size / context_parallel_size:CP 的并行度

内存与稳定性优化

  • CPU RAM 高效加载:对于单设备无法容纳的模型,启用 cpu_ram_efficient_loadingSHARDED_STATE_DICT(位于 FullyShardedDataParallelPlugin 中)至关重要。
  • 有效批次大小:有效批次大小计算公式为 micro_batch_size * gradient_accumulation_steps * dp_world_size,其中 dp_world_size = (dp_shard_size * dp_replicate_size) / tp_size
  • 学习率缩放:随着有效批次大小的增大,学习率应线性缩放或采用平方根缩放,以保持训练稳定性。
  • 梯度检查点:通过在反向传播时重新计算中间激活来以计算换取内存。它可以将激活内存降低 60-80%,但训练时间约增加 20-30%。

Sources