显存不足时大模型训练的切分策略,核心思路是把模型、梯度、优化器状态拆开,按需分布到多张卡上;没有万能方案,只有按硬件和模型规模选组合。当一张显卡的显存被模型参数占满,报错“CUDA out of memory”时,很多人第一反应是换更大显存的卡,但在2026年,单卡显存增长远慢于模型参数膨胀,切分训练已经成为主流选择。
显存吃紧的真相:模型、梯度和优化器谁在占地方
你在跑一个10亿参数的模型时,以为显存是被参数挤爆的,其实参数只占一小块地方,FP16精度下,10亿参数约占2GB显存,真正让显存见底的是三件套:梯度、Adam优化器状态和混合精度下的master weights,把这四样加起来,训练阶段每10亿参数通常需要吃掉16GB到20GB显存,其中参数占比不到15%。
显存不足”这个问题的本质,不是模型本身放不下,而是训练动态数据把显存塞满了,这也能解释为什么一个模型能用单卡做推理,却没法用单卡做训练。
业内专家指出,理解这一点是选择切分策略的前提:如果你只分模型参数、不动优化器状态,显存压力几乎没有缓解,该崩还是崩。
多卡训练显存不够怎么办:先分清三种切分粒度
当单卡放不下完整训练状态时,大家讨论的“切分”通常分三个层面,各有各的切法。
数据并行把模型复制了无数份,反而更吃显存
数据并行是最古老的并行方式:每张卡存一份完整模型、一份完整梯度,每步只同步梯度,它的优点是不用改模型结构,缺点是显存占用随着卡数线性增长,当你一开始就报显存不足时,单纯加卡做数据并行没有用,因为每张卡仍然需要容纳完整训练状态。
张量并行把矩阵切开,适合层特别宽的模型
张量并行把一层内的矩阵运算拆成多块,由多张卡协同算完一个层,再传给下一层,GPU之间同步频繁,通信开销大,一般只在单机多卡、有高速NVLink或IB互联的环境里用,它对Transformer的MLP层和Attention头比较友好,但对CNN这类结构收益不大。
流水线并行按层切,显存省了,但卡间有空转
流水线并行把模型按层切成几段,每张卡只扛其中一段的前向和反向,显存占用随切段数下降,但段与段之间存在气泡时间,切得越多,利用率越低。
这三种是基础切法,实际训练中,几乎没人只用其中一种,而是把它们和梯度切分工具结合起来用。
主流切分策略的显存收益对比
| 策略 | 显存收益来源 | 通信成本 | 典型场景 |
|---|---|---|---|
| 数据并行 | 无(模型不切) | 每步全量梯度同步 | 小模型多卡提速 |
| 张量并行 | 单层参数分片 | 每层前向反向多次同步 | 大模型单机多卡 |
| 流水线并行 | 按层分段 | 段间激活传输 | 模型层数极多 |
| ZeRO阶段1 | 切分优化器状态 | 梯度通信前压缩 | 多卡训练显存不够时首选 |
| ZeRO阶段2 | 切分梯度+优化器状态 | 通信量约等于DDP | 单机多卡性价比高 |
| ZeRO阶段3 | 参数+梯度+优化器全切分 | 每层参数广播 | 模型大到单卡完全装不下 |
| FSDP | 与ZeRO阶段3类似 | 按需全收集参数 | PyTorch原生生态 |
| CPU offload | 把部分状态挪到内存 | CPU-GPU搬运 | 单卡显存不足的兜底方案 |
从这张表能看出,多卡训练显存不够怎么办的答案,往往不是某一种并行,而是“张量并行/流水线并行打底,再用ZeRO或FSDP把优化器状态摊到更多卡上”。
DeepSpeed ZeRO 和 PyTorch FSDP 哪个好?这里给出实际选择逻辑
这是2026年出现频率最高的问题之一,ZeRO出自DeepSpeed库,FSDP是PyTorch官方实现,两者在理念上高度相似:都认为每张卡只保存模型状态的一部分,用到时再拼回来。
ZeRO的阶段划分,决定了显存节省的上限
ZeRO把训练状态分成参数、梯度、优化器状态三组,分成三个阶段切分,阶段1只切优化器状态,阶段2额外切梯度,阶段3把参数也切开,阶段3下,单卡显存需求理论上随卡数线性下降,这也是它能跑动千亿级模型的原因。
ZeRO的代价在通信,阶段3下,每层前向计算前需要广播参数,反向后需要reduce梯度,通信次数远多于普通数据并行。
卡间互联速度不够快时,ZeRO阶段3可能比阶段2更慢,多数情况下,单机8卡用阶段2就能解决显存问题;只有跨机或模型超过单卡容量十多倍时才上阶段3。
FSDP把分片逻辑内建到PyTorch里,上手成本低
FSDP的切分逻辑和ZeRO阶段3基本一致,但它的优势在于和PyTorch生态无缝衔接,accelerate、transformers等库都有原生支持,如果你已经跑在PyTorch官方代码栈里,不想引入额外依赖,FSDP是更低风险的选项。
判据只有一个:你愿不愿意换训练框架
从效果看,两者在相同卡数和模型规模下的显存占用差异很小,行业共识认为,选择依据在于工程生态:DeepSpeed提供了更多内存优化插件和量化工具,适合改造成本高的老代码;FSDP适合新项目,因为它跟着PyTorch版本迭代,踩坑时更容易在官方社区找到答案。
要回答“DeepSpeed ZeRO 和 PyTorch FSDP 哪个好”,可以从这几个问题入手:
- 你的代码有没有用到原生的
DataParallel?有的话FSDP迁移更顺。 - 是否需要混合使用张量并行和流水线并行?DeepSpeed的集成方案更成熟。
- 团队成员熟悉哪套配置?维护比选型更重要。
4090训练大模型显存不足时,优先打开CPU offload
手头只有一张24GB显存的RTX 4090,想跑7B甚至13B模型训练,很多人的第一反应是“把batch size调小”,但batch size调到1仍然报错的情况很常见,因为训练7B模型所需的优化器状态已经超过40GB。
这时能续命的选项是CPU offload,它把优化器状态甚至梯度放在内存里,显卡只保留当前计算需要的参数副本,代价是CPU-GPU之间的PCIe带宽远低于显存带宽,每一步训练都变慢,但至少能跑起来。
实际操作中,配合accelerate的配置可以这样打开:
- 用
accelerate config进入交互配置,选择“CPU offload”。 - 把
offload_optimizer_device设为cpu。 - 开启
offload_param_device,让参数也按需搬运。 - 如果内存也吃紧,再开
offload_state_dict,把checkpoint写到磁盘。
在4090上打开offload后,7B模型的训练能从“完全跑不了”变成“跑得慢但稳定”,这不是大模型训练显存优化的最终方案,却是成本最低的过渡方案。
实操:从报错到跑通的四步切分路径
如果你现在正被“CUDA out of memory”卡住,按照下面四条路径排查,比盲试参数更有效。
第一步,用一张卡排除数据问题
先把batch size设为1,关闭梯度累积,确认模型本身能不能做一次前向反向,这一步仍然崩溃的话,说明模型结构或输入尺寸有问题,不关切分的事。
第二步,估算训练状态总量
用小batch跑一下,记录显存峰值,减去模型推理显存,得到优化器状态和梯度的占用,再用transformers的model.num_parameters()和显存占用公式估算总量,确认到底缺多少显存。
第三步,按卡数选策略
- 单卡且缺的显存不到一半,优先开CPU offload,动模型结构。
- 2到4张卡,用DeepSpeed ZeRO阶段2,或FSDP默认配置,显存普遍够用。
- 4到8张卡仍不够,把张量并行和流水线并行叠加进来。
- 超过8张卡或机房内网带宽受限,不要硬上ZeRO阶段3,改用模型并行加offload。
第四步,观察通信耗时占比
训练日志里的step time如果明显分成“计算时间+等待时间”,并且等待时间占比超过三成,说明切分粒度过细,调大zero_optimization.stage3_gather_16bit_weights_on_model_save这类参数,或在FSDP里改sharding_strategy为SHARD_GRAD_OP,往往比继续加卡更解决问题。
关于显存不足时模型切分的三个常见问题
切分训练会影响最终模型效果吗?
不会,切分只改变参数存储位置和通信方式,不改变前向反向的计算逻辑,相同随机种子和相同超参数下,切分训练与单卡训练的loss曲线基本一致,唯一需要留意的是,使用CPU offload后,浮点数在内存和显存间搬运可能带来极小的舍入偏差,但不影响模型收敛结果。
张量并行为什么在消费级显卡上表现一般?
张量并行要求每层计算时多张卡反复同步,消费级显卡之间通常只有PCIe连接,没有NVLink,带宽远低于数据中心卡,层数越多,同步次数越多,通信耗时越明显,在两张4090上跑张量并行,多数场景比单卡offload更慢,所以消费级显卡优先考虑ZeRO或FSDP,不碰张量并行。
首发原创文章,作者:王坚,如若转载,请注明出处:https://idctop.com/article/625648.html





