Skip to content

Q130 · 显存优化技术 ​

团队有两张标称 24 GB 显存的 GPU,想训练一个虚构的 10 亿参数模型。有人算“参数用 FP32,一份权重约 4 GB,两张卡共 48 GB,怎么可能装不下?”实际开始训练,单卡却在反向传播前后报 OOM(out of memory,显存不足)。原因是训练不仅要放权重,还要放梯度、优化器历史状态、前向中间结果和临时缓冲;普通数据并行让两张卡各自保存完整副本,并不能简单把两卡显存合并成一大块。

解决时不要只背“用 ZeRO”。先识别哪一类东西占空间:前向中间结果太大,可以少存、反向时重算;每卡有重复的参数/梯度/优化器状态,可以跨卡分片;合适的计算和张量可用较短的数据格式;还可以把部分状态放到 CPU 内存或 NVMe 存储。每种节省都把代价移到别处——额外计算、跨卡通信、数值验证或设备间搬运。本题关注这些优化机制与选择顺序;模型为什么单卡放不下的完整账本,见训练工程挑战。PyTorch Activation Checkpointing · DeepSpeed ZeRO

本文数字全部是教学假设:10 亿个可训练参数,FP32 权重与梯度,Adam 有两份 FP32 历史统计;每卡处理一份固定长度的小批次。为方便算术,1 GB = 10^9 字节。24 GB 是名义设备容量,驱动、框架和临时分配还要占空间;算出低于 24 GB 也不等于真实训练一定能跑。

术语:显存里放着哪些东西 ​

术语或符号初学者可以怎样理解在本文中的对应物
参数 / 权重模型学到、训练时不断更新的数字10 亿个 FP32 数,每个 4 字节,约 4 GB
梯度反向传播算出的“每个参数应如何调整”的信号同样假设为 10 亿个 FP32 数,约 4 GB
Adam 优化器状态为平滑和调整更新幅度,Adam 保存的一阶、二阶历史统计两份各 4 GB,合计 8 GB;本例未额外算 FP32 主权重
激活 / activation前向计算途中产生、反向求梯度可能要用的中间张量假设一个小批次的峰值占 8 GB
工作缓冲算子、通信、框架等在计算时临时申请的空间为教学固定记 2 GB,真实值会变
FP32 / BF16 / FP16每个数分别用 32 位或 16 位保存;16 位通常是 2 字节混合精度只让选定运算和张量变小
前向 / 反向前向算预测与损失;反向沿计算过程求梯度激活在前向产生,反向可能需要它
数据并行 / DDP多卡各处理不同样本,但通常各保留整套训练状态两张卡各自持有 4+4+8 GB,不是平分
ZeRO / FSDP在多个数据并行进程之间切分训练状态的方案两卡各保存一部分状态,需要按需通信
Offload / 卸载把部分状态从 GPU 显存挪到 CPU 内存或 NVMe留出显存,但可能等待数据搬回
峰值显存一次训练步骤中占用最高的时刻,不是空闲时读数反向或参数聚合时可能比平时高

这里的 GB 是近似十进制计数,没有考虑张量对齐、梯度累积、模型结构、注意力实现和不同优化器。若使用某些混合精度 Adam 配方,还可能保留一份 FP32 主权重,使模型状态比本例更多;因此不能把下文的 26 GB 作为 10 亿参数模型的普适答案。ZeRO 原论文对另一种常见的 16 位计算加 FP32 主权重口径给出每参数 16 字节的推算,说明字节数依赖具体配方。ZeRO 原论文

先把单卡 26 GB 的账算出来 ​

本例每卡项目算法估算占用
FP32 参数10 亿 × 4 字节4 GB
FP32 梯度10 亿 × 4 字节4 GB
Adam 一阶、二阶状态2 × 10 亿 × 4 字节8 GB
已保留的激活假设值,取决于批量和序列长度8 GB
工作缓冲假设值,实际可能更高2 GB
合计4 + 4 + 8 + 8 + 226 GB

两张卡若使用普通数据并行,每张仍是约 26 GB,因此都可能超过各自的 24 GB 名义容量。这比把 26 × 2 = 52 GB 与总容量 48 GB 比较更有用:OOM 发生在某张卡的具体峰值上。若用同一模型更长的序列,8 GB 激活假设会变;若换优化器,8 GB Adam 状态假设也会变。定位时应先在代表性长序列上测每卡峰值,并区分张量已占用与缓存分配器保留的显存。PyTorch CUDA 内存监测说明

激活重算、ZeRO/FSDP 分片、混合精度和 CPU Offload 各自减少不同显存,并带来计算、通信、数值或传输代价

图的四格是同一显存账本上的四种不同取舍,不是必须依次执行的步骤。分片格中的两张 GPU 示意状态被分摊,不表示各卡都还有完整三份;真实 ZeRO 阶段决定到底分哪一项。混合精度格只表示部分计算张量可由 32 位变为 16 位,绝不意味着把总显存直接除以二。

激活重算:少留前向中间结果,反向再算 ​

正常训练要保留某些前向结果,以便反向传播计算梯度。Activation Checkpointing(激活检查点,也叫激活重算)让程序在选定一段网络的前向结束后,只保存重建该段所需的边界张量,暂不长期保存段内许多中间激活;反向走到此处时,重新运行该段前向计算,得到所需中间值,再继续求梯度。这里的“checkpoint”是计算边界,不是把训练状态写到磁盘以便断点续训。PyTorch 官方直接把它定义为“用额外计算换内存”。PyTorch torch.utils.checkpoint

在本例中,若实测原来 8 GB 激活经过合适的分段后降到约 3 GB,而其他项不变,则教学预算变成 4 + 4 + 8 + 3 + 2 = 21 GB/卡,名义上进入 24 GB 以内。“3 GB”是示范性测量目标,不是 PyTorch 保证的压缩率:仍要留边界输入、输出和未被重算的张量;不同层、微批量、序列长度会让结果变化。代价是反向时多做计算,训练步骤可能变慢。若重算段在第二次执行时读取了不同的全局状态、随机行为处理不当或含副作用,甚至可能与原前向不等价,PyTorch 文档对此有明确警告;需要核对损失和梯度,而不只看 OOM 是否消失。PyTorch Activation Checkpointing

适用情形是激活占比高,例如长序列、较大的每卡微批量。若账本主要被参数和 Adam 状态占满,重算激活的收益有限;这时应考虑分片或卸载。它也可与分片组合,但节省的仍是同一激活项,不能把两种方案宣称的“百分比”直接相加。

ZeRO / FSDP:把重复的训练状态分给多张卡 ​

**ZeRO(Zero Redundancy Optimizer)**在数据并行组内逐阶段切掉多卡重复保存的状态。Stage 1 只分优化器状态;Stage 2 再分梯度;Stage 3 再分参数。阶段越高,单卡平时需要保留的完整副本越少,但计算某层时仍须拿到所需权重,并在反向后合并或分发梯度,因此需要跨卡通信和临时缓冲。DeepSpeed 官方给出的三个阶段正是这个递进关系。DeepSpeed ZeRO 教程 · ZeRO 原论文

以两卡理想均分、仍用 FP32 与同一激活假设计算,下面每一行都从原来的 26 GB 单独启用一种阶段;“其他缓冲 2 GB 不变”仅为比较,真实 all-gather 缓冲可能增大。

每卡方案参数梯度Adam 状态激活其他理想化合计
普通数据并行4488226 GB
ZeRO Stage 1448 ÷ 2 = 48222 GB
ZeRO Stage 244 ÷ 2 = 248220 GB
ZeRO Stage 34 ÷ 2 = 2248218 GB

FSDP(Fully Sharded Data Parallel)是 PyTorch 的分片数据并行实现。它的 FULL_SHARD 会分片参数、梯度和优化器状态,机制上可类比 ZeRO Stage 3;SHARD_GRAD_OP 对梯度和优化器状态做分片,参数在计算中的保留时机又与 FULL_SHARD 不同,可类比 ZeRO Stage 2 的省内存方向。两者不是同一个配置项或完全相同的时序,不能把上表数字当作 PyTorch FSDP 的承诺。FSDP 文档明确说明,FULL_SHARD 会在前向前 all-gather 所需参数,并在反向后以 reduce-scatter 分片梯度;这些术语分别指从各卡拼出计算所需权重、归并并重新分发梯度份额。PyTorch FSDP 文档

分片也有失败路径:从上表看 Stage 3 “只需 18 GB”,真实运行仍可能在某层权重临时聚合、通信预取或优化器更新时越过 24 GB;一个特别大的未拆分模块也会造成短时峰值。此时要看峰值发生在哪个阶段、哪个卡、哪个模块,调整分片单元与通信桶、减少微批量或结合激活重算。不要只拿训练空闲时的 18 GB 截图证明方案可行。跨机网络慢时,Stage 3 的节省也可能被通信等待抵消;要同时记录每步耗时与吞吐。PyTorch FSDP 文档

混合精度:让合适的张量更短,而非全部减半 ​

FP32 每个数约 4 字节,FP16 / BF16 每个数约 2 字节。混合精度让适合低精度的计算与张量使用 16 位格式,而数值敏感的运算或优化器状态仍可保持 32 位。PyTorch 的 AMP(Automatic Mixed Precision,自动混合精度)通过 autocast 为不同运算选择合适精度;FP16 训练还常需要梯度缩放,降低很小的梯度变成零的风险。BF16 与 FP16 都是 16 位,但表示范围不同,实际数值行为及硬件支持需分别验证。PyTorch AMP · PyTorch AMP recipe

只看本例激活这一项的假设性上限:若那 8 GB 全是可用 16 位替代的 FP32 数,且其余状态维持原口径,这一项理论上变 4 GB,合计 4 + 4 + 8 + 4 + 2 = 22 GB/卡。但真实 AMP 会选择性转换算子,中间可能需要额外拷贝、FP32 主权重或不同精度的梯度;8 GB 激活也未必全由可减半的张量构成。22 GB 只是“只改一项”的算术示范,不是开启 autocast 后的预报值。若与重算一起用,它们可能都作用于同一激活,所以不能再把“省 4 GB”和“省 5 GB”不加分析地合并。

混合精度的失败不一定是 OOM,更可能是损失变为 NaN(非数字)、出现无穷大、梯度溢出或验证质量下降。遇到这种情况,要定位是哪些运算和数据触发,检查梯度缩放、精度类型及模型数值范围,必要时让敏感运算回到 FP32;不能因为显存下降就宣布训练正确。硬件或框架不支持某种 dtype(数据类型)时,也不能只改一个配置字符串就认为已生效。PyTorch AMP

Offload:把空间换到 CPU 内存或磁盘 ​

CPU Offload 把某些训练状态放在 GPU 之外,需要时搬到 GPU,或直接在 CPU 上完成优化器更新。DeepSpeed 配置文档分别列出优化器状态卸载和参数卸载;优化器状态可卸到 CPU 或 NVMe,参数卸载的支持与阶段有关。NVMe 是较慢但容量更大的非易失存储设备,适合容量压力极高且可接受传输代价的配置;其路径、缓冲和主机容量必须实际满足要求。Offload 并不是删除状态,数据只是搬家。DeepSpeed ZeRO 配置 · ZeRO-Infinity 论文

若只做一个概念算术:假设把本例每卡 8 GB Adam 状态完全移出 GPU,而其余四项维持,GPU 常驻账本从 26 变为 4 + 4 + 0 + 8 + 2 = 18 GB/卡,CPU 端则至少要容纳移走的状态。但这是理想常驻量,不含传输、预取、CPU 优化器和暂存缓冲的峰值;具体 DeepSpeed Offload 还要配合支持的 ZeRO 阶段,不能把这行直接当成一个可运行配置。若用 ZeRO Stage 1 先把两卡优化器状态各分成 4 GB,再卸载,则每卡搬走的是已分给它的 4 GB,不应重复声称又省 8 GB。CPU 内存不足会把 OOM 转移到主机;总线或 NVMe 太慢会让 GPU 等数据、每步时间变长。DeepSpeed ZeRO 配置

Offload 适合剩余显存不足、但主机内存和传输条件允许,且团队接受吞吐下降或能通过预取掩盖等待时使用。它与 ZeRO 分片、混合精度和重算可以组合,但要在同一口径下重新做账,并实测而非累加宣传数字。

把四种办法放回同一训练任务 ​

若这两张卡的 26 GB 预算要尽快变得可运行,我会先用最短的代表性训练步骤测峰值、分清是激活还是静态状态主导,再试一项有明确靶点的改动。假设实际确认激活是 8 GB,激活重算后实测约 3 GB,那么理想账是 21 GB;若仍需留更大余量,可试 ZeRO Stage 2:参数 4、梯度 2、Adam 状态 4、激活 3、其他 2,理想合计 4 + 2 + 4 + 3 + 2 = 15 GB/卡。这是一条可能成功的配置路径,不是保证单卡峰值 15 GB:分片临时聚合和运行时缓冲要再测。DeepSpeed ZeRO 教程 · PyTorch Activation Checkpointing

每改一项,都记录至少四类结果:每卡训练步骤峰值、每步时长或有效 token 吞吐、损失/梯度是否正常、与未优化小样本基线的结果是否一致。PyTorch 的 max_memory_allocated() 可看张量分配峰值,max_memory_reserved() 可看缓存分配器管理的峰值;外部工具看到的设备占用可能还包括框架外分配,不能把这几个数混为同一个指标。PyTorch CUDA 内存说明

若“15 GB 预算”仍报 OOM,先看是否是最长序列、某层 all-gather 或临时缓冲在反向时造成尖峰;不能直接断言 ZeRO 无效。若显存没问题但每步慢很多,看激活重算的额外计算、分片通信、CPU 搬运分别耗时多少。若出现 NaN,暂停对比 FP32 小样本基线和数值日志,不能以“跑得下”替代“训得对”。这样选择技术的依据是具体哪块显存最紧和团队能接受哪种代价,而不是四个开关全部打开。PyTorch FSDP · PyTorch AMP

面试时怎么回答 ​

我会先把显存分成参数、梯度、优化器状态、激活和临时缓冲,测训练步骤的单卡峰值,再按占比选技术。激活大就用 activation checkpointing,少存中间张量、反向重算,换来额外计算;多卡重复的模型状态大,就用 ZeRO 或 FSDP 分片,Stage 1 先分优化器、Stage 2 再分梯度、Stage 3 再分参数,换来通信和临时聚合;混合精度让合适的运算与张量用 16 位,但要保留数值检查,不能把全部显存简单除二;offload 把状态移到 CPU 或 NVMe,换来传输与主机容量压力。比如 10 亿 FP32 参数、Adam、8 GB 激活的虚构预算是每卡约 26 GB,两张 24 GB 卡做普通数据并行仍装不下。重算激活与 Stage 2 分片后理论账可到约 15 GB,但必须用真实峰值、吞吐和训练损失验证。

若追问“为什么两张 24 GB 卡不能直接训练一个需 26 GB 的副本”,答:普通数据并行每卡都保留完整训练状态,容量不能自动相加;必须真正做分片或改变每卡保存量。若问“ZeRO Stage 3 为什么还会 OOM”,答:按层计算时要临时聚合参数,并有通信桶和工作区,上表只算理想常驻量。若问“混合精度节省一定最大吗”,答:不一定;优化器仍可能是 FP32,激活也未必全能减半,还要看数值稳定。

资料依据 ​

最后更新2026-09-26
难度P1
频率high
阅读20 min
主题llm-training / gpu-memory / activation-checkpointing
觉得有帮助?把这个链接转给正在求职的朋友 · 用 Ctrl + K 全站搜索其它题