☰
大模型训练显存优化:FSDP状态分片与并行网格实战指南
2026/9/28 14:48:10 网站建设 项目流程

做预训练的人基本都经历过这样的场景:代码在单卡上跑得好好的,模型一换大,迎面就是 CUDA out of memory。更难受的是,你隐约知道要“切”,但到底切什么、怎么切、切完图上是什么样子,脑子里一片浆糊。我在被 13B 模型的优化器状态反复折磨之后,才把 “FSDP 状态分片” 和 “整张并行网格” 这两件事彻底串起来。这篇东西适合两类人:一是模型快要超出单卡显存、正准备往多卡训练迁移的工程师;二是已经用 FSDP 跑通小模型,但吞吐上不去、想搞清楚并行维度到底怎么排的人。我会按自己踩坑的顺序来写:先算清楚显存账,再讲 FSDP 到底切了什么,然后把 TP、PP、DP 怎么合成一张网格、怎么落到启动配置里讲明白,最后给一份可以直接照抄的排查速查表。

1. 先想明白:大模型预训练为什么非切不可

1.1 显存账本:一个 7B 模型在单卡上到底吃了多少

很多人以为显存不够只是因为参数太大,其实真相是:预训练阶段真正吃满显存的,往往不是参数本身,而是优化器状态。

按最常见的混合精度 + AdamW 组合来算,每个参数要占用大约 16 字节:

  • 前向用的 FP16/BF16 参数副本:2 字节
  • 更新用的 FP32 主权重副本:4 字节
  • Adam 一阶动量(m):4 字节
  • Adam 二阶动量(v):4 字节
  • 反向传播用的梯度副本:约 2 字节

7B 模型一乘就是 7 × 10^9 × 16,约 112 GB。这张账算完,80GB 显存的 A100/H100 连“状态”都放不下,更别提激活值。13B 直接翻到 208 GB 以上。所以“切分”不是可选项,而是一道数学题:状态总量除以显存,决定了你至少需要多少张卡。

我自己的习惯是先算这一层账,再考虑训练策略。卡数不够,后面所有并行方案都是空中楼阁。

1.2 三种切法:切数据、切层、切算子

显存账算完,下一步是搞明白一个分布式训练系统里其实只有三种基本切法,所有花里胡哨的并行方案都是它们的组合。

切数据是最好理解的一种:每张卡放一份完整(或完整状态)的模型副本,每个 step 喂不同的 batch,最后把梯度同步一下。传统 DDP 做的就是这件事,但它要求每张卡能放下完整模型状态,所以在大模型场景下很快就失效了。

切层叫做流水并行(Pipeline Parallel,PP):把 Transformer 的几十上百层切成几段,每段挂在不同的卡上。数据像工厂流水线一样依次流过各段,代价是“流水线气泡”——首尾两端的卡经常空转。PP 的通信量不大,但负载均衡很难做完美。

切算子叫张量并行(Tensor Parallel,TP):把一次矩阵乘法的权重按行或列切碎,多张卡合力算同一个算子。这是对带宽和延迟最敏感的一种切法,因为每个 Linear 前向后向都要跨卡同步一次,所以 обычно 只放在 NVLink 这种高速互联的范围内。

预训练里的“切”,从来不是选一种,而是三种按比例叠加。FSDP 在概念上是切数据的一种升级版,但它的切入点不是“复制模型”,而是“把模型状态本身也切碎”。

2. FSDP 状态分片:把参数、梯度、优化器状态切成碎片

2.1 FSDP 是什么:ZeRO-3 在 PyTorch 里的实现

FSDP 全称 Fully Sharded Data Parallel,最早的设计思路来自 DeepSpeed 的 ZeRO 系列。ZeRO 的三个阶段是一层层递进的:ZeRO-1 只切优化器状态,ZeRO-2 把梯度也切了,ZeRO-3 把参数、梯度、优化器状态全部切碎,每张卡只持有完整世界的一份 N 分之一。FSDP 的 FULL_SHARD 策略,本质上就是 ZeRO-3 在 PyTorch 里的官方实现。

FSDP 的实现很有一套:它把模型的参数先扁平化(flatten)成一维大张量,再切成一个个 FSDP unit。前向计算的时候,需要哪部分权重就把那部分 shard 通过 all-gather 拼回完整权重;反向传播算完梯度之后,再用 reduce-scatter 把梯度按同样的分片方式散回去。之后每个 rank 只需要用自己手里那一小片参数和梯度去更新对应的优化器状态。

为什么要 flatten?因为 all-gather 这种集合通信喜欢“一大块”数据,几十个小 tensor 分别去通信,会带来大量消息开销。把参数压成一维大张量后,一次通信能传几 GB,效率高得多。这也是 FSDP 性能优化里很关键的一个设计决策。

实际操作里还有个很重要的开关:auto_wrap_policy。FSDP 默认不会把整个模型包成一个巨型 FSDP unit,而是建议按 Transformer 层来包,每层一个 unit。这样每个单位的数据量适中,通信能和计算重叠,而且前向算完一层就能释放一层的完整权重,内存水位会被压得很低。

2.2 三种分片策略,到底怎么选

PyTorch FSDP 提供了几种 ShardingStrategy,我平时常用的就下面这几个:

策略分片范围显存节省单 step 通信模式适合场景
NO_SHARD什么都不切几乎不省all-reduce 梯度单卡能放下模型,只是想要分布式加速
SHARD_GRAD_OP梯度 + 优化器状态中等reduce-scatter + all-gather显存差一点,但不想牺牲通信效率
FULL_SHARD参数 + 梯度 + 优化器状态最多all-gather + reduce-scatter7B 以上,显存吃紧的常规选择
HYBRID_SHARD节点内 FULL_SHARD,节点间 DDP多节点内集合通信 + 节点间 all-reduce跨节点带宽不足,又想全局状态分片

我个人的默认选择很简单:单机多卡跑 7B,直接用 FULL_SHARD;一旦超过一个节点,先看清楚节点间是 400G IB 还是 25G 网卡,前者可以继续 FULL_SHARD,后者最好换成 HYBRID_SHARD,或者引入 TP、PP 把跨节点的通信量压下来。不要迷信“FSDP 是万能解药”,它对互联的要求比你想象的高。

2.3 通信账本:FSDP 的代价到底贵在哪

FSDP 的显存收益人人都知道,但很少有人算过它的通信账单。假设 FSDP 组里有 N 个 rank,模型的字节数是 M,那么前向阶段每个 unit 都要做一次权重 all-gather,反向阶段要做一次梯度 reduce-scatter。一个训练 step 跑完,每张卡收发的大致数据量约等于:

通信量 ≈ 2 × M × (N − 1) / N

N 很大的时候,这个数接近 2M。也就是说,模型权重的字节数乘以 2,就是一个 step 里每张卡要“吞下去”的通信数据。

7B 模型用 FP16 训练,M 约 14GB,64 卡纯 FSDP 的话,每张卡每个 step 要搬接近 28GB 的数据。节点内 NVLink 能扛,一旦到了跨节点 25GbE(约 3GB/s),光通信就要 9 秒以上,训练就彻底废了。这就是为什么大模型训练里几乎不可能只靠纯 FSDP 撑起超大集群——不是显存不够,是通信链路先崩溃。

2.4 实战里最容易踩的 FSDP 配置坑

第一坑:把整个模型包成一个 FSDP unit。很多新手搜到示例代码,直接FSDP(model)一把梭。这样做的后果是 all-gather 的峰值就是全模型大小,显存水位被瞬间拉满,7B 模型甚至可能在 8 卡上 OOM。正确做法是按 Transformer 层包,让单位足够小。

第二坑:只开 FSDP 不开激活检查点(activation checkpointing)。显存账本里我只算了参数、梯度和优化器状态,实际训练里激活值同样能吃几十 GB。FSDP 把状态从 112GB 压到 14GB 之后,激活值反而成了主要矛盾。7B 以上我基本必开 activation checkpointing,用一点重计算换显存,很划算。

第三坑:混淆了 CPU offload 的代价。cpu_offload=true确实能把显存占用进一步压低,但代价是每层前向都要从 CPU 取参数,训练速度经常掉 30% 到 50%。我只在“差最后几 GB 显存”的时候开它,正常情况下宁可缩小 batch size。

第四坑:checkpoint 保存方式没想清楚。默认的 FULL_STATE_DICT 会把全量参数汇聚到 rank 0,再写出去,大模型经常直接把 CPU 内存撑爆,而且保存时间极长。用 SHARDED_STATE_DICT 按 rank 分别落盘,配合 torch.distributed.checkpoint 加载,才是大模型该有的做法。

3. 把策略画成一张网格:TP、PP、DP 怎么合成一个整体

3.1 设备网格到底是什么

“一张网格”不是比喻,而是一个真真实实的数据结构。假设你有 64 张卡,把它们想象成一个三维网格:(DP=4, PP=2, TP=8)。每个进程在这个网格里都有一个三维坐标 (dp_id, pp_id, tp_id),三个维度的乘积必须等于总的卡数 4 × 2 × 8 = 64。

网格的意义在于:它把“模型怎么切”和“卡之间怎么连”这两件事统一定义清楚了。FSDP 的通信组由 DP 维度决定,TP 的通信组落在节点内的 8 张卡上,PP 的通信发生在管道相邻段之间。网格确定了,通信拓扑、权重摆放、数据切分方式全都能推导出来。

PyTorch 生态里,要叠加 TP + FSDP 的时候会用到torch.distributed.tensor.DeviceMesh,或者直接在 Megatron 这类框架里传--tensor-model-parallel-size和--pipeline-model-parallel-size。不管用什么工具,底层画的都是这张网格。

我觉得把并行策略想成网格最大的好处是:你不必在脑子里维护一堆乱七八糟的通信组,只需要在纸上画一个三维立方体,每一维的长度标清楚,所有卡的角色就一目了然了。

3.2 画网格之前,先收集四个输入

网格不是拍脑袋定的,我每次设计都会先把四个参数摆出来。

第一个是模型状态总字节数。用 1.1 节的 16 字节每参数估算,7B 约 112GB,70B 约 1.12TB。这个数决定了至少要多少卡。

第二个是 GPU 数量和单卡显存。比如 8×80GB 和 64×80GB,方案完全不一样。

第三个是卡间拓扑,这是最容易忽略的。节点内 NVLink、节点间 IB、普通以太网,三种链路的带宽差一个数量级。TP 必须放在 NVLink 内,FSDP 能跨节点但要会算账,PP 可以把大传输切成一次性的激活而不是每层都通信。

第四个是模型结构细节。PP 的段数要能整除总层数,TP 维要能整除注意力头数和 MLP 结构里的并行度。比如 70B 的 LLaMA 系模型有 80 层,TP=8 跟 32 头整除得很好,但 TP=6 就会很尴尬。

3.3 两个典型网格,一步一步步推出来

先看最简单的:7B 模型,8 张 A100,单机。

状态总量 112GB,8 卡平均分到 14GB,加上激活检查点大概 20-30GB 的激活占用,80GB 卡完全放得下。单机内 NVLink 带宽充足,FSDP all-gather 的 28GB 通信量也能扛住。所以这张网格我直接画成 (DP=8, PP=1, TP=1),纯 FULL_SHARD,配置简单,生态支持最好,速度也不差。

再看一个重一点的:70B 模型,8 个节点、每节点 8 卡,一共 64 张 A100 80GB。

先算显存:1.12TB 状态总量,64 卡均摊是 17.5GB,容量上没问题。但通信上纯 FSDP 的账很难看:每卡每 step 要搬约 275GB 数据,横跨节点的部分会把 IB 带宽打穿。这时候就需要画网格了。

我把网格画成 (DP=4, PP=2, TP=8)。TP=8 放在节点内,NVLink 扛住权重矩阵并行;PP=2 把 80 层切成两段,每段 40 层,放在不同节点上;剩下 4 个节点副本走 DP/FSDP 维度。这样 FSDP 的 all-gather 范围从 64 卡缩到 4 卡,每步通信量锐减到约 26GB,同时 TP 的 all-reduce 又全部落在 NVLink 内。显存方面,由于 TP 和 FSDP 叠加,每个 rank 持有参数的 1/32,状态量约 35GB,加上激活也能装进 80GB。

网格里每一维为什么是这个数,背后都有计算支撑,不是随手填的。

3.4 网格设计时最常见的两个幻觉

第一个幻觉是“FSDP 卡越多越省心”。确实,分片越多单卡显存越低,但通信量是随 rank 数线性涨的。我在 2.3 节算过,纯 FSDP 的通信量接近 2M,翻一倍卡数,每个 step 的通信量也跟着翻倍。所以卡多了以后,正确思路是用 TP/PP 把“必须跨节点通信的那部分”切出去,而不是一味扩大 FSDP 组。

第二个幻觉是“TP 越大越好”。TP 确实能把单层权重切成小块,但代价是每个算子前后都要 all-reduce 同步。TP 维超过节点内的 8 张卡之后,all-reduce 就要走跨节点网络,延迟直接爆炸。而且小 hidden size 的模型切 TP=8 本身就有碎片问题,同步次数太多,收益反而被通信延迟吃光。

4. 实操:从 7B 起步,把网格落到启动命令里

4.1 最小可跑的 FSDP 启动方案

理论说再多,不如一个能跑的命令。先给单机 8 卡的最小方案:

torchrun --nproc_per_node=8 train_fsdp.py

核心 Python 部分大致是这个结构:

import torch import torch.distributed as dist from torch.distributed.fsdp import ( FullyShardedDataParallel as FSDP, ShardingStrategy, ) from torch.distributed.fsdp.wrap import ( transformer_auto_wrap_policy, ) dist.init_process_group("nccl") # 按 Transformer 层作为 FSDP unit policy = transformer_auto_wrap_policy( transformer_layer_cls={LlamaDecoderLayer} ) model = build_model_7b() model = FSDP( model, sharding_strategy=ShardingStrategy.FULL_SHARD, auto_wrap_policy=policy, device_id=torch.cuda.current_device(), mixed_precision=bf16, )

如果用的是 Hugging Face 生态,accelerate 的 FSDP 配置会更直观,YAML 大概长这样:

compute_environment: LOCAL_MACHINE distributed_type: FSDP fsdp_config: fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP fsdp_transformer_layer_cls_to_wrap: LlamaDecoderLayer fsdp_sharding_strategy: FULL_SHARD fsdp_backward_prefetch: BACKWARD_PRE fsdp_forward_prefetch: false fsdp_offload_params: false fsdp_state_dict_type: SHARDED_STATE_DICT fsdp_min_num_params: 1000000 mixed_precision: bf16

注意我在这里用的是 bf16 而不是 fp16。A100 以上设备训练大模型,bf16 的指数位更宽,loss 不会动不动就溢出,是最省心的选择。

4.2 网格和分组在代码里怎么对应

网格不能只存在于脑子里。还是拿 64 卡、(DP=4, PP=2, TP=8) 的例子,8 个节点物理排布大概是这个样子:

DP 副本PP=0PP=1
DP=0node0,TP 组覆盖 rank 0-7node1,TP 组覆盖 rank 8-15
DP=1node2,TP 组覆盖 rank 16-23node3,TP 组覆盖 rank 24-31
DP=2node4node5
DP=3node6node7

在这张表里,FSDP 的 all-gather 只发生在同一行、同一列的 4 个节点之间;TP 的 all-reduce 只在同一节点内的 8 卡之间;PP 的激活传递则是 node0 → node1 单向传输。用 Megatron 类框架,对应的命令行参数就是:

--tensor-model-parallel-size 8 \ --pipeline-model-parallel-size 2 \ --data-parallel-size 4

网格最怕的是“逻辑上对,物理上乱”。如果 TP 组没有放在同一节点内,网格设计得再漂亮也无济于事。所以每次启动新实验,我第一件事就是检查 rank 到物理节点的映射,确保网格坐标和实际拓扑一致。

4.3 训练现场该盯哪些指标

网格画完、代码跑起来,不代表万事大吉。我盯训练健康度主要看几个东西:吞吐量(tokens/s 或 samples/s)、单卡 GPU 利用率、显存水位、loss 曲线。

GPU 利用率这个指标最容易骗人。某张卡利用率低,很可能不是它在摸鱼,而是在等通信。用nvidia-smi看到卡上显存占用大起大落、利用率周期性跌零,基本就是通信瓶颈的特征。想精确定位,用torch.profiler抓一段训练循环,看看 NCCL kernel 的时间占比,超过 30% 就要回头调整网格了。

吞吐和收敛要一起看。有些人为了提升吞吐疯狂加大 global batch size,结果 loss 掉不下去、训练发散,白跑几天。梯度累积不是免费的,它改变的是优化器的更新频率和噪声特性,不是简单的“多跑几千步就行”。

5. 常见翻车现场与排查速查表

5.1 还没开始就 OOM:先查 all-gather 峰值

我见过很多次“明明算着每卡只要 14GB,一跑就 OOM”的情况,多半是 FSDP unit 太大导致的。全模型一个 unit 的话,前向第一步就要 all-gather 出完整的全模型权重,峰值瞬间拉满。解决办法是把 unit 粒度降到 Transformer 层,并打开激活检查点。如果还是差一点,再考虑把 batch size 减半,而不是急着加卡。

5.2 吞吐上不去:卡越多,反而越来越慢

这种问题十有八九出在通信拓扑和网格不匹配。纯 FSDP 跨节点 all-gather 是第一个怀疑对象;TP 组跨节点是第二个。我会先看一眼带宽利用率,再用 HYBRID_SHARD 替代 FULL_SHARD,或者把 TP 组限制在节点内、把跨节点维度改成 PP。这里有个反直觉的点:减少参与 all-gather 的卡数,比增加卡数更能提速。

5.3 checkpoint 加载失败:多半是状态字典类型不匹配

用 SHARDED_STATE_DICT 存了档,加载的时候一个 FULL_STATE_DICT 读进来,直接报错。这是最典型的分片状态和整体状态混用问题。统一用分布式 checkpoint 的 API 读写,不要混着来。另外保存前一定记得检查各个 rank 的元数据一致性,文件散落不均匀,恢复的时候你是看不出来的。

5.4 排查速查表

现象可能原因检查手段常见修复
启动即 OOMFSDP unit 太大,all-gather 峰值高打印各 unit 参数量按 Transformer 层 wrap,降低 min_num_params
第一个 step OOM激活值峰值过高开启 activation checkpointing减小 micro batch size
吞吐随卡数增长停滞跨节点 FSDP 通信瓶颈torch.profiler 看 NCCL 占比改用 HSDP,或加 TP/PP 压通信
GPU 周期空闲流水线气泡太大看 pipeline schedule micro batch 数增加 micro batch,拉平气泡
checkpoint 保存 OOMFULL_STATE_DICT 汇聚到单卡查看 rank 0 内存改用 SHARDED_STATE_DICT
微调阶段 loss 不降网格没变,但 batch、LR 没适配对比基线的 loss 曲线调整学习率和 warmup

最后再分享一个我自己的习惯:动任何大模型训练之前,先在白纸上把网格画出来,标清每一维长度、每张卡的角色、每次通信走的链路。这个动作花不了十分钟,但能避免绝大多数“配置写错、跑了一星期才发现”的惨剧。FSDP 是内存问题的解药,但网格设计才是吞吐问题的根本。模型越大,越要靠这一张图把通信、显存、计算安排得明明白白。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询