ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

大模型训练显存估算与混合精度训练:FP16与BF16选型指南

2026/10/8 12:03:09 拓冰建站 浏览量
大模型训练显存估算与混合精度训练:FP16与BF16选型指南 1. 大模型训练显存估计与混合精度训练详解显存不够这件事几乎每个做大模型训练的人都经历过。你兴冲冲配好了环境写好训练脚本一跑起来直接给你甩一个CUDA out of memory连个商量余地都没有。更让人头疼的是有时候你把 batch size 调到 1 了它还是炸。这时候你才意识到大模型训练里显存不是“省着点用”就能解决的问题它是一道需要提前算清楚的数学题。这篇内容就是围绕这道数学题展开的。我会把大模型训练中显存到底被谁吃掉了、每一项占多少、怎么在动手之前就估算出来讲清楚然后重点拆解混合精度训练——FP16 和 BF16 到底怎么选、为什么现在主流都推 BF16、INT8 和 BF16 的模型又差在哪里。适合正在做或者准备做大模型训练、微调的工程师也适合想搞明白“为什么我的卡跑不动”的开发者。看完你至少能做到两件事一是拿到一个模型能大致算出它要多少显存二是知道混合精度训练该怎么配、坑在哪里。2. 显存到底被谁吃掉了逐项拆解2.1 显存占用的四大块很多人以为显存就是用来存模型参数的这个理解只对了四分之一。在大模型训练场景下显存主要被四部分占用模型参数、梯度、优化器状态、激活值。前三个是“静态”的跟模型结构强相关一旦模型定了基本就定了激活值是“动态”的跟 batch size、序列长度直接挂钩这也是为什么你把 batch size 调到 1 还可能炸的原因。先给一个最直观的结论用 Adam 优化器做全参数训练时一个参数量为 P 的模型光是参数、梯度、优化器状态这三项就需要大约16P 字节的显存混合精度下。这个 16 是怎么来的下面一步步算。2.2 参数、梯度、优化器状态的显存账先看纯 FP32 训练的情况。假设模型有 P 个参数模型参数本身每个参数 4 字节占4P梯度每个参数对应一个梯度也是 4 字节占4PAdam 优化器状态Adam 会为每个参数维护一阶动量momentum和二阶动量variance各 4 字节占8P加起来就是 4P 4P 8P 16P 字节。也就是说一个 10B100 亿参数的模型纯 FP32 训练光这三项就要 160GB 显存一张 80GB 的卡根本放不下。那混合精度训练为什么能省因为参数和梯度可以用 FP16/BF16 存每个只占 2 字节。但注意Adam 的优化器状态通常还是 FP32 的而且为了数值稳定一般还会保留一份 FP32 的 master weight主权重。所以混合精度下FP16/BF16 参数2PFP16/BF16 梯度2PFP32 master weight4PAdam 一阶动量FP324PAdam 二阶动量FP324P合计 2P 2P 4P 4P 4P 16P 字节。你会发现混合精度下这三项还是 16P没省这是因为省下来的参数和梯度空间被 master weight 和 FP32 优化器状态吃回去了。注意很多人以为混合精度训练能大幅降低静态显存其实它主要省的是激活值显存和计算开销静态部分并没有想象中省那么多。真正想省静态显存得靠 ZeRO、LoRA 这类技术。2.3 激活值真正的大头激活值是前向传播过程中每一层产生的中间结果反向传播时要用到所以必须存着。它的显存占用跟batch size × 序列长度 × 隐藏维度 × 层数成正比。公式化一点对于 Transformer 结构激活值大致跟batch_size × seq_len × hidden_size × num_layers成正比前面还有个跟具体实现相关的系数。这就是为什么序列长度翻倍显存经常直接翻好几倍甚至更多。你跑 512 长度没事跑到 2048 就炸不是模型变大了是激活值涨了。激活值可以通过梯度检查点gradient checkpointing来换显存——不存中间激活反向时重新算一遍用时间换空间通常能把激活值显存降到原来的几分之一。2.4 一个可落地的显存估算公式把上面几块合起来给你一个实操中够用的估算公式混合精度 Adam 梯度检查点总显存 ≈ 16P参数梯度优化器 激活值 临时缓冲 激活值 ≈ batch_size × seq_len × hidden_size × num_layers × k其中 k 是个经验系数跟是否用梯度检查点、注意力实现方式有关。不开梯度检查点k 可能到 10 以上开了之后能降到 2 左右。临时缓冲一般留个 1-2GB 余量。举个例子一个 7B 模型P 7×10^9静态部分16 × 7×10^9 112GB看到没光静态就 112GB 了单卡 80GB 根本放不下。所以 7B 模型全参数训练至少得上多卡 ZeRO或者用 LoRA 这类只训练部分参数的方法。这也是为什么现在个人和小团队做微调LoRA 几乎是标配。3. 混合精度训练FP16 与 BF16 的取舍3.1 混合精度到底“混合”了什么混合精度训练的核心思路很简单计算用低精度存储和更新用高精度。前向和反向的矩阵运算用 FP16 或 BF16 来做速度快、显存省但参数的更新、优化器状态、以及一些对数值敏感的累加操作用 FP32 来保证精度。具体来说训练时会同时维护两份权重一份 FP16/BF16 的用于前向反向计算一份 FP32 的 master weight用于接收梯度更新。每次迭代FP16 的梯度算出来后转成 FP32 更新 master weight再把 master weight 转回 FP16 给下一轮用。这样既享受了低精度计算的速度又避免了低精度累积误差把模型训崩。3.2 FP16 的数值范围陷阱FP16 有 1 位符号、5 位指数、10 位尾数。它的动态范围大概是 6×10^-5 到 65504。问题就出在这个范围上太小会下溢成 0太大会溢出成 inf。训练中最典型的问题就是梯度下溢。很多梯度值非常小FP16 表示不了直接变成 0参数就更新不动了。另一个问题是梯度爆炸一旦某个梯度超过 65504直接变 inf整个训练就废了。为了解决这个问题FP16 训练通常要配一个loss scaling损失缩放机制。思路是把 loss 乘上一个很大的数比如 2^16这样反向传播出来的梯度也同比放大就不容易下溢了更新参数前再除回来。动态 loss scaling 会自动调整这个缩放系数发现溢出就减小一段时间没溢出就增大。3.3 BF16 为什么成了主流BF16 有 1 位符号、8 位指数、7 位尾数。它的指数位和 FP32 一样多所以动态范围和 FP32 基本一致能表示 10^-38 到 10^38 这个量级。代价是尾数只有 7 位精度比 FP16 低。但对深度学习来说动态范围比精度更重要。梯度下溢、溢出这些问题BF16 基本不会遇到所以BF16 通常不需要 loss scaling训练更稳定代码也更简单。这就是为什么现在新出的卡A100、H100、以及各类国产加速卡都在推 BF16主流框架的默认混合精度也逐步从 FP16 转向 BF16。代价是 BF16 的精度损失在理论上比 FP16 大但在实际训练中由于有 FP32 master weight 兜底这个损失通常可以接受。实测下来BF16 训练的收敛曲线和 FP32 非常接近而 FP16 如果 loss scaling 没调好反而容易出问题。3.4 FP16 与 BF16 对比速查维度FP16BF16符号位11指数位58尾数位107动态范围约 6e-5 ~ 65504约 1e-38 ~ 3e38精度较高较低是否需要 loss scaling通常需要通常不需要训练稳定性一般好硬件支持广泛较新硬件实操建议如果你的卡支持 BF16优先用 BF16省心。如果只能用 FP16务必开启动态 loss scaling并且监控梯度是否溢出。4. INT8 和 BF16 模型的区别4.1 它们解决的不是同一个问题热搜里“int8和bf16模型的区别”这个问题很多人搞混了一个点BF16 主要是训练和推理时的计算精度格式而 INT8 主要是推理阶段的量化格式。它们的目标不一样。BF16 模型指的是用 BF16 精度训练或存储的模型它的权重和激活还是浮点数只是精度比 FP32 低。INT8 模型指的是把权重有时还有激活量化成 8 位整数来存储和计算是一种量化手段。4.2 显存和精度的权衡从显存角度看INT8 每个参数只占 1 字节BF16 占 2 字节FP32 占 4 字节。所以一个 7B 模型FP32约 28GBBF16约 14GBINT8约 7GBINT8 在显存上优势明显这也是为什么端侧部署、大模型推理都喜欢用 INT8 量化。但代价是精度损失INT8 量化会引入量化误差需要通过校准calibration来尽量减少。BF16 的精度损失则小得多基本可以忽略。4.3 使用场景怎么选训练阶段用 BF16或 FP16混合精度不用 INT8。INT8 训练目前还不成熟容易训崩。推理阶段追求精度用 BF16几乎无损。推理阶段追求省显存/省成本用 INT8甚至 INT4但要接受一定精度下降。端侧/边缘设备优先 INT8 及更低比特量化。一句话总结BF16 是“降精度但保范围”INT8 是“降精度也降范围但省得更多”。两者不是替代关系而是不同阶段的不同工具。5. 实操从估算到跑起来的完整流程5.1 第一步动手前先算显存拿到一个模型先别急着跑。按这个顺序估一遍确认参数量 P看模型配置文件或论文算静态显存全参数训练按 16P 估LoRA 微调按 (2P 少量) 估估激活值batch_size × seq_len × hidden_size × num_layers × kk 取 2开梯度检查点到 10不开加 1-2GB 缓冲跟你的卡对比决定要不要上多卡、ZeRO、LoRA5.2 第二步配置混合精度以主流框架为例混合精度一般通过一个开关或上下文管理器开启。核心配置项# 伪代码示意具体 API 以你用的框架为准 training_config { precision: bf16, # 优先 bf16不支持则 fp16 loss_scaling: dynamic, # fp16 时开启bf16 可关 gradient_checkpointing: True, # 激活值吃紧时开启 gradient_accumulation_steps: 4, # 显存不够时用累积换 batch }关键点gradient_accumulation_steps是个好东西。显存不够时把大 batch 拆成几个小 batch 分步算梯度再累积等效于大 batch但显存占用按小 batch 算。5.3 第三步跑起来后的监控训练跑起来后盯这几个指标显存占用用nvidia-smi或框架自带的内存监控看峰值有没有逼近上限loss 曲线BF16 应该和 FP32 很接近FP16 如果 loss 突然变 NaN多半是 loss scaling 出问题梯度范数如果梯度范数异常大或异常小检查精度配置5.4 显存不够时的排查顺序按这个顺序试基本能解决 90% 的 OOM降 batch size开梯度检查点用梯度累积换 batch换 LoRA 等参数高效微调上 ZeRO / 模型并行换更省显存的注意力实现6. 常见问题与避坑经验6.1 常见问题速查表问题现象可能原因解决方向一跑就 OOM静态显存超了上 LoRA / ZeRO / 多卡batch1 还 OOM激活值太大开梯度检查点、降序列长度loss 变 NaNFP16 梯度溢出换 BF16 或调 loss scaling训练极慢精度配置不当确认用了混合精度显存忽高忽低碎片化设置内存分配策略、固定输入形状多卡显存不均并行策略问题检查 ZeRO stage 和切分方式6.2 几个踩过的坑坑一以为混合精度能省一半显存。前面算过了静态部分混合精度并不省省的是激活值和计算。别指望开了混合精度就能把 7B 模型塞进单卡。坑二FP16 不开 loss scaling。这个错误新手常犯结果就是训练半天 loss 不动因为梯度全下溢成 0 了。用 FP16 一定记得开。坑三BF16 无脑用。虽然 BF16 稳定但有些老卡不支持强行用会报错或者走模拟路径反而更慢。用之前确认硬件支持。坑四忽略激活值。很多人只算参数显存结果一跑就炸。记住激活值跟序列长度强相关长文本任务尤其要注意。坑五梯度累积和 batch norm 冲突。如果用梯度累积注意 batch norm 的统计量会不准大模型一般用 layer norm 所以问题不大但心里要有数。6.3 一个实用的小技巧如果你不确定该用 FP16 还是 BF16写个判断逻辑先检测硬件是否支持 BF16支持就用 BF16不支持再退回 FP16 并开启动态 loss scaling。这样一套代码能适配不同环境省得来回改配置。# 伪代码示意 if hardware_supports_bf16(): precision bf16 use_loss_scaling False else: precision fp16 use_loss_scaling True这个判断在实际部署里特别有用因为你的训练环境和推理环境可能不是同一批卡。7. 显存优化的进阶方向7.1 ZeRO把静态显存摊到多卡上ZeROZero Redundancy Optimizer的核心思想是把参数、梯度、优化器状态这些本来每张卡都存一份的东西切分到多张卡上。ZeRO 分三个阶段stage 1 切优化器状态stage 2 再切梯度stage 3 连参数也切。切得越狠单卡显存越省但通信开销越大。实操中7B 模型全参数训练ZeRO stage 2 或 3 基本是标配。stage 3 省显存最明显但对通信要求高卡间带宽不够会拖慢训练。7.2 LoRA只训练一小部分参数LoRA 的思路是在原模型旁边挂小的低秩矩阵只训练这些小矩阵原模型参数冻结。这样需要存梯度和优化器状态的参数量大幅减少显存占用能降一个数量级。对于微调场景LoRA 几乎是性价比最高的选择。7.3 梯度检查点用时间换空间梯度检查点不存中间激活反向传播时重新算。显存能省很多代价是训练速度慢 20%-30%。显存吃紧时这是最直接的救命手段。7.4 这些技术怎么组合实际项目里这些技术经常组合使用。比如LoRA BF16 梯度检查点能在单卡上微调 7B 甚至 13B 模型。再大就得加 ZeRO 和多卡。组合的原则是先上最省事的混合精度、梯度检查点不够再上 LoRA还不够再上 ZeRO 和多卡。我个人在实际操作中的体会是显存估算这件事宁可提前多算十分钟也别跑起来炸了再回头查。尤其是团队协作时把显存预算写进训练配置文档能省掉大量沟通成本。混合精度这块BF16 能上就上别在 FP16 的 loss scaling 上浪费太多时间。最后再分享一个小技巧训练脚本里加一行显存峰值打印每次跑完记录一下几次之后你对这类模型的显存直觉就建立起来了比任何公式都准。