Skip to content

ZeRO 显存优化与 3D 并行策略详解

内容整理自学习笔记,仅供面试备考参考;不构成录用、培训或考试承诺。

1. 什么是 3D 并行?

3D 并行是三种分布式并行策略的组合,能够以非常高效的方式训练大型模型。这三种策略分别是:

策略简介优缺点
Data Parallel (DP)每张卡保存完整模型,数据拆分训练简单易实现,但显存占用大
Tensor Parallel (TP)张量按维度拆分到不同 GPU减少单卡计算量,通信频繁
Pipeline Parallel (PP)模型按层拆分到不同 GPU显存友好,存在流水线气泡

2. 三种并行策略对比

2.1 Data Parallel(数据并行)

每张卡保存一个完整模型副本,每次迭代将 batch 数据等分为 micro-batch,各卡独立计算梯度,通过 AllReduce 求梯度均值后各自更新参数。

# 2 张 GPU 的数据并行
GPU0: [L0 | L1 | L2]  参数: [a0 | b0 | c0]  [a1 | b1 | c1]
GPU1: [L0 | L1 | L2]  参数: [a0 | b0 | c0]  [a1 | b1 | c1]

2.2 Tensor Parallel(张量并行/横向并行)

每个张量被切分为多个分片,各分片在不同 GPU 上独立并行处理,步骤结束时再同步/规约结果。这也被称为横向并行(张量并行)。

# 张量并行示意
GPU0: [L0 | L1 | L2] 中的部分张量分片
GPU1: [L0 | L1 | L2] 中的部分张量分片

2.3 Pipeline Parallel(管道并行/纵向并行)

模型在多个 GPU 上按层垂直拆分,每个 GPU 处理管道的一个阶段,并行处理小批次数据。

# 8 层模型,2 张 GPU
GPU0: [L0 | L1 | L2 | L3]
GPU1: [L4 | L5 | L6 | L7]

3. 为什么需要 ZeRO?

虽然 Data Parallel 应用最广泛,但其要求每张卡存储完整模型副本,显存大小成为制约模型规模的主要因素

关键洞察: 是否可以让每张卡只存 1/N 的模型参数,合并后仍是一个完整模型?随着卡数增加,单卡显存占用降低,可训练的模型也更大。

ZeRO 系列技术是一种显存优化的数据并行方案,专为训练超大规模语言模型设计。

4. ZeRO 核心思想

去除数据并行中的冗余参数,使每张卡只存储一部分模型状态,从而显著减少显存占用。

5. 显存分配分析

ZeRO 将每张卡的显存内容分为两类:

类别组成占比
模型状态参数(Parameters)、梯度(Gradients)、优化器状态(Optimizer States)优化器状态占 75%
剩余状态激活值、临时缓冲区、显存碎片取决于模型和 batch size

示例: GPT-2 1.5B 参数,fp16 权重仅需 3GB 显存,但模型状态实际耗费 24GB!优化器状态是头号显存杀手。

6. ZeRO 优化三阶段

针对模型状态,ZeRO 使用分片策略——每张卡只存 1/N 的模型状态量,系统内仅维护一份完整状态。

阶段分片内容显存减少倍数通信量变化
ZeRO-1 (Pos)优化器状态分区4x与数据并行相同
ZeRO-2 (Pos+g)优化器状态 + 梯度分区8x与数据并行相同
ZeRO-3 (Pos+g+p)优化器状态 + 梯度 + 参数分区与 Nd 成线性关系通信量增加约 50%

示例: k=12,参数量 Φ=7.5B,Nd=64 张 GPU —— 随着 ZeRO 阶段深入,显存优化效果非常显著。

7. ZeRO Offload:显存不足,内存来补

ZeRO-Offload 将训练阶段的某些模型状态下放到 CPU 内存,利用廉价的内存扩展可用空间。

7.1 单层一次迭代的计算流程

混合精度训练下,前向计算(FWD)需要上一层的激活值和本层参数,反向传播(BWD)同样需要激活值和参数计算梯度。使用 Adam 优化器时,参数量为 M,权重为 2M(fp16)或 4M(fp32)。

节点分布:

  • GPU 端: FWD-BWD Super Node(融合节点),计算密集
  • CPU 端: Update Super Node(融合节点),与 Adam 状态交互

8. ZeRO Offload 计算流程

  1. GPU 阶段: 进行前向和反向计算,将梯度传给 CPU。为提高效率,GPU 可在反向传播阶段边计算新梯度边将填满的 bucket 传输给 CPU——当反向传播结束时,CPU 已拥有最新梯度值
  2. CPU 阶段: 进行参数更新,将更新后的参数传回 GPU
mermaid
graph LR
    A[GPU: FWD-BWD] -->|gradient 16| B[CPU: Update Super Node]
    B -->|parameter 16| A

提示: 这种 GPU/CPU 协同设计,将计算密集部分留在 GPU,将参数更新这种相对轻量但与优化器状态紧密耦合的操作放在 CPU,实现高效卸载。