大模型训练绕不开两个硬骨头:显存不够和精度怎么选。我见过太多团队把模型并行度调来调去,最后发现是显存估算从一开始就错了;也见过有人一上来就开FP16,结果loss曲线直接炸成烟花。这篇就围绕大模型训练显存估计和混合精度训练这两件事,把我在实际项目里踩过的坑、算过的账、调过的参数,原原本本摊开讲一遍。不管你是刚接触大模型训练的新手,还是已经跑过几轮微调的老手,只要你在为“为什么OOM”“BF16和FP16到底选哪个”“INT8能不能用来训练”这些问题头疼,下面的内容应该能帮你省下不少试错时间。
1. 显存估计:先把账算清楚再动手
1.1 显存到底被谁吃掉了
很多人一看到OOM就去加卡、降batch size,但从来没认真算过显存到底花在哪。大模型训练的显存占用可以拆成四大块:模型参数、梯度、优化器状态、激活值。前三个是“静态开销”,跟batch size关系不大;激活值是“动态开销”,随batch size和序列长度线性增长。
先看静态部分。假设模型有 ( \Phi ) 个参数,用Adam系列优化器做全参数训练:
- 模型参数(FP32):( 4\Phi ) 字节
- 梯度(FP32):( 4\Phi ) 字节
- Adam的一阶动量:( 4\Phi ) 字节
- Adam的二阶动量:( 4\Phi ) 字节
- 如果用了混合精度,还需要额外保存一份FP16/BF16的参数副本:( 2\Phi ) 字节
加起来就是 ( 4+4+4+4+2 = 18\Phi ) 字节。也就是说,一个10B参数的模型,光静态显存就要吃掉180GB。这就是为什么单卡80GB根本放不下10B模型的全参数训练——还没算激活值呢。
注意:这里说的是全参数训练。如果你做的是LoRA这类参数高效微调,优化器状态只覆盖LoRA旁路的那部分参数,静态开销会小一到两个数量级。
激活值的估算稍微复杂一点。它跟网络结构、序列长度、batch size、是否用梯度检查点都有关系。一个粗略的经验公式是:
[ \text{激活显存} \approx b \cdot s \cdot h \cdot L \cdot k ]
其中 ( b ) 是batch size,( s ) 是序列长度,( h ) 是隐藏维度,( L ) 是层数,( k ) 是一个跟具体结构相关的系数,通常在10到20之间。这个公式只能用来做数量级判断,精确值必须靠实测。
1.2 手把手做一次显存估算
我拿一个具体的配置来演示。假设你要训练一个7B参数的模型,配置如下:
| 项目 | 数值 |
|---|---|
| 参数量 | 7B |
| 精度 | BF16混合精度 + FP32主权重 |
| 优化器 | AdamW |
| 序列长度 | 2048 |
| batch size | 8 |
| 层数 | 32 |
| 隐藏维度 | 4096 |
| 梯度检查点 | 开启 |
第一步,算静态显存。
- FP32主权重:( 7 \times 10^9 \times 4 = 28 ) GB
- BF16参数副本:( 7 \times 10^9 \times 2 = 14 ) GB
- FP32梯度:( 7 \times 10^9 \times 4 = 28 ) GB
- Adam一阶动量:28 GB
- Adam二阶动量:28 GB
静态合计:( 28+14+28+28+28 = 126 ) GB。
第二步,算激活显存。
开启梯度检查点后,激活显存大约降到原来的 ( 1/\sqrt{L} ) 左右。不开检查点的粗略估计是:
[ 8 \times 2048 \times 4096 \times 32 \times 15 \approx 3.2 \times 10^{13} \text{ 字节} \approx 32 \text{ GB} ]
开检查点后大概降到 6~8 GB 量级。这个数字波动很大,实际以profiler为准。
第三步,加总。
静态126GB + 激活7GB ≈ 133GB。单张80GB卡肯定放不下,至少需要2张卡做ZeRO-2或者ZeRO-3切分。如果切到8张卡上,每卡静态开销约16GB,加上激活和通信buffer,单卡占用大概25~30GB,跑起来比较舒服。
这个估算过程看起来简单,但实际项目里我见过太多人跳过这一步,直接拍脑袋上8卡,结果发现显存利用率只有40%,白白浪费算力。
1.3 显存优化的几个实用手段
算清楚账之后,如果发现显存不够,有几个手段可以按优先级尝试:
梯度检查点(Gradient Checkpointing)是最划算的。它用计算换显存,把激活显存降到原来的 ( 1/\sqrt{L} ) 甚至更低,代价是反向传播时多算一次前向,训练速度大概慢20%~30%。对于显存紧张的场景,这个代价完全可以接受。
ZeRO系列切分是必选项。ZeRO-1切优化器状态,ZeRO-2再切梯度,ZeRO-3把参数也切了。切得越狠,通信开销越大。我的经验是:能放下就用ZeRO-2,放不下再上ZeRO-3,因为ZeRO-3的通信量会显著拖慢训练速度。
降低精度要谨慎。把优化器状态从FP32降到BF16能省一半静态显存,但训练稳定性会变差,尤其是小模型上更容易出问题。我一般只在显存极度紧张且模型足够大(>13B)时才考虑这个选项。
减小batch size或序列长度是最后的手段。batch size太小会影响梯度估计的稳定性,序列长度太短则可能截断重要上下文。如果非要做,优先减batch size,因为序列长度往往跟任务语义强相关。
2. 混合精度训练:BF16和FP16到底怎么选
2.1 从FP32到FP16再到BF16的演进逻辑
要理解混合精度,先得理解为什么需要它。FP32有23位尾数和8位指数,数值范围大约是 ( 10^{-38} ) 到 ( 10^{38} ),精度很高但存储和计算开销大。FP16把尾数砍到10位,指数也砍到5位,存储减半、计算速度提升,但数值范围缩到了 ( 10^{-5} ) 到 ( 10^4 ) 左右。这个范围太窄了,训练时梯度很容易下溢成0或者上溢成inf。
BF16是Google提出来的折中方案。它保留了FP32的8位指数,只把尾数砍到7位。所以BF16的数值范围和FP32一样大,不会溢出,但精度比FP16低。对于深度学习训练来说,数值范围比精度更重要,因为梯度动态范围极大,溢出是致命问题,而精度损失可以通过其他手段补偿。
我用一个生活化的类比来解释:FP32像一把精确到毫米的卷尺,量程1米;FP16像精确到厘米的卷尺,量程只有30厘米;BF16像精确到厘米的卷尺,量程也是1米。量衣服的时候,量程不够比精度不够更让人抓狂。
2.2 BF16和FP16的实测对比
我在同一个7B模型上分别用BF16和FP16跑过SFT训练,配置完全一致,结果差异很明显:
| 对比项 | BF16 | FP16 |
|---|---|---|
| 是否需要loss scaling | 不需要 | 需要 |
| 训练稳定性 | 高,几乎不溢出 | 中,需要调初始scale |
| 收敛速度 | 略慢 | 略快 |
| 最终loss | 基本一致 | 基本一致 |
| 硬件支持 | Ampere及以上 | 几乎所有现代GPU |
| 显存占用 | 相同 | 相同 |
BF16最大的优势是不需要loss scaling。FP16训练必须配一个动态loss scaler,初始scale设大了会溢出,设小了梯度下溢,调参本身就是一门玄学。我遇到过好几次FP16训练跑了几千步突然loss爆掉,查半天发现是scaler的growth interval设得太激进。
BF16的劣势是尾数少,理论上精度低。但在大模型上这个差异几乎可以忽略,因为大模型本身参数量大,对单个参数的精度不敏感。反而是在小模型(<1B)上,BF16的精度损失会稍微明显一些。
实操建议:如果你的卡支持BF16(Ampere架构及以上),无脑选BF16。如果是更老的卡只支持FP16,那就老老实实配loss scaler,初始scale设成 ( 2^{16} ),growth interval设成2000步,backoff factor设成0.5。
2.3 混合精度训练的完整配置
混合精度不是简单地把dtype改成bf16就完事了,它涉及一整套配置。我以PyTorch的AMP(Automatic Mixed Precision)为例,把关键配置拆开讲。
模型权重的主副本必须保持FP32。这是混合精度的核心。前向和反向用BF16/FP16算,但优化器更新权重时用的是FP32主副本。如果主副本也是低精度,训练几百步后权重就会因为累积舍入误差而漂移。
哪些算子用低精度,哪些保持FP32。矩阵乘法、卷积这类计算密集型算子用低精度收益最大;softmax、layer norm、loss计算这类对数值范围敏感的算子必须保持FP32。PyTorch的AMP会自动处理这些,但如果你手写kernel,就得自己判断。
梯度累积时的精度处理。如果你用梯度累积来模拟大batch,累积的梯度建议用FP32存,最后一步再转成低精度做all-reduce。我见过有人用BF16累积梯度,跑了半天发现梯度全是0,就是因为累积过程中下溢了。
一个典型的AMP配置长这样:
from torch.cuda.amp import autocast, GradScaler scaler = GradScaler(enabled=use_fp16) # BF16不需要scaler for batch in dataloader: optimizer.zero_grad() with autocast(dtype=torch.bfloat16): outputs = model(**batch) loss = loss_fn(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()如果用BF16,把GradScaler的enabled设成False,或者直接用torch.bfloat16的autocast上下文,不需要scaler。
2.4 混合精度训练的常见坑
坑一:loss scaling的初始值。FP16训练时,初始scale设成 ( 2^{16} ) 是个比较稳的起点。设太小梯度下溢,设太大直接溢出。我一般会先跑100步观察scale的变化趋势,如果scale一直在降,说明初始值设大了;如果一直不涨,说明设小了。
坑二:BatchNorm和LayerNorm的精度。这两个归一化层对数值范围很敏感,必须用FP32算。PyTorch的AMP会自动把LayerNorm的统计量计算放在FP32,但如果你自己实现了归一化层,记得手动指定dtype。
坑三:优化器状态的精度。Adam的动量和方差建议用FP32存。用BF16存的话,动量更新时的累积误差会慢慢累积,训练后期loss会抖动。这个坑我在一个13B模型上踩过,前10K步一切正常,后面loss开始周期性震荡,查了一周才发现是优化器状态精度的问题。
坑四:梯度裁剪和低精度的交互。梯度裁剪通常在FP32梯度上做,如果你在低精度梯度上裁剪,裁剪阈值需要相应调整。我一般先把梯度转成FP32再裁剪,避免精度问题。
3. INT8训练:能不能用,什么时候用
3.1 INT8和BF16的本质区别
热搜词里有人问“int8和bf16模型的区别”,这个问题问到了点子上。BF16是一种浮点格式,有指数位和尾数位,能表示非常大和非常小的数;INT8是定点格式,只有8个bit表示整数,范围是-128到127。两者的数值表示能力完全不同。
BF16适合训练,因为训练需要处理动态范围极大的梯度。INT8适合推理,因为推理时的激活值和权重分布相对集中,可以通过量化校准把范围压到INT8能表示的区间内。用INT8做训练目前还不是主流,因为梯度的动态范围太大,量化误差会严重破坏训练稳定性。
3.2 INT8量化的基本原理
INT8量化的核心是找到一个缩放因子(scale)和零点(zero point),把浮点数映射到整数区间。公式是:
[ x_{int8} = \text{round}\left(\frac{x_{fp32}}{s}\right) + z ]
其中 ( s ) 是scale,( z ) 是zero point。推理时反量化回来:
[ x_{fp32} \approx s \cdot (x_{int8} - z) ]
量化的关键是确定 ( s ) 和 ( z )。常见的方法有对称量化和非对称量化。对称量化把zero point固定为0,适合权重这种分布对称的数据;非对称量化让zero point可学习,适合激活值这种分布偏移的数据。
3.3 INT8在训练中的实际应用
虽然INT8训练不主流,但在某些场景下确实有用。比如INT8优化器状态,把Adam的动量和方差量化成INT8存储,能省75%的优化器显存。代价是训练稳定性下降,需要配合error feedback机制来补偿量化误差。
另一个场景是INT8梯度通信。在多卡训练时,梯度all-reduce是通信瓶颈。把梯度量化成INT8再通信,通信量降到1/4,代价是引入量化噪声。这个方案在带宽受限的集群上比较有吸引力。
我实测过INT8优化器状态在7B模型上的效果:显存从126GB降到约70GB,但loss比BF16高0.02左右,而且训练后期有轻微震荡。如果你的显存极度紧张且能接受一点精度损失,可以试试;否则还是老老实实用BF16。
注意:INT8训练目前没有成熟的自动混合精度框架支持,需要手动实现量化逻辑。如果你不是特别清楚自己在做什么,建议不要在生产环境用INT8训练。
4. 显存和精度的联合调优实战
4.1 一个完整的调优案例
我拿一个真实项目来串一遍。需求是微调一个13B模型,硬件是4张A100 80GB,目标是在保证收敛的前提下尽量用大batch size。
第一轮:全参数BF16 + ZeRO-2。
静态显存:( 13 \times 18 = 234 ) GB,切到4卡上每卡约58.5GB。加上激活值和通信buffer,单卡占用约70GB,刚好卡在80GB边缘。batch size只能设到4,再大就OOM。
第二轮:加梯度检查点。
激活显存从约12GB降到约3GB,单卡占用降到约62GB。batch size可以提到8,训练速度因为重算前向慢了约25%,但吞吐量(样本/秒)反而提升了。
第三轮:换ZeRO-3。
静态显存切得更细,每卡约45GB,加上激活约48GB。batch size提到16,但通信开销明显增加,单步时间从1.2秒涨到1.8秒。吞吐量跟第二轮差不多,但显存余量更大,可以再提序列长度。
最终选择:第二轮方案。ZeRO-3的通信开销不值得,梯度检查点的性价比最高。
4.2 调优的决策树
基于上面的经验,我整理了一个显存和精度调优的决策流程:
- 先算静态显存,确定最少需要几张卡
- 如果单卡放不下静态显存,上ZeRO-2;还不够就ZeRO-3
- 静态显存放下后,看激活显存。如果OOM,先开梯度检查点
- 梯度检查点开了还OOM,再考虑降batch size或序列长度
- 精度优先选BF16,硬件不支持再选FP16
- INT8只在显存极度紧张且能接受精度损失时考虑
这个顺序的逻辑是:优先用通信换显存,其次用计算换显存,最后才牺牲训练质量。降batch size和序列长度会影响模型效果,能不用就不用。
4.3 监控和profiling
调优不能靠猜,得靠数据。我一般用两个工具:PyTorch Profiler看算子和显存分配,nvidia-smi看整体占用。Profiler能告诉你显存峰值出现在哪个算子、哪个阶段,比只看总数有用得多。
一个实用的技巧是:在训练循环里每隔100步打印一次显存占用,观察它的变化趋势。如果显存随步数缓慢增长,说明有内存泄漏(通常是某个tensor没释放);如果显存突然跳变,说明某个batch的序列长度异常。
5. 常见问题速查
5.1 显存相关
| 问题 | 可能原因 | 解决方法 |
|---|---|---|
| 训练启动就OOM | 静态显存超了 | 加卡或上ZeRO |
| 跑几十步后OOM | 激活显存累积或泄漏 | 开梯度检查点,检查tensor释放 |
| 显存占用远低于预期 | 估算公式系数偏大 | 以profiler实测为准 |
| 多卡显存不均 | ZeRO切分不均衡 | 检查参数切分策略 |
5.2 精度相关
| 问题 | 可能原因 | 解决方法 |
|---|---|---|
| FP16训练loss变NaN | loss scale太大 | 降低初始scale,加backoff |
| BF16训练loss不降 | 学习率不匹配 | BF16可以适当调大学习率 |
| 训练后期loss震荡 | 优化器状态精度不够 | 优化器状态用FP32 |
| 梯度全是0 | 梯度下溢 | 检查是否用了低精度累积梯度 |
5.3 我踩过的几个坑
坑一:以为BF16不需要任何配置。实际上BF16虽然不需要loss scaler,但学习率需要重新调。BF16的梯度精度比FP32低,同样的学习率下更新步长会有偏差。我一般会把学习率调大10%~20%。
坑二:梯度检查点开在所有层上。梯度检查点不是免费的,每开一层就多一次前向重算。我一般只开在显存占用最大的那几层,比如attention层,FFN层视情况开。全开的话训练速度会慢40%以上。
坑三:忽略通信buffer的显存。多卡训练时,all-reduce需要额外的通信buffer,大小跟梯度大小相当。4卡ZeRO-2训练13B模型时,通信buffer大概占5~8GB,估算时容易漏掉。
坑四:用INT8存优化器状态但没加error feedback。没有error feedback的INT8量化会让小梯度直接变成0,训练后期loss完全不动。加error feedback后情况好转,但实现复杂度不低。
6. 一些个人经验
显存估计这件事,公式只能给你一个起点,真正的数字必须靠实测。我现在的习惯是:新模型先跑一个batch,用profiler把显存分布打出来,再决定并行策略。这样比拍脑袋估算靠谱得多。
混合精度方面,BF16已经是默认选项了。除非硬件不支持,否则没必要纠结FP16。FP16的loss scaling调参成本太高,省下来的那点收敛速度不值得。
INT8训练我持保留态度。推理量化已经很成熟了,但训练量化还在早期。如果你的显存真的紧张到必须用INT8,那可能说明模型规模跟硬件不匹配,考虑换个更小的模型或者用LoRA这类参数高效方法,比硬上INT8更划算。
最后分享一个小技巧:训练时把torch.cuda.memory_summary()的输出存到日志里,出问题时翻日志比重新跑一遍快得多。这个习惯帮我省过好几次通宵调试。