大模型训练显存估算与混合精度实践:从OOM到从容开训
准备训练一个 7B 模型前我犯过一个让同事笑掉大牙的错误看着模型参数只有 14GBBF16就信心满满地在一张 40GB 的卡上起了训练脚本结果第一个 step 都没跑完就 OOM。后来我才真正搞明白大模型训练的显存占用从来不是模型多大就占多大参数只是显存账单上的第一行后面跟着梯度、优化器状态、激活值、临时缓冲区每一行都不容小觑。这篇就把我踩坑换来的经验和计算套路完整讲清楚重点是怎么在开训练前准确估算显存以及混合精度训练为什么能省显存、又为什么不是无脑换成低精度那么简单。适合准备训练或微调大模型、经常被 OOM 卡到怀疑人生的同学参考。1. 参数不是唯一账单训练时显存的五笔大额支出1.1 模型参数与梯度最直观但未必最大的部分很多人估算显存的第一反应是看模型参数量。这个习惯没大错但如果只算这一项后面一定会翻车。以 7B 模型为例模型权重本身占有的显存FP32 精度下是 7B × 4 字节 ≈ 28GBBF16/FP16 下是 14GBFP8 下甚至只需要 7GB。这个数字是静态的模型加载后一直躺在显存里看得见摸得着。训练和推理最本质的区别是训练还要保存梯度。反向传播会把每个参数的梯度算出来存起来等 optimizer 更新时统一使用。梯度的大小和参数本身一样也是每个参数一份。如果梯度也用 FP32 存7B 模型又加 28GBBF16 则是 14GB。到这里7B 模型在 BF16 训练下光参数加梯度就是 14 14 28GB看起来还在 40GB 卡的能力范围内。但千万别高兴太早接下来的优化器状态才是真正的大头。1.2 优化器状态隐藏的显存大头以最常用的大模型优化器 AdamW 为例它会给每个参数额外维护两个状态一阶矩momentum和二阶矩variance。这两个状态几乎总是以 FP32 保存因为低精度会让累积统计的信息逐轮漂移。每个参数多出 4 4 8 字节。7B 模型就是 56GB这个数字比模型参数本身还大一倍。于是单纯参数 梯度 Adam 优化器状态在混合精度场景下就已经是 14 14 56 84GB。大部分单卡 A100 80GB 已经放不下了更别提还得给激活值留位置。网上有个流传很广的说法混合精度训练下每个参数大约占 16 字节。这个数字怎么来的如果是显式保留 BF16/FP16 参数副本2B加上对应精度梯度2B再加上 FP32 Master Weight4B和 Adam 状态8B正好是 16B/参数。如果你走 PyTorch AMP 路线模型参数保持 FP32内部在计算时临时转成低精度做算子运算存储构成是 FP32 参数 4B FP32 梯度 4B Adam 状态 8B总账也是 16B/参数。殊途同归7B 模型的固定训练开销就是约 112GB这不是单卡能解决的问题。1.3 动态激活值与临时缓冲区规模取决于 batch 与序列长度固定开销之后真正让显存估算变得复杂的是激活值。前向传播时每一层算出来的中间张量都要保存下来供反向传播计算梯度使用。模型参数量固定后激活值的大小主要由 batch size、序列长度、层数、隐藏维度和算子实现方式决定。曾经有人在 40GB 卡上强训 7B 模型发现即使把 batch size 调到 1依然 OOM 到怀疑人生。原因就在这里序列长度一拉长激活值会呈指数级叠加。比如 7B 模型层数 32、hidden size 4096batch size 1、序列长度 2048 时仅 attention score 矩阵这一项在不使用 FlashAttention 的情况下就要存大约 8.6GB序列长度翻倍到 4096这一项直接飙到 34GB。你没看错只是一个中间张量就比整个模型参数还大。除此之外还有实际训练中不可避免的临时缓冲区AllReduce 集合通信的暂存空间、CUDA context 初始化、PyTorch 显存分配器预留的空闲块。在分布式训练里NCCL 通信用的 buffers 是按最大消息大小预分配的也占用一定显存。这些零零碎碎加起来往往会让 nvidia-smi 显示的数字比公式估算高出一截。2. 显存估算实操7B/13B/70B 模型到底吃多少显存2.1 固定开销一句话公式既然固定开销大头就是参数、梯度、优化器状态这三件套我们可以先建立一个最实用的估算公式纯 FP32 Adam每参数 16B参数 4B 梯度 4B Adam 状态 8B混合精度 Adam每参数约 16B低精度参数 2B 低精度梯度 2B FP32 Master Weight 4B Adam 状态 8B混合精度 SGD如果非要用 SGD每参数约 8B低精度参数 2B 低精度梯度 2B FP32 Master Weight 4B没有额外优化器状态我用这个公式试算过几个常见规模的模型结果和实际训练时观察到的固定占用非常接近模型规模优化器固定显存开销8×A100 80GB 下每卡平摊7BAdam混合精度约 112GB14GB13BAdam混合精度约 208GB26GB70BAdam混合精度约 1120GB140GB显然单节点不够注意固定开销还不含激活值、通信 buffer 和框架自身开销所以看到 7B 模型需要 112GB 这个数字时第一反应应该是单卡无论如何都必须分布式而不是想着硬塞进去。2.2 激活值从倍率法到 Attention Score 专项激活值比较难给一个精确到小数点的公式因为不同架构差异很大。我的习惯是做两层估算。第一层是粗估把激活值近似为 batch_size × sequence_length × num_layers × hidden_size 乘一个常数系数。这个系数在不同模型里变化很大GPT 这类稠密 Transformer 在 FP16 下通常取 8~24 倍具体取决于是否使用 FlashAttention、LayerNorm 实现的精度、有没有用 activation checkpointing。第二层是重点排查 Attention Score 这一项它是最容易爆雷的地方。传统 Transformer 实现里每个 head 的 Q 和 K 相乘后得到的 score 矩阵会被完整保存用于反向传播。这一项的大小是B × num_heads × L × seq_len × seq_len × 2 字节拿 7B 模型举例num_heads 32层数 32batch 1seq_len 2048结果就是约 8.6GBseq_len 4096 时约 34GB。如果模型或框架用了 FlashAttention 这类 flat kernelAttention Score 不会显式物化到显存激活值立刻能砍掉一大截。这也是为什么大模型训练框架近两年全面转向 FlashAttention——不只是为了加速更重要的是为了显存很多人低估了 score 矩阵的显存杀伤力。2.3 分布式训练视角ZeRO 分片后每卡显存怎么算固定开销在一张卡上放不下就得靠分布式把显存摊开。DeepSpeed ZeRO 是最常见的做法分三个阶段理解特别直观ZeRO Stage 1只把优化器状态切到各卡。每卡固定开销 参数全量 梯度全量 优化器状态/N。7B 模型 8 卡下大约是 14 14 56/8 35GB勉强能住进 40GB 的卡。ZeRO Stage 2梯度也分片。每卡固定开销 参数全量 梯度/N 优化器状态/N。7B 模型 8 卡大约是 14 14/8 56/8 22.75GB宽裕很多。ZeRO Stage 3参数、梯度、优化器状态全部分片。每卡固定开销 ≈ 112/8 14GB再留激活值空间40GB 甚至 32GB 卡都有机会跑 7B 级模型的轻量训练。我自己在做 7B 微调时最常用的配置就是 2~4 张 80GB 卡 ZeRO Stage 2/3 activation checkpointing。先算固定开销再给激活值留出 1.5~2 倍 buffer基本不会 OOM。3. 混合精度训练的内在机制不是随意砍精度3.1 FP16/BF16/FP32 的角色分配很多人第一次接触混合精度训练时会误以为把模型换成 FP16 然后跑就行。如果真这么干训练很可能会发散或者loss卡死不动。混合精度的本质不是全员降低精度而是让不同的部分停留在适合它的精度。我从三个精度各自的特性说起FP32标准单精度指数位 8 位尾数位 23 位。动态范围大精度足够但慢显存消耗大。FP16半精度指数位 5 位尾数位 10 位。动态范围窄最大只能到 65504小于约 6e-5 的数值就会迅速损失精度次正规数甚至直接下溢为零。但快Tensor Core 在处理 FP16 矩阵乘法时吞吐量大约是 FP32 的两倍以上。BF16Brain Float指数位 8 位尾数位 7 位。它把 FP32 的动态范围完整保留下来代价是尾数只剩 7 位精度比 FP16 还低但由于动态范围足够大不需要担心梯度下溢到零。同样是 2 字节存储同样能吃到 Tensor Core 加速。在 LLM 训练中权重和梯度通常用 BF16/FP16 参与矩阵乘法optimizer 状态和 master weight 保留 FP32少部分对精度敏感的算子继续 FP32 计算。所以混合精度一词的重点是混合不是低精度。3.2 Master Weights为什么低精度更新必须有个 FP32 存根FP16 只有一个大约三到四位有效十进制数的表示能力。假设权重是 1.0Adam 更新量是 0.0001FP16 根本无法表达这个差值更新会被舍入吞掉等于没训。BF16 更夸张尾数只有 7 位相对精度约 1%参数更新很可能消失。因此混合精度训练必须额外保留一份 FP32 的权重副本Master Weights。优化器在 FP32 副本上完成加法和更新再把更新后的 FP32 值转存回低精度参数用于前向传播。如果你用 PyTorch AMPautocast 的实现策略通常是让模型参数主体还是 FP32只是在算子执行时自动选择合适的低精度数据流这样模型参数本身某种程度上就扮演了 Master Weight 的角色。在 DeepSpeed ZeRO 里FP32 Master Weight 通常会跟着 optimizer 状态一起分片保存这也是为什么 Stage 1 对显存的削减这么明显。3.3 Loss Scaling把梯度从看不见拉回看得见FP16 量级下反向传播的梯度经过层层连乘后会变得特别小很容易低于 FP16 能表示的最小正数直接变零。一旦梯度变零底层参数学习停滞上层靠梯度信号也带不动整个模型训练失效。解决办法是在损失函数上放大一个 scale。前向算出 loss 后乘以一个很大的常数比如 65536 甚至更大再走反向传播。这样梯度也随之放大落到 FP16 的动态范围内能被正常表示和回传。反向传播结束后优化器更新前再把放大后的梯度缩小回去。loss scale 不是一个定死的值PyTorch 的 GradScaler 会动态调整如果检测到梯度出现 Inf/NaN就把 scale 缩小如果连续很多步都正常就适当调大 scale。BF16 因为动态范围和 FP32 相同基本不存在下溢到零的问题所以用 BF16 通常可以省掉 loss scaling 的环节这是一个极其重要的实操差异。3.4 哪些算子在混合精度下要保留 FP32实践中并不是让所有算子都去跑低精度有几类算子对数值精度特别敏感切成 FP16/BF16 很容易出问题LayerNorm / RMSNorm均值和方差的计算涉及小数值累加低精度会引入较大误差。PyTorch 的 autocast 会自动让 LayerNorm 留在 FP32。Softmax指数运算在 FP16 下容易溢出或精度失真同样适合留在 FP32。FlashAttention 内部在计算 softmax 时也会用 FP32 中间值做在线归一化。CrossEntropy Loss / LogSoftmax这类归约操作通常保留 FP32 计算loss 最终返回 FP32。任何涉及小学习率和极小更新量的场景如果你用的是 FP16 且没有 Master Weight梯度下溢和更新舍入会在前几步就暴露问题。这也是为什么在 LLM 类架构里真正被低精度化的主要是线性层矩阵乘法而归一化层、损失层和部分注意力计算仍然保持 FP32 或内部使用 FP32 累加。4. AMP 实战与 BF16 两条路线4.1 标准 AMPautocast GradScalerPyTorch 里最标准的 FP16 混合精度训练长这样import torch scaler torch.cuda.amp.GradScaler() model MyModel().to(cuda) optimizer torch.optim.AdamW(model.parameters(), lr1e-4) for batch in dataloader: optimizer.zero_grad() with torch.autocast(device_typecuda, dtypetorch.float16): loss model(batch) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()关键点有三个autocast 负责把 forward 里的线性层、注意力打分自动切到 FP16GradScaler 负责把 loss 放大后回传scaler.step 前内部会先检查梯度是否 Inf/NaN再决定要不要真正执行 optimizer.step。我建议顺手打开 torch.backends.cuda.matmul.allow_tf32 和相关优化选项但注意它影响的是 FP32 矩阵乘法里是否启用 TF32。这和 AMP 的 FP16 是两回事别混在一起调参。4.2 BF16 路线省事但不等于无脑如果训练环境是 H100、A100 或新的消费级 Ada 架构显卡大多数情况下可以直接用 BF16。很多 Hugging Face Transformer 和预训练框架跑 LLM 时BF16 是默认选择。原因很简单不需要 loss scaling动态范围与 FP32 一致训练稳定性更好。但 BF16 也有自己的坑。尾数只有 7 位极端情况下梯度很小但非零的量也可能超出 7 位尾数的表示能力导致梯度信息丢失。好在 Adam 的动量机制天然对梯度做归一化小梯度经过累积后仍然能产生合理更新所以实践中训练不易崩。真正要注意的是误差累积如果你在 FP32 下有一个严重依赖高精度累积的操作比如某些自定义 CUDA kernel切到 BF16 后可能莫名其妙地 loss 波动。遇到这类情况最稳妥的办法是把该算子显式留在 FP32。4.3 梯度裁剪、checkpoint 保存的几个细节我见过不少同学把 FP16 训练跑通后卡在了一个奇怪的问题上梯度裁剪明明设了 max_norm 1.0为什么裁剪后效果完全不对原因很典型loss scale 还在梯度上没脱掉。梯度裁剪必须在 unscale 之后再进行scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update()如果用 BF16 或者不开 GradScaler直接 clip_grad_norm_ 没有这个问题但一旦将来切回 FP16 就会踩坑。checkpoint 保存也要额外上心。混合精度训练里 optimizer.state_dict 包含 FP32 的动量、方差和 master weight 信息loss scale 也有一个时间状态。保存时记得把 scaler.state_dict 一并存进去torch.save({ model: model.state_dict(), optimizer: optimizer.state_dict(), scaler: scaler.state_dict(), lr_scheduler: scheduler.state_dict(), step: step, }, checkpoint_path)恢复时先 scaler.load_state_dict再继续训练。要是漏掉 scaler 状态loss scale 会从头开始可能在前几步产生极端梯度导致暂时性的稳定下降。5. 组合拳把显存估算变成可落地的训练方案5.1 Activation Checkpointing 与 FlashAttention 的实战收益算完固定开销后下一步是压缩激活值。最直接的手段是 activation checkpointingPyTorch 里对任意模块包一层torch.utils.checkpoint.checkpoint或者对 Hugging Face 模型直接调用model.gradient_checkpointing_enable()。它的原理很反直觉前向传播时不去保存所有中间激活只保存每层的输入等少量必要张量反向传播时再重新跑一次前向把激活现场重建出来。代价是训练时间明显变长大约增加 20%~40% 的计算量收益是激活值内存大幅下降。叠加 FlashAttention 后Attention Score 不再物化两个方案合起来通常能把激活值压到未优化前的 1/4 甚至更低。对于 7B 模型、batch size 1、seq_len 2048 的场景我实测激活值可以从十几 GB 降到 2~3GB 的量级。我在实际使用中的体会是先开 FlashAttention再看是否需要 gradient checkpointing。因为 FlashAttention 对代码侵入性低、基本免费而 gradient checkpointing 会增加计算时间所以只在显存真的紧张时候开。如果你用的是 DeepSpeed 或 PyTorch FSDP框架往往自带 activation checkpoint 的开关直接调用接口即可。5.2 用 PyTorch 自检 API 验证真实占用估算终究是估算理论算完必须拿实测兜底。PyTorch 提供了非常好的内存审计工具核心是这三行torch.cuda.reset_peak_memory_stats() # 跑一个真实 step peak_allocated torch.cuda.max_memory_allocated() peak_reserved torch.cuda.max_memory_reserved()memory_allocated是当前真正被张量占用的显存memory_reserved是 PyTorch 从显卡那拿走的显存包括缓存块。两者之间的差就是分配器预留下来的空闲缓存这解释了为什么 nvidia-smi 看到显存占用比代码里张量需求大得多。我习惯在训练脚本里每隔固定步数打印一下torch.cuda.memory_summary()特别是在开启动态 batch、动态序列长度时能第一时间发现某个模块猛然吃显存。定位到具体算子后再用 torch.profiler 打印算子级别的内存统计确认是不是 Attention Score 或某个 FNN 中间张量在爆。5.3 从估算到选卡我踩坑后的固定流程这几年我摸出来的套路基本是四步先按每参数 16 字节算出固定开销这是下限任何显存优化都不能让这部分低于这个数。根据 batch size 和序列长度估算激活值。如果序列长度超过 1024默认先把 FlashAttention 打开再按激活值是参数固定开销的 10%~20% 做初步猜测。确定是否要开 gradient checkpointing通常一开激活值从十几个 GB 掉到个位数 GB。这种情况下 40GB 单卡 ZeRO Stage 1/2 就能跑 7B 微调。留出 20%~30% 余量给 CUDA context、NCCL buffer、allocator 缓存和临时张量。余量不存在的项目上线必 OOM。我看着很多人的训练计划是在显存上限和batch size 太小导致训练不稳定之间反复横跳。如果你发现 batch size 因为显存被迫压到 1也别纠结梯度累积可以解决有效 batch 大小问题代价只是反向传播次数变多训练效果并不会因此遭受致命打击。最后一个值得记住的细节是不要仅凭一张卡的显存大小判断能不能训练模型。7B 模型 40GB 卡看似可行实际必须配合 ZeRO/FSDP 和激活值优化才能跑起来13B 模型 4×80GB 卡也并不是必须能跑还要看序列长度和实现精度。每次开新训练前走一遍上面的估算流程比反复试错 OOM 快得多。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →