一张 GPU 放不下模型时,多张 GPU 怎么分工?

数据并行、张量并行、流水线并行与 ZeRO 到底在拆什么

目录

  第五课算过:7B 模型采用混合精度 AdamW 训练时,仅参数、梯度和优化器状态就可能接近 112 GB,还没有算激活值。

  一张 80 GB GPU 放不下,就要让多张 GPU 分担。但“多卡训练”不是一种固定做法,关键要先回答:究竟要拆什么?

数据并行拆 Batch,张量并行拆一层里的计算,流水线并行拆模型层,ZeRO/FSDP 拆参数、梯度和优化器状态。

多 GPU 并行方式

四种多 GPU 分工方式,点击查看原始 SVG

1、为什么需要多卡

  多卡通常解决两类问题:一是单卡放不下,二是单卡算得太慢。

问题 目标
显存容量不足 把模型状态或激活分散到多卡
训练时间太长 让多卡同时处理更多计算

  多一张 GPU 不等于速度直接翻倍。GPU 之间需要传输数据,通信会占时间;切分不均匀还会让部分 GPU 等待。

2、数据并行:拆 Batch

  Data Parallel 最直观:每张 GPU 都保存完整模型,但处理不同的样本。

GPU 0:完整模型 + Batch A
GPU 1:完整模型 + Batch B
GPU 2:完整模型 + Batch C
GPU 3:完整模型 + Batch D

  每张卡分别完成 Forward 和 Backward,然后通过 All-Reduce 汇总梯度,使所有模型副本保持一致。

  优点是实现简单、吞吐量高;缺点是每张卡仍要保存完整模型。如果模型本身已经放不下,纯数据并行无能为力。

3、张量并行:拆一层计算

  Tensor Parallel 把同一层中的大矩阵拆到多张 GPU。例如一个线性层的权重矩阵可以按列或按行切开:

\[Y=XW,\qquad W=[W_0\;W_1]\]

  GPU 0 计算 $XW_0$,GPU 1 计算 $XW_1$,最后合并结果。Attention 的多个 Head、FFN 的大矩阵也可以按类似思路切分。

  它能让单层超大参数分散到多卡,但层内经常需要通信,因此通常要求 GPU 之间有较快连接。

4、流水线并行:拆模型层

  Pipeline Parallel 把不同层放到不同 GPU:

GPU 0:Layer 1~8
GPU 1:Layer 9~16
GPU 2:Layer 17~24
GPU 3:Layer 25~32

  数据像流水线一样依次通过各张卡。为了减少后面的 GPU 空等,Batch 会再切成多个 Micro Batch,让不同阶段同时工作。

  流水线的难点是 Bubble:开始和结束时,总有部分阶段没有工作。如果各阶段计算量不均匀,等待会更加明显。

5、ZeRO/FSDP:拆训练状态

  普通数据并行会在每张卡复制完整参数、梯度和优化器状态。ZeRO 与 FSDP 的思路是:既然这些副本重复,为什么不分片保存?

ZeRO 阶段 主要分片内容
Stage 1 优化器状态
Stage 2 优化器状态 + 梯度
Stage 3 优化器状态 + 梯度 + 参数

  FSDP 与 ZeRO Stage 3 的核心目标相近:平时每张卡只保存自己负责的参数分片,需要计算某层时临时 All-Gather,用完再释放或重新分片。

  这样能显著降低单卡显存,但会增加通信量,并让训练流程更复杂。

6、现实中如何组合

  大型训练通常不是四选一,而是组合使用:

节点之间:Data Parallel / FSDP
节点内部:Tensor Parallel
模型很深:再加入 Pipeline Parallel

  这种组合有时被称为 3D Parallelism。选择方案时要同时看模型大小、Batch Size、GPU 数量、显存容量和互联带宽。

通信到底在传什么

操作 直观含义 常见场景
All-Reduce 汇总后让每张卡都拿到结果 数据并行同步梯度
All-Gather 收集各卡分片,拼成完整数据 FSDP 临时获取参数
Reduce-Scatter 求和并把结果分片发回 分片梯度
Send / Receive 从一个阶段传给下一阶段 流水线并行

  计算越快,通信越可能成为瓶颈。因此 GPU 数量、网络拓扑和切分策略必须一起考虑。

四种方式对比

方法 拆什么 单卡有完整模型吗 主要代价
数据并行 Batch 是 梯度同步
张量并行 层内矩阵 否 频繁层内通信
流水线并行 模型层 否 Bubble 与调度
ZeRO / FSDP 训练状态 否或临时聚合 参数/梯度通信

记住这 5 件事

  1. 多卡训练既为了解决显存,也为了提高吞吐。
  2. 数据并行每张卡都有完整模型,只拆 Batch。
  3. 张量并行拆一层,流水线并行拆不同层。
  4. ZeRO/FSDP 重点消除训练状态的重复副本。
  5. 并行度越高,越需要权衡计算、通信和等待时间。

第六课复习总图

大模型第六课复习总图

第六课复习总图,点击查看原始 SVG

  下一课,我们可以继续看:FP32、FP16、BF16 和混合精度为什么能影响速度、显存与训练稳定性?