外观
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 计算流程
- GPU 阶段: 进行前向和反向计算,将梯度传给 CPU。为提高效率,GPU 可在反向传播阶段边计算新梯度边将填满的 bucket 传输给 CPU——当反向传播结束时,CPU 已拥有最新梯度值
- CPU 阶段: 进行参数更新,将更新后的参数传回 GPU
mermaid
graph LR
A[GPU: FWD-BWD] -->|gradient 16| B[CPU: Update Super Node]
B -->|parameter 16| A提示: 这种 GPU/CPU 协同设计,将计算密集部分留在 GPU,将参数更新这种相对轻量但与优化器状态紧密耦合的操作放在 CPU,实现高效卸载。