大模型训练显存估算与混合精度实战:从账本到调优

发布时间:2026/9/29 15:51:08
大模型训练显存估算与混合精度实战:从账本到调优 大家在各种技术群里问得最多的问题一上来通常就是这张 24G 的卡能不能跑 7B那两张 80G 的卡能不能上 13B说实话这种问题会让人心头一紧——不是问题本身有问题而是提问的人大概率还不清楚显存到底消耗在了哪里。大模型训练的显存估计不是一个拍脑袋的经验公式它是一笔可以算得很细的账而混合精度训练之所以能成为标配恰恰是在这笔账最肥的两块肉上下刀子的技术。这篇文章我就把这两件事铺开揉碎从显存账本的每一项构成开始再到混合精度训练背后的数值原理、Loss Scaling、BF16 与 FP16 的取舍最后给出一套可以对着自己配置直接套用的显存估算方法和踩坑经验。文章面向的是准备训练或微调大语言模型的工程师和学生如果你已经跑通过小模型、现在想上更大参数量或者卡着显存调 batch size那这篇就是给你整理的实战笔记。下面所有规律和数据都来自我实际跑过的训练任务以及对比过源码实现后的结论不是抄来的表。1. 显存账本训练时 GPU 内存的四个大头是怎么花掉的很多人以为显存主要被模型权重占着于是拿参数量乘以 2FP16或者乘以 4FP32就完事了。真正训练起来才发现这样算出来的数字连实测的一半都不到。模型权重只是显存账单的第一行后面还排着梯度、优化器状态、激活值这三头大象。1.1 不仅是权重梯度本身占一份还得按参数对齐在一轮正常的反向传播里每个模型参数都会产生一个梯度梯度的数据形状和参数完全一致。也就是说权重有多少字节梯度至少也要占同样多的字节。如果你用 FP16 做混合精度训练权重可能是 FP162 字节但梯度为了保证数值稳定通常会攒在 FP324 字节里这块开销直接翻倍。这里很多人会踩一个坑用单卡跑小模型的时候不觉得一到多卡并行就会怀疑“怎么每个进程显存占比这么高”。原因在于梯度在多卡场景下不是每个 rank 各存一份完事的AllReduce 通信后每个进程都得保留完整的梯度用于下一步更新。显存账本的第一个修正就是梯度是实打实按参数数量逐字节计算的没有侥幸空间。以 7B 模型为例单参数 FP32 梯度就是 7B × 4 字节 28GB。注意这一项还没算权重本身。如果你用的是 AdamW还需要再往下看。1.2 优化器状态Adam 那一堆动量变量最容易被忽视优化器状态是显存估计里最容易被漏掉、同时占比最大的一项。SGD 很省钱只需要一阶动量但 AdamW 是现在训练大模型的绝对主流它要保存一阶动量 m 和二阶动量 v两个变量都按参数形状存储。如果优化器状态用 FP32那么 7B 模型仅优化器状态就是 7B × 4 字节 × 2 56GB。加上权重和梯度光这三样就 98GB 往上了。这也是为什么 NVIDIA 在 Megatron-LM 里做分布式优化器状态切分DataParallel 的 ZeRO 变体时第一步动手砍的就是这块肉。很多工程师第一次看到 ZeRO 的分配图才反应过来原来 FP32 的 Adam 状态才是吃掉显存的“大房东”16GB 显存的卡想训 7B 模型不把优化器状态切走根本不可能。做一个简单的对比表大家感受更直观7B 参数量单位 GB存储项目数据精度单卡原始占用未优化模型权重FP16/FP3214 / 28梯度FP3228Adam 一阶动量FP3228Adam 二阶动量FP3228权重总和-98FP32 权重下不要小看这一节因为后面所有的显存估计技巧——混合精度、梯度检查点、ZeRO——本质上都在“先弄清楚谁是大头再想办法把大头削掉”。1.3 激活值层数深了比权重还吓人第四个大头是激活值也就是前向传播时每一层保留下来、供反向传播使用的中间张量。很多人初学时以为只有权重和梯度占显存结果用 batch size 稍微调大一点就 OOM原因基本都在激活值。激活值的规模怎么理解Transformer 每一层的激活值大小大约是batch_size × 序列长度 × 隐藏维度 × 若干倍数而且很多模型层数有 32 层甚至 80 层每一层都要保存。拿一个典型配置举例batch size 为 8、序列长度 2048、隐藏维度 4096、层数 32仅 LayerNorm、Attention 计算中的 QKV、MLP 中间结果这些要用于反向传播的张量累积起来会是权重和梯度的数倍。一种应对是减小 batch size但这会让吞吐率掉得很厉害另一种是梯度检查点Gradient Checkpointing它把前向计算时的激活值丢掉反向传播时再重新算一遍用算力换显存。这个我在后面会专门讲。在进入混合精度细节之前大家脑子里先有一个总账的概念显存 模型权重 梯度 优化器状态 激活值 通信与临时缓冲如果算出来跟实测差一大截优先怀疑激活值和优化器状态。2. 混合精度训练原理FP16 快但娇气BF16 稳但挑剔混合精度训练不是简单地把模型从 FP32 换成 FP16它包含三个核心设计Master WeightFP32 权重副本、Loss Scaling损失缩放以及 BF16/FP16 的取舍。这套组合拳打下来显存能降一半训练速度还能显著提升。但每一步背后都有数值上的原因不搞清楚就会出现“loss 死活不降”或者“直接爆 NaN”的疑难杂症。2.1 Master Weights为什么更新时必须有一份 FP32 副本这大概是混合精度训练里最反直觉的一点为了省显存用了 FP16反而要额外存一份 FP32 的权重副本那不是更费显存吗道理是这样的FP16 的精度大约只有 3 位有效十进制数字能表示的最小正常数范围也远不如 FP32。当用 Adam 这类自适应学习率优化器做更新时权重更新的量级可能很小比如 1e-6而 FP16 在数量级接近 1 的数上最小可表示的增量大约在 1e-3 左右。这种情况下更新步长会被舍入误差直接吞掉等于没更新。所以实际做法是优化器更新在 FP32 的 Master Weight 上进行每一步用 FP32 权重更新完再向下转换成 FP16 给前向与反向传播使用。这个设计是“省了运算带宽多了存储副本”。显存上确实多了一份 FP32 权重但省下了 FP32 的梯度存储和优化器状态里的算力开销总体来说收益远大于代价。而且如果用了 ZeRO-1 或 ZeRO-2FP32 权重副本还能被进一步切分。每个新手跑混合精度时都觉得“我换成 FP16 训就够了”实际上正确的环路是FP32 权重Master Weight复制成 FP16 权重。FP16 权重做前向传播。损失函数计算后在 FP32 下得到 loss再缩放到 FP16 范围内。FP16 反向传播得到 FP16 梯度。FP16 梯度转回 FP32参与优化器更新。更新后的 FP32 权重再次转成 FP16进入下一轮。这套结构下真正在“算”的部分基本都是半精度而“存”的部分保留了一份高精度保底。没有 Master Weight混合精度训练在长序列和大模型上几乎必然出现收敛退化。2.2 Loss Scaling对抗梯度下溢的保命机制FP16 表示数有一个特点范围比较窄超出 65504 就会变成 Inf比 6e-5 还小的数则会直接变零。反向传播里梯度出现大量小数值是常态如果误差直接变成 0底层参数根本学不动。这就是为什么 PyTorch 的 GradScaler 默认会做损失缩放先把 loss 乘上一个较大的缩放因子常见是 65536 或 2 的整数次幂再对缩放后的 loss 做反向传播等效于把梯度整体放大若干倍避免小梯度在 FP16 下被下溢清空。我曾经调一个多任务模型的 loss发现某个任务分支的梯度一直特别小开混合精度后那个分支彻底“躺平”。后来把 GradScaler 的 init_scale 从默认值调大几个数量级然后观察梯度的统计分布分支的更新量才恢复正常。这件事给一个经验如果发现混合精度训练下 loss 下降曲线比 FP32 平缓很多第一件事不是调学习率而是检查 loss 的 scale 因子有没有在工作以及梯度是不是大面积在这个量级下处于下溢区域。PyTorch 里 GradScaler 的机制还会动态调整缩放因子连续若干步没有出现 Inf/NaN就把 scale 调大一旦出现溢出就跳过这一步然后调小 scale。这套自动机制比较省心但大模型训练里有时梯度本身就容易爆scale 会来回震荡。这时候我会选择把 max_scale 调高同时配合梯度裁剪grad clip让 loss 缩放和梯度裁剪配合而不是彼此打架。2.3 BF16 与 FP16精度范围与硬件支持是分水岭BF16Brain Floating Point用 8 位指数和 7 位尾数表示数指数范围和 FP32 一样宽但尾数精度非常低。FP16 是 5 位指数和 10 位尾数数值更精细但范围窄容易上溢。对训练来说BF16 最大的优势在于不用担心下溢——因为它的指数范围和 FP32 一致代价是尾数精度低常规距离度量时误差明显。这个差异在模型比较敏感时会造成一个麻烦BF16 的舍入误差在 loss 收敛后期可能表现为 loss 震荡踩点不准。但现代大语言模型动辄几千亿参数训练时对极小尾数的敏感度远低于对量级范围的需求加上分布式训练中通信量巨大BF16 的宽范围更适合大规模场景。PyTorch 从 1.12 起把 AMPAutomatic Mixed Precision的 BF16 混精策略支持得很完整用torch.autocast(dtypetorch.bfloat16)就能开启。还有一个实际参考点NVIDIA 从 Ampere 架构开始支持 BF16 计算但效率不完全一致而华为昇腾、寒武纪等国产加速卡对 BF16 的支持也在往主推方向发展。选型策略很简单如果是 Ampere 及之后的 NVIDIA 卡BF16 通常比 FP16 更省心如果是老卡V100 等老老实实用 FP16 加 GradScaler。我在 A800 上实测过相同规模的 7B 微调BF16 和 FP16 的收敛差异在可接受范围内但 BF16 下基本没有出现因 scale 调整导致的“垃圾步”省了不少 debug 时间。以下是对照表方便查阅特性FP16BF16指数位58尾数位107表示范围窄最大 65504与 FP32 相同下溢风险高需 Loss Scaling低基本无下溢精度相对较高相对较低硬件要求大部分卡可用Ampere 及之后 / 新一代加速卡典型策略PyTorch GradScalerautocast 开 bfloat16如果你的训练对数值特别敏感比如做强化学习的人类反馈对齐或者跑小模型做科学计算类任务混合精度不要盲目上 BF16最好先用小规模数据跑几个 step对比 loss 曲线与 FP32 的偏差。3. 显存估算实操从公式推导到逐层核对现在进入本文最硬核的部分给定一个模型参数量和训练配置如何不靠猜算出大概需要多少显存。下面这套方法我一直在用它可以帮你快速判断当前硬件够不够以及在多卡下怎么切分最合理。3.1 粗算公式按参数与优化器状态直接估对于 Transformer 类模型训练态下最基本的粗算公式可以写成训练显存 ≈ 模型权重字节数 梯度字节数 优化器状态字节数 激活值峰值逐项展开权重参数量 × 精度字节数FP16 为 2FP32 为 4混合精度下按 FP16 算更接近实际但别忘了 Master Weight 的存在梯度参数量 × 4通常是 FP32优化器状态AdamW 时为参数量 × 4 × 2SGD 类为参数量 × 4激活值取决于 batch size、序列长度、隐藏维度、层数、注意力头数以及是否开梯度检查点拿 7B 模型、AdamW、FP16FP32 混合精度来粗算FP16 权重14GBFP32 Master Weight28GBFP32 梯度28GBAdam 状态56GB以上合计 126GB还没算激活值。看到这个数字单张 80GB 的 A100/H100 其实是放不下的——除非上 ZeRO比如把优化器状态切到多卡。这也是为什么 7B 以上的训练基本告别单卡甜点区至少需要多卡 ZeRO-1 才能舒服地跑。粗算公式最大的价值是帮你在选硬件阶段排除方案。如果粗算结果已经超过可用显存 20% 以上别折腾单卡了直接考虑多卡并行或者用 LoRA 等参数高效微调。3.2 细算激活值一条可逐层推导的链激活值细算是显存估计里最需要耐心的一步因为不同层不同位置保存的张量形状都不一样。以标准自注意力 Transformer 编码器层为例逐层拆输入到 LayerNorm需要保存输入 x形状是N × L × HN 是 batchL 是序列长度H 是隐藏维度。多头注意力中在计算出 Q/K/V 之后如果实现里保存了注意力分数矩阵那还需要保存N × 头数 × L × L的矩阵。MLP 中第一个线性层产生的中间激活形状一般是N × L × 4H标准的 FFN 是 4 倍隐藏维度。Dropout 的掩码有时也会被存下来。每一层都有属于自己的残差连接与 LayerNorm 参数对应的激活都要计入。把每一项加总再乘上模型层数 L_layer就是单样本激活的总项。然后乘以 batch size或者按实际采样逻辑算总 batch就能得到激活值总体积。拿一个小配置演示一下N4、L1024、H1024、12 层、8 头、FFN 维度 4096。每层激活中光 MLP 中间激活就是 4 × 1024 × 4096 × 4 字节 / 1024^3 ≈ 64MB光这一项 12 层就 768MB再加上 QKV、注意力分数矩阵、dropout 掩码整体下来轻松上 2GB 以上。如果你把 batch size 提升到 16激活值直接乘以 4立刻超过 8GB——这就是为什么很多人把 batch 调大一倍就直接 OOM。开梯度检查点之后前向过程不保留这些大中间激活反向传播时重新前向计算一次激活显存能降到原来的 1/3 甚至更低代价是训练时间大约增加 30% 到 50%。所以实际调参的时候我会先问自己是时间更重要还是显存更重要如果是多机多卡研究实验开梯度检查点往往很划算。3.3 实测核对用 nvidia-smi 和 PyTorch 逐段对账公式算出来只是一张理论报表真实的显存占用还会因为框架缓存、CUDA context、通信库而多出几个 GB。我的习惯是先用公式估算一个基准值然后实际跑两三个 step 用nvidia-smi和 PyTorch 的torch.cuda.max_memory_allocated()做一个对账。import torch # 在训练循环开始时记录内存在结束时取峰值 used torch.cuda.max_memory_allocated() print(fPeak allocated memory: {used / 1024**3:.2f} GB)注意nvidia-smi显示的进程占用通常会比torch.cuda.max_memory_allocated()高因为前者包含 CUDA context、cuDNN 工作空间、NCCL 持久通信缓冲等。差异一般在这个范围CUDA context 在 PyTorch 初始化后约 500MB 到 1GBNCCL 每个通信组会预留固定的缓冲多卡下每卡可能多出几百 MB 到 1GB。如果实测比公式高很多优先看是不是本机的 NCCL 环境变量比如NCCL_BUFFSIZE设置的缓冲过大或框架开了显存预分配。这里给一个建议做这类核对时把 batch size 从 1 开始慢慢加画出“显存随 batch 变化的曲线”。如果曲线斜率异常陡峭基本是激活值在作怪如果起点就很高、斜率很平则是权重、梯度和优化器状态占比过高。排查方向完全不同效率差很多。3.4 ZeRO 与 SDP 并行对显存的影响实际分配方式决定最终需求多卡场景下分布式并行对显存的影响不是简单的除法。理解这一点对准确估计非常有帮助。最常见的组合是 PyTorch FSDP 或 DeepSpeed ZeRO它们做的事情本质相同把优化器状态、梯度、甚至权重在各 rank 间切分切得越狠单卡峰值越低但通信量也会上升。ZeRO-1只切分优化器状态。7B 模型在 8 卡下Adam 状态从 56GB 摊到每卡 7GB单卡显存压力大幅下降。ZeRO-2切分优化器状态 梯度。ZeRO-3全部切分包括权重。单卡可以只保留一层权重适合模型大到根本装不进单卡的情况但通信开销显著增加。我之前用 DeepSpeed ZeRO-3 在 4 张 32GB 卡上微调过一个 13B 模型单卡显存峰值大约 26GB跑得非常勉强但能跑通。换成 8 卡后并没有想象中每卡降到 13GB因为每卡还要处理通信缓冲和临时显存。这里有一个经验ZeRO 对显存的帮助存在边际递减效应卡越多单卡显存减少的比例越不明显通信墙反而会成为新瓶颈。Tensor Parallel张量并行则会把权重、激活按切分维度拆到各卡上单卡存一部分权重计算和通信模式完全不同。所以显存估算要结合具体并行策略来算不能只套公式。4. 训练微调的调优实践混合精度之外还能从哪里挤出显存如果说前两章是“算账”那这一章就是“省钱”。混合精度帮你削掉优化器状态和梯度的部分压力但真正要把一个 7B 或 13B 塞进有限的硬件还需要组合拳。4.1 梯度累积等效大 batch 的显存杠杆很多人一上来就追求大 batch size结果显存爆了被迫调小又发现梯度噪声大、收敛不稳。梯度累积是在这种两难局面下的标准解法每个 micro batch 单独前反向传播不立即做优化器更新而是把梯度累加到一个缓冲里攒够多个 micro batch 之后再统一更新权重。这样你享受了大 batch size 的稳定更新又不需要一次把所有样本都塞进 GPU。显存开销主要看单个 micro batch 的激活值而梯度缓冲只是额外一份与参数等大的张量。在 7B 模型上用 batch size 4 累计步数 8等效 batch 是 32训练效果与显存压力和解 32 完全不可同日而语。代码上的一个细节是PyTorch 默认在调用loss.backward()时会自动把梯度累加进.grad所以梯度累积其实不需要额外维护缓冲区只需控制 optimizer.step() 的调用频率。唯一要注意的是如果开了 GradScaler不能每个 micro batch 都调用scaler.scale(loss).backward()后马上scaler.step(optimizer)必须等完整的累积步数满后再统一step否则缩放系数会被反复“消耗”逻辑容易乱。推荐写法是每个 micro batch 结束时scaler.scale(loss).backward()等到累积步数满后再scaler.step(optimizer)和scaler.update()。4.2 梯度检查点与激活重算用 30% 算力换 70% 显存梯度检查点在上一节已经提到这里展开讲实际操作。它并不是简单地“不保存所有激活”而是小心地选择哪些层保留激活、哪些层丢弃。常见的策略是每 N 层设置一个检查点反向传播时若需要第 m 层的激活但该激活没被保存就从最近的检查点重新执行前向到第 m 层重新算出激活值。PyTorch 里使用非常方便from torch.utils.checkpoint import checkpoint def transformer_block_with_checkpoint(block, x, *args): return checkpoint(block, x, *args)实测下来在 13B 模型、序列长度 4096 的场景下开启梯度检查点后显存峰值从原来的 31GB 降到 11GB 左右时间成本大约增加 40%。如果你的节点网卡很快、kernel 实现不错这个时间成本在实际工程中完全可以接受尤其当它能帮你把模型从“跑不了”变成“能跑”时这一点时间太值了。4.3 动态显存分配器与 offloadL2 优化或 CPU 卸载的取舍除了上面两个大招式还有一些小优化点日常运维中非常实用。PyTorch 的显存分配器默认是缓存式的它在分配和释放时会保留一部分显存复用表现为显存高水位不会立刻下降。如果训练后加载又卸载不同大小的模型显存会“碎”得厉害这时候可以用PYTORCH_CUDA_ALLOC_CONFexpandable_segments:True开启扩展段分配实测对碎片化场景能降低几十个百分点。CPU offload 则是另一个方向的极端做法把优化器状态甚至梯度卸载到 CPU 内存用 PCIe 带宽换显存。DeepSpeed ZeRO-Offload 就是这么干的。它对单卡小显存场景特别有效但 CPU 与 GPU 之间的搬运延迟会导致训练吞吐下降。如果你有 64GB 内存且不着急出结果这是个不错的兜底方案如果你追求吞吐和迭代速度别开 offload宁可减 batch 或者上多卡。4.4 常见 OOM 排查路径不是所有爆显存都该去加卡排错思路和调优思路同样重要。一次训练 OOM启动nvidia-smi看到显存占满不少人第一反应是加卡实际上可能是进入了一个可以快速修复的坑。我建议按下面的顺序排查看报错栈如果是CUDA out of memory先定位是发生在前向、反向还是 optimizer step这决定了是激活值、梯度还是优化器状态的锅。用torch.cuda.memory_summary()看分配明细确认是不是分配器缓存导致的“虚高”。从头看模型是否用了非必要的超大临时张量比如某些框架实现里为了方便而把整个序列的中间结果全部 ReLU 前缓存。顺手检查是否开了eval模式却在构建计算图。我记得有一次发现某些 PyTorch 版本在 LoRA 微调时没冻结 base model结果反向传播把整个大模型梯度都算了显存怎么可能不爆。把这条路径走一遍绝大多数 OOM 都不是硬件不够的问题而是代码或配置问题。反过来如果确认是激活值确实承担不住 batch size优先开梯度检查点和梯度累积而不是急着砍模型尺寸。5. 从理论到部署几个值得长期保留的显存优化习惯这一章更像是长期工程实践里的“洗手习惯”没什么高深理论但坚持下来能避免很多半夜才出现的玄学 bug。5.1 给每次实验留下一张显存快照训练脚本里我会固定输出一张统计表包含权重大小、梯度大小、优化器状态大小、激活峰值、实际峰值、batch size、是否开启梯度检查点。每次调参后自动更新。这样一周后回看能很清楚地知道显存是被哪个变量吃走而不是靠记忆去猜。nvidia-smi --query-gpumemory.used,memory.total --formatcsv只要坚持记录三轮实验你对自己模型的显存画像会变得非常清晰。以后看到别人的配置就能直接判断“他这配置八成得 OOM”省去很多试错。5.2 版本升级时留意默认行为的变化PyTorch、transformers 和 DeepSpeed 的版本升级往往会改变显存行为。举个例子transformers 的某些版本默认开启了use_cacheTrue做生成时会把键值缓存保留下来显存占用和序列长度线性相关训练时如果忘了关长序列场景下会有意外开销。而 DeepSpeed 版本升级后 ZeRO 的通信配置有时会改变默认 offload 策略导致显存表现前后差异巨大。所以每次升级依赖都要拿一两个基准跑一遍别迷信“版本升级一定更好”。这个教训我吃过亏有一次仅仅从 transformers 4.28 升到 4.31同样的脚本显存多了 2GB最后定位到是默认的注意力实现从 eager 变成了 SDPA行为差异导致的。5.3 混合精度的最终选择优先在训练开始前确定训练到一半切换精度策略是噩梦级的决定。一旦开启混合精度优化器状态和权重副本的组织方式已经定下来了中途从 FP16 切 BF16所有状态都得重新初始化而且此前适配 GradScaler 的 scale 记录全部作废。所以确定精度策略要在训练开始前就敲死包括是否开 Loss Scaling、是否开梯度检查点、batch size 和序列长度组合能不能在目标显存内稳定运行。如果你拿不准就先用小规模数据跑 5 个 step对比 FP32、FP16、BF16 三套配置的 loss 下降情况和显存峰值。这个对比实验成本很低但能把后期的问题消灭在萌芽阶段。很多人在这一步偷懒结果训练跑了两天开始出 NaN回滚时间成本非常不划算。6. 写在最后一个可复用的最小显存配置清单文章写到这里最后给大家一份我自己的显存配置工作流。每次接手一个新模型训练或微调任务我基本照着这个清单过一遍查参数量确认权重、梯度、优化器状态三项的理论值。用目标 batch size 和序列长度粗算激活值判断是否需要梯度检查点。决定精度策略新卡用 BF16老卡用 FP16 GradScaler。把 batch size 从 1 开始往上试探记录显存曲线找出线性区和拐点。确定是否上 ZeRO 或 FSDP按卡数、通信拓扑和显存预算做最优选择。跑 5 步基准核对显存实际值与理论值的偏差确认 buffer 大小是否合理。长期实验固定记录显存快照。这套流程本身就是我一路踩坑换来的。早期做分布式训练我老是纠结怎么让理论值与nvidia-smi完全一致后来才明白对账的意义不是追求精确相等而是定位异常和标准差。显存估计和混合精度训练都是工程性极强的事情没有太多玄学把账算清把机制摸透剩下的就是按部就班地跑实验、看曲线、调配置。如果你正在为一张卡能不能跑某个模型纠结建议先按上面的账本老老实实算一遍再决定是买卡、租卡还是老老实实上 LoRA。