资讯详情

大模型训练显存估计与混合精度训练:BF16、FP16与int8量化详解

📅 2026/10/7 20:10:11 | 华诺云谱 👁 阅读
大模型训练显存估计与混合精度训练:BF16、FP16与int8量化详解
1. 大模型训练显存估计与混合精度训练详解显存不够这件事几乎每个做大模型训练的人都经历过。你写好模型、准备好数据、启动训练脚本结果第一轮迭代还没跑完终端就甩出一行CUDA out of memory然后进程直接挂掉。更让人头疼的是你明明算过参数量觉得一张 80G 的卡应该绰绰有余但实际跑起来就是不够。问题出在哪答案通常藏在显存估计和混合精度训练这两个环节里。这篇文章面向的是正在做或准备做大模型训练的工程师和研究者不管你是在单卡上微调一个 7B 模型还是在多卡集群上预训练一个更大的模型显存估计和精度选择都是绕不开的基本功。我会从显存的组成结构讲起把每一块显存开销的来源拆清楚然后深入混合精度训练的原理和实操细节重点对比 BF16 和 FP16 的差异以及 int8 和 bf16 模型到底有什么区别。整篇内容基于我在实际训练中踩过的坑和总结的经验尽量做到你读完就能拿去算自己的场景。2. 显存到底被谁吃掉了2.1 显存开销的四大组成部分很多人估计显存的时候只算模型参数比如 7B 模型用 FP16 存储那就是 7×10⁹ × 2 字节 ≈ 14GB。然后一看 A100 有 80GB觉得稳了。结果一跑就炸。原因很简单模型参数只是显存开销的一部分而且往往不是最大的那部分。在实际训练中显存主要被以下四块吃掉模型参数Parameters权重本身占用的显存。梯度Gradients反向传播时为每个参数计算的梯度通常和参数同精度、同大小。优化器状态Optimizer States这是最容易被低估的部分。以 Adam 为例它需要为每个参数维护一阶矩估计动量和二阶矩估计方差如果再用 FP32 存储那就是参数量的两倍。激活值Activations前向传播过程中每一层的中间输出需要保留到反向传播时使用。这部分和 batch size、序列长度强相关往往是大头。把这四块加起来才是真正的显存需求。下面我用一个具体的例子来算。2.2 以 7B 模型为例的完整显存计算假设我们有一个 7B70 亿参数的模型使用 Adam 优化器混合精度训练参数和梯度用 FP16优化器状态用 FP32不做任何显存优化。逐项计算模型参数7B × 2 字节FP16 14GB。梯度7B × 2 字节FP16 14GB。优化器状态Adam 需要 FP32 的动量和方差各 7B × 4 字节 28GB两份共 56GB。再加上一份 FP32 的参数副本用于参数更新时的数值稳定又是 28GB。所以优化器相关总共 84GB。激活值这部分变化很大。粗略估算对于 Transformer 结构激活值显存大约和batch_size × seq_len × hidden_dim × num_layers成正比。以 batch size 为 1、序列长度 2048、hidden dim 4096、32 层为例激活值大约在几 GB 到十几 GB 之间取决于是否使用梯度检查点gradient checkpointing。把上面加起来14 14 84 激活值 ≈ 112GB 激活值。一张 80GB 的卡根本放不下。这就是为什么实际训练中必须做显存优化——要么用 ZeRO 之类的分片策略要么用 LoRA 等参数高效微调方法要么把优化器状态换成更省内存的方案。注意上面这个计算是“朴素”情况下的上限。实际中通过混合精度、梯度检查点、ZeRO Stage 1/2/3、8-bit Adam 等手段可以把显存压到原来的几分之一甚至十几分之一。但前提是你得先知道显存花在哪了才能有针对性地优化。2.3 激活值显存的估算方法激活值这块最容易被忽略但它在长序列训练中往往是瓶颈。一个实用的估算公式是激活值显存 ≈ batch_size × seq_len × hidden_dim × num_layers × 精度字节数 × 系数其中系数取决于具体实现通常在 10 到 20 之间因为每层有多个中间张量比如 attention 的 Q/K/V、FFN 的中间层等。以 7B 模型为例hidden_dim4096num_layers32seq_len2048batch_size1FP16 精度1 × 2048 × 4096 × 32 × 2 × 15 ≈ 8GB这只是一个粗略估计。实际中如果开启梯度检查点激活值可以降到原来的 1/√num_layers 左右代价是增加约 30% 的计算时间。这是一个典型的时间换空间的取舍。2.4 快速估算表与经验公式为了方便你快速估算我整理了一个经验对照表以 FP16 混合精度、Adam 优化器为例模型规模参数显存梯度显存优化器状态激活值粗略合计无优化1B2GB2GB12GB1-2GB17-18GB7B14GB14GB84GB8-12GB120-124GB13B26GB26GB156GB15-20GB223-228GB70B140GB140GB840GB80-100GB1200GB从表里可以清楚看到优化器状态才是真正的显存杀手。这也是为什么 ZeRO Stage 2 和 Stage 3 要把优化器状态和梯度分片到多张卡上——不分片根本放不下。一个快速心算的口诀是混合精度 Adam 的情况下每 1B 参数大约需要 16-20GB 显存不含激活值。7B 就是 112-140GB13B 就是 208-260GB。记住这个数量级你在选卡和配集群的时候心里就有底了。3. 混合精度训练的核心原理3.1 为什么需要混合精度默认情况下深度学习框架用 FP32单精度浮点数来存储和计算。FP32 有 23 位尾数、8 位指数数值范围大约在 10⁻³⁸ 到 10³⁸ 之间精度足够但代价是每个数占 4 字节而且计算吞吐量相对较低。混合精度训练的思路是在大部分计算中用低精度FP16 或 BF16只在关键环节保留 FP32。这样做有两个好处显存减半参数、梯度、激活值用 2 字节存储显存占用直接砍半。计算加速现代 GPU如 A100、H100对 FP16/BF16 有专门的 Tensor Core 加速矩阵乘法的吞吐量可以是 FP32 的数倍。但低精度也有风险数值范围窄容易溢出或下溢。所以混合精度训练不是简单地把所有东西换成 FP16而是有一套精细的机制来保证数值稳定性。3.2 FP16 与 BF16 的本质区别FP16 和 BF16 都是 16 位浮点数但它们的位分配完全不同格式符号位指数位尾数位数值范围精度FP321823±10³⁸高FP161510±65504中BF16187±10³⁸低关键差异在于指数位。BF16 的指数位和 FP32 一样是 8 位所以它的数值范围和 FP32 相同不会轻易溢出。代价是尾数只有 7 位精度比 FP16 低。FP16 的尾数有 10 位精度更高但指数只有 5 位最大只能表示 65504训练中很容易溢出。打个比方FP16 像一把刻度很细但量程很短的尺子BF16 像一把刻度粗但量程很长的尺子。训练中梯度值可能非常小10⁻⁸ 级别也可能突然变得很大所以量程比刻度更重要。这就是为什么现在大模型训练普遍首选 BF16。3.3 混合精度训练的完整流程混合精度训练的标准流程以 PyTorch 的 AMP 为例大致如下维护一份 FP32 的主权重副本模型参数在内存中保留 FP32 版本用于参数更新。前向传播用低精度将 FP32 权重量化到 FP16/BF16用低精度做矩阵乘法得到低精度的激活值。损失缩放Loss Scaling这是 FP16 训练的关键步骤。由于 FP16 下溢问题严重需要把损失值放大一个倍数如 2¹⁶使得梯度也相应放大避免变成 0。反向传播后再把梯度缩回原比例。反向传播用低精度计算低精度的梯度。参数更新用 FP32把低精度梯度转成 FP32加到 FP32 主权重上。BF16 因为数值范围和 FP32 一致通常不需要损失缩放流程更简单。这也是 BF16 在实际使用中更省心的原因之一。提示PyTorch 中可以通过torch.cuda.is_bf16_supported()检查当前 GPU 是否支持 BF16。A100、H100、RTX 30 系及以上基本都支持。如果硬件不支持 BF16就只能用 FP16 损失缩放。3.4 损失缩放的动态调整机制损失缩放不是固定倍数而是动态调整的。PyTorch 的GradScaler会自动做这件事初始缩放因子通常设为 2¹⁶。每次迭代检查梯度中是否有inf或nan。如果有说明缩放过头导致溢出了就把缩放因子减半并跳过这一步的参数更新。如果连续若干步默认 2000 步没有出现溢出就把缩放因子加倍逐步逼近最优值。这个机制的好处是自适应不需要手动调。但要注意如果训练中频繁出现溢出说明模型或数据可能有问题比如学习率太大、数据中有异常值等不能只靠损失缩放来掩盖。4. 混合精度训练的实操配置4.1 PyTorch AMP 的标准写法下面是一个典型的 PyTorch 混合精度训练代码骨架我加了详细注释说明每一步的意图import torch from torch.cuda.amp import autocast, GradScaler model MyModel().cuda() optimizer torch.optim.Adam(model.parameters(), lr1e-4) # BF16 不需要 scalerFP16 需要 use_bf16 torch.cuda.is_bf16_supported() scaler GradScaler(enablednot use_bf16) for batch in dataloader: optimizer.zero_grad() # 前向传播自动选择低精度 with autocast(dtypetorch.bfloat16 if use_bf16 else torch.float16): outputs model(batch) loss loss_fn(outputs, targets) if use_bf16: # BF16 直接反向传播 loss.backward() optimizer.step() else: # FP16 需要缩放损失 scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()这段代码里几个关键点autocast上下文管理器会自动决定哪些操作用低精度、哪些保持 FP32。比如矩阵乘法用低精度但 softmax、layer norm 等对数值敏感的操作会保持 FP32。GradScaler只在 FP16 下启用。BF16 因为不需要损失缩放直接反向传播即可。scaler.step(optimizer)内部会先检查梯度是否有效再决定是否更新参数。4.2 哪些操作必须保持 FP32混合精度不是所有操作都降到低精度。以下操作通常需要保持 FP32否则容易出问题Softmax涉及指数运算FP16 下容易溢出。Layer Normalization需要计算均值和方差低精度下误差累积明显。损失函数特别是交叉熵涉及 log 运算低精度下不稳定。参数更新必须用 FP32否则累积误差会导致训练发散。PyTorch 的autocast已经内置了这些规则你不需要手动指定。但如果你自己写了自定义算子就要注意把这些操作排除在低精度之外。4.3 梯度检查点的配合使用梯度检查点Gradient Checkpointing是另一个省显存的大杀器。它的原理是前向传播时只保存部分层的激活值其余层的激活值在反向传播时重新计算。这样激活值显存可以大幅降低代价是增加一次前向传播的计算量。在 PyTorch 中使用很简单from torch.utils.checkpoint import checkpoint class MyModel(nn.Module): def forward(self, x): # 对显存占用大的层使用 checkpoint x checkpoint(self.layer1, x) x checkpoint(self.layer2, x) return x实测下来梯度检查点通常能把激活值显存降到原来的 30%-50%而训练速度只慢 20%-30%。在显存紧张的情况下这个取舍非常划算。注意梯度检查点和混合精度可以同时使用但要注意checkpoint函数在 FP16 下需要额外处理确保重计算时的数值一致性。PyTorch 较新版本已经支持在autocast下正常使用。5. int8 与 bf16 模型的区别5.1 量化与混合精度是两回事很多人把 int8 和 bf16 混为一谈觉得都是“降低精度省显存”。其实它们属于两个不同的技术路线混合精度BF16/FP16仍然是浮点数只是位数减少。训练时用低精度计算但保留 FP32 主权重。主要目的是加速计算和减少显存。量化int8/int4把浮点数映射到整数表示。通常用于推理阶段把训练好的模型压缩成低比特格式减少存储和推理成本。简单说BF16 是训练时的精度选择int8 更多是推理时的压缩手段。当然现在也有 int8 训练的研究如 8-bit Optimizer但主流训练还是用 BF16/FP16。5.2 int8 量化的基本原理int8 量化的核心是找到一个缩放因子scale和零点zero point把浮点数线性映射到 [-128, 127] 的整数区间int8_value round(float_value / scale) zero_point反量化时float_value (int8_value - zero_point) × scale这个过程的精度损失主要来自两个方面一是舍入误差二是缩放因子的选择。如果缩放因子选得不好小数值会被压缩成 0大数值会饱和。实际中常用的方案是分通道量化per-channel quantization即每个通道用不同的缩放因子而不是整个张量共用一个。这样能更好地适应不同通道的数值分布。5.3 训练场景下如何选择在实际训练中我的建议是优先用 BF16只要硬件支持BF16 是训练的首选。数值范围大不需要损失缩放省心。硬件不支持 BF16 时用 FP16配合动态损失缩放也能稳定训练但需要多留意梯度溢出。int8 用于推理或特定优化器比如 bitsandbytes 的 8-bit Adam可以把优化器状态从 FP32 压到 int8显存直接省 75%。这在显存极度紧张时非常有用。场景推荐精度理由预训练BF16数值稳定速度快微调BF16 或 FP16视硬件支持而定推理int8/int4压缩模型提升吞吐优化器状态int88-bit Adam省显存精度损失可接受5.4 8-bit Adam 的实操与效果8-bit Adam 是 bitsandbytes 库提供的一个优化器它把 Adam 的动量和方差从 FP32 量化到 int8 存储使用时再反量化。实测下来7B 模型的优化器状态从 84GB 降到约 21GB省了 75% 的显存而训练效果几乎无损。使用方式import bitsandbytes as bnb optimizer bnb.optim.Adam8bit( model.parameters(), lr1e-4, betas(0.9, 0.999) )需要注意的是8-bit Adam 在训练初期可能会有轻微的损失波动但通常在几百步后就恢复正常。如果你的任务对数值极其敏感建议先在小规模上验证。6. 常见问题与排查技巧6.1 显存溢出OOM的排查思路OOM 是最常见的问题。排查时按以下顺序检查确认模型参数量用sum(p.numel() for p in model.parameters())算一下看看和预期是否一致。检查 batch size 和序列长度这两个参数对激活值影响最大。先试着把 batch size 降到 1看是否能跑通。确认优化器状态如果你用的是 Adam显存开销是 SGD 的三倍以上。考虑换 8-bit Adam 或 ZeRO。开启梯度检查点如果激活值是瓶颈这个最有效。检查是否有内存泄漏比如在训练循环中不断累积张量而没有释放。6.2 损失变成 NaN 怎么办FP16 训练中损失变 NaN 很常见原因通常是梯度溢出。处理步骤检查损失缩放因子是否过大可以手动调小初始值。检查学习率是否太大试着降低 10 倍。检查数据中是否有异常值如 inf、nan。如果用的是 FP16考虑换 BF16。BF16 几乎不会出现溢出问题。6.3 混合精度训练速度反而变慢这种情况通常有几个原因GPU 不支持 Tensor Core老卡如 GTX 10 系没有 FP16 加速混合精度反而增加转换开销。频繁的精度转换如果代码中频繁在 FP32 和 FP16 之间转换开销会抵消加速收益。batch size 太小Tensor Core 需要足够大的矩阵才能发挥优势batch size 太小时加速不明显。6.4 常见问题速查表问题现象可能原因解决方法OOM优化器状态太大换 8-bit Adam 或 ZeROOOM激活值太大开启梯度检查点减小 batch size损失 NaNFP16 梯度溢出换 BF16 或降低损失缩放因子训练变慢硬件不支持 Tensor Core检查 GPU 型号必要时回退 FP32精度下降量化损失过大检查缩放因子改用量化感知训练损失不收敛学习率与精度不匹配降低学习率检查 warmup 设置7. 一些实操中的经验体会显存估计这件事理论计算和实际占用往往有出入因为框架的实现细节、CUDA 的内存分配策略、碎片化等因素都会影响。我的习惯是先用理论算一个下限然后实际跑一个小 batch 看真实占用再按比例推算。比如先用 batch size1 跑通看torch.cuda.max_memory_allocated()返回多少然后根据目标 batch size 线性外推激活值大致线性但要注意 attention 的 O(n²) 部分。混合精度方面BF16 确实比 FP16 省心太多。我早期用 FP16 训练时经常要盯着损失缩放因子调来调去换了 BF16 之后基本没再管过数值稳定性问题。唯一的代价是 BF16 的精度略低但在大模型训练中这点精度损失对最终效果的影响微乎其微。最后分享一个小技巧如果你在单卡上调试可以用torch.cuda.memory_summary()打印详细的内存分配报告它会告诉你每一块显存被什么占用包括已分配、已缓存、碎片等。这个工具在排查 OOM 时非常有用比盲目猜测高效得多。
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。

↑