FP16一开就loss炸了,TF32到底要不要开,大模型微调怎么把显存降下来,这类问题最近被反复问到。大模型训练和微调的成本越来越高,大家第一个想动手优化的就是混合精度训练,想用FP16配合TF32把显存占用和算力消耗压下来。这个思路是对的,尤其是当数据以token形式进入模型时,每一个token都要完整走一遍前向和反向,显存和算力消耗几乎跟token总数挂钩。这里我把自己踩过的坑、验证过的方案都摊开讲一讲,希望能帮你把这套流程跑通。
1. 混合精度训练到底在解决什么问题?
1.1 FP16、FP32、TF32的关系先搞明白
想用好混合精度训练,第一步不是写代码,是把浮点格式搞明白。一个浮点数在计算机里由符号位、指数位、尾数位组成,位数不同,能表达的范围和精度就不同:
- FP32:1位符号 + 8位指数 + 23位尾数,总共32位。这是深度学习之前默认的精度,也叫单精度。
- FP16:1位符号 + 5位指数 + 10位尾数,总共16位。指数位只有5位,导致它的动态范围特别小,最大只能到65504,最小正常数是6.1e-5左右。
- TF32:这个最容易被误会,很多人以为它是32位格式,实际上它是NVIDIA Ampere架构开始引入的Tensor Core计算格式,内部是1位符号 + 8位指数 + 10位尾数,总共19位,但在32位容器里存储和处理。指数范围跟FP32一致,所以动态范围很大,代价是尾数只有10位,做乘法时有效精度其实和FP16一个水平。
- BF16:1位符号 + 8位指数 + 7位尾数,总共16位。指数位和FP32完全一样,动态范围合适,但尾数位只有7位,有效精度比FP16还低。
很多训练卡在FP16上,就是因为低估了FP16的窄范围。打个比方:FP32像一把量程30厘米、刻度到毫米的尺;FP16量程缩到6厘米,刻度还是十分之一毫米;BF16量程够大但刻度变粗;TF32是量程大、刻度中等偏粗的一个变种。正常数值在FP32里没问题,切到FP16后要么溢出变成inf,要么小数部分直接归零,一下就把训练弄崩了。
1.2 为什么说省显存就能降低token消耗
模型训练里,输入文本被切分成token后,通常以[batch_size, seq_len]的索引矩阵进入网络,经过embedding后变成[batch_size, seq_len, hidden_size]的高维张量。从这一层开始,每一层注意力机制、每一层FFN的输入输出和中间激活,规模都跟batch内的token总数成正比。
也就是说,在显存受限的情况下,能塞进GPU的token总数决定了训练效率。想要多塞token,无非两条路:增大batch size,或者拉长序列长度,但都需要更多显存。混合精度训练把权重和中间激活的一部分从32位砍到16位,同等显存预算下能直接多装一批token。这里的"减少token消耗",指的是模型每处理一个token所占用的显存和算力成本下降了,并不是说模型输出变短了。
这个账在推理场景更好算。多轮对话或者长上下文场景里,模型会为每个已经出现的token缓存一组Key和Value,也就是KV Cache。假设hidden_size是4096,模型有32层,用FP32存KV Cache的话,每个token每层要存两个向量,大约4096×2×2×32=512KB。同样的模型切成FP16,只算KV Cache一个token就能省256KB。上下文一旦拉到几万token,差距就非常明显了。
1.3 Tensor Core真正把算力红利吃进去
省显存只是混合精度的一半好处,另一半是速度。NVIDIA早期从Volta架构开始加入Tensor Core,专门为低精度矩阵乘设计。FP16输入在Tensor Core上的理论吞吐通常是FP32的好几倍,实际训练中配合CUDA库优化,跑BERT、GPT这类Transformer模型,很多算子都能获得实质性加速。
这里面有个容易被忽略的点:Tensor Core不是自动对所有操作生效的。在PyTorch中,要启用TF32需要显式设置开关,FP16则需要通过AMP这样的机制把前向计算切到FP16路径。否则代码根本没走到Tensor Core上,白瞎了硬件能力。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 方案选型:FP16、TF32、BF16,到底用哪个?
2.1 三种精度横向对比
我把常用的三种方案放在一起对比,方便做技术选型。
| 项目 | FP16 | BF16 | TF32 |
|---|---|---|---|
| 总位数 | 16位 | 16位 | 19位(32位容器) |
| 指数位 | 5位 | 8位 | 8位 |
| 尾数位 | 10位 | 7位 | 10位 |
| 最大数值 | 65504 | 约3.4e38 | 约3.4e38 |
| 最小正常数 | 约6.1e-5 | 约1.2e-38 | 约1.2e-38 |
| 典型用途 | 混合精度训练、推理 | 大模型训练 | 矩阵乘、卷积加速 |
| 适用硬件 | Volta及以后 | Ampere及以后 | Ampere及以后 |
| 需要loss scaling | 需要 | 不需要 | 不需要 |
BF16虽然总位数也是16位,但指数位保留了FP32的水平,动态范围非常大,所以训练中很少遇到溢出问题,正是这个特性让它在大模型训练里迅速流行起来。但代价是尾数位只有7位,有效精度大约3位十进制数字,部分任务如果对数值特别敏感,还是会有精度损失。
TF32的尾数位有10位,比BF16稍好,但它不是通用的张量存储格式,主要作用在矩阵乘法和卷积上,层归一化、Softmax、损失函数这些算子仍然以FP32计算。
2.2 为什么不推荐直接全FP16
很多人一开始图省事,想把模型所有参数直接切成FP16,结果跑两步就loss爆炸,然后得出结论说FP16不能用。其实这个尝试方向就错了。
FP16的问题不只是数值范围窄,还有一个更隐蔽的点:权重更新时,学习率通常只有1e-4到1e-5量级,FP32下正常的小数值,在FP16里可能直接变成0,梯度根本没法更新。
混合精度训练的思路不是让所有东西都变成FP16,而是让关键路径保持FP32。权重本身以FP32保存,前向和反向的计算过程用FP16跑,梯度计算完累积到FP32的master weights上,优化器更新的是FP32权重,更新完再转回FP16供下一轮计算使用。也就是说,只有计算和中间激活是FP16,真正存模型参数的主副本一直是FP32。这个"混合"指的就是两层精度并存,不是一刀切。
2.3 硬件条件决定方案选择
我的建议是:如果是A100、H100、4090、3090这些卡,训练新模型优先考虑BF16,不用跟loss scaling纠缠,省心稳定;如果手头是V100、T4这类老卡,或者对精度特别敏感、想尽量少改代码,就走FP16+loss scaling;如果模型不大,只是想白嫖一点加速且严格不想动训练逻辑,TF32是最低成本的方案。
这里特别提醒一下:TF32只在Ampere及以后的架构上真正有硬件支持。老卡上设置allow_tf32不会报错,但也不会生效,可能开了跟没开一样。
3. 实操:把训练脚本切成FP16+TF32
3.1 PyTorch AMP推荐写法
PyTorch里最省事的混合精度方案是AMP,全称Automatic Mixed Precision。PyTorch 1.6之后内置支持,不需要额外装库。以训练一个Transformer模型为例:
python复制import torch
from torch.amp import GradScaler, autocast
device = "cuda"
model = create_model().to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=2e-4)
# 开启TF32,只对Ampere及以上架构生效
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
scaler = GradScaler("cuda")
for step, (input_ids, labels) in enumerate(train_loader):
input_ids = input_ids.to(device)
labels = labels.to(device)
optimizer.zero_grad(set_to_none=True)
with autocast("cuda"):
logits = model(input_ids)
loss = criterion(logits, labels)
scaler.scale(loss).backward()
if max_grad_norm is not None:
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm)
scaler.step(optimizer)
scaler.update()
这段代码是整个混合精度训练的核心,每一步都值得细看:
with autocast("cuda"):只影响上下文里的算子。PyTorch会自动判断哪些算子适合降精度,哪些数值敏感的算子(比如Softmax、LayerNorm、交叉熵)仍然用FP32,不需要手动标记。scaler.scale(loss).backward():对loss乘一个大数,默认初始值是65536,把计算出来的梯度放大到FP16可表达的范围内,防止下溢变成0。scaler.unscale_(optimizer):在梯度裁剪之前必须调用,它把梯度从放大后的值还原成真实值。如果你不调用而直接clip,裁的是放大后的梯度,梯度裁剪直接失效。scaler.step(optimizer):内部会检查unscaled之后是否有inf或NaN,如果有就跳过这轮参数更新,避免把坏梯度写进权重。scaler.update():每个batch结束后根据最近一段时间的梯度情况调整scale,如果一直正常,scale会周期性增大,一旦出现溢出就自动回退。
3.2 TF32的两个开关容易漏
TF32在PyTorch里的开关其实不止一个。最常用的是这两个:
python复制torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
第一个影响矩阵乘法,训练Transformer的线性层主要走这条;第二个影响cuDNN里的卷积。很多人只开了第一个,发现卷积网络加速不明显,原因就在这里。
有个小坑需要提醒:PyTorch官方对cuDNN的TF32限定在有限场景,部分情况下精度损失可能比预期大。如果你的任务对精度特别敏感,可以先只开matmul的TF32,对比验证集指标再决定要不要开卷积的TF32。
怎么确认TF32真的生效了?最直接的办法是找一个批量矩阵乘,分别开关TF32跑一下耗时对比。也可以用NVIDIA Nsight Compute做profile,查看kernel是不是走在了Tensor Core的TF32路径上。如果开了之后速度和精度都没有任何变化,先检查显卡型号。
3.3 配合gradient checkpointing和DeepSpeed,进一步压显存
混合精度只是第一步。实际训练大模型的时候,光有AMP还不够,因为显存大头中,中间激活值占了很大比例。可以使用gradient checkpointing把前向过程的激活值丢掉一部分,反向传播时再重新计算:
python复制model.gradient_checkpointing_enable()
这个方法以时间换显存,会让训练变慢一些,但能让batch size或者序列长度大幅提升。我的经验是,跟FP16一起用,训练吞吐反而经常比原来更高,因为显存宽裕之后batch可以堆更大,通信和kernel启动开销被摊薄了。
如果你用DeepSpeed训练,其实可以不用手写AMP,直接在配置里开:
json复制{
"fp16": {
"enabled": true,
"auto_cast": true,
"loss_scale": 0,
"initial_scale_power": 16,
"loss_scale_window": 1000,
"hysteresis": 2,
"min_loss_scale": 1
}
}
如果想用BF16,配置长这样:
json复制{
"bf16": {
"enabled": true
}
}
DeepSpeed的好处是默认处理好loss scaling,还支持ZeRO把优化器状态、梯度、参数分片到多卡,进一步降低单卡显存压力,让相同Token总量下整体显存消耗明显下降。
4. 常见问题与排查技巧实录
4.1 一开FP16就loss爆炸,排查顺序很重要
这个问题几乎每个人都会遇到,我从自己的踩坑经验出发,总结了一套排查顺序:
| 现象 | 可能原因 | 排查方法 |
|---|---|---|
| 训练几步后loss变NaN | 梯度溢出,loss scale一直回退不了 | 调低initial_scale_power,比如16改成10或8 |
| 参数更新不生效,loss长期不降 | scaler检测到inf/NaN后持续跳过step | 打印scaler的scale值,观察是否频繁回退 |
| 梯度裁剪失效 | 在unscale之前调用了clip_grad_norm_ | 确保先unscale_,再clip |
| 学习率微调后直接崩 | FP16对学习率更敏感 | 尝试降低lr,或改用warmup更长的schedule |
| 数据里存在异常值 | 某些样本产生极大梯度 | 检查输入数据分布,定位异常样本 |
| 一直不收敛,但不报错 | FP16精度不足以支撑当前任务 | 换BF16试一下,排除精度天花板问题 |
有一个比较隐蔽的问题:很多人把scaler.update()放在optimizer.step()之前,导致scale更新时机不对,出现间歇性不收敛。正确顺序是:backward -> unscale_ -> clip -> step -> update。
还有一次,我排查了很久发现是loss.item()打印出来的值本身没问题,但backward之后梯度里有NaN。后来在unscale之后打印了梯度norm,才定位到是模型里某一层依赖了超出FP16范围的中间结果。这提醒我一件事:遇到问题先确认是loss的问题、梯度的问题还是权重的问题,不要上来就怀疑AMP框架。
4.2 TF32开了没效果,或者精度下降明显
TF32开了没效果,最常见的原因前面说过:显卡太老,不支持。第二个常见原因是算子类型限制,TF32只对矩阵乘法和卷积生效,如果你的模型主要耗时在别的算子上(比如大量自定义算子、频繁的维度变换),加速自然不明显。
如果开了TF32之后精度下降明显,我的建议是分开验证。先在验证集上对比FP32和TF32的指标,如果误差在千分位级别,通常可以接受;如果差很多,就需要看模型哪一部分对精度敏感。实践中,科学计算类的任务往往比NLP/CV任务对精度更敏感,这类项目建议训练时用TF32,推理时回到FP32,或者直接换BF16,因为BF16保留了FP32的动态范围,稳定性更好。
4.3 版本迁移的坑:torch.cuda.amp / apex / torch.amp
很多老项目还在用NVIDIA的apex,或者老写法torch.cuda.amp.autocast。如果你的项目还能跑,不一定要急着改,但新项目强烈建议直接用新方式:from torch.amp import autocast, GradScaler。新版写法统一了设备和DDP场景的体验,还在老写法上做了一些改进。
旧代码升级到新版时,需要注意几个不兼容点:
GradScaler("cuda")这里要传device参数,旧代码很多不传。autocast("cuda")里也要指定device,不再默认取全局设备。- apex的
amp.initialize和amp.scale_loss的写法跟PyTorch原生amp完全不一样,迁移时不要混着用。
另外,PyTorch 2.x时代很多人在用torch.compile加速。AMP和torch.compile一起使用时,一般把autocast放在模型的forward外面一层包住即可,不需要深入到内部算子手动标记。真遇到编译报错,先把compile关掉,排除精度问题之后再回来排查编译配置,别直接怀疑是精度开关的问题。
最后一点个人体会
这套方案跑通之后,我最大的感受是:训练成本高不高,很多时候不是看模型参数量,而是看每个token要占多少显存、要烧多少FFLOPS。混合精度训练是把单位token成本打下来的第一板斧,配合gradient checkpointing和合理的batch策略,经常比单纯换大卡、调模型结构更立竿见影。
最后分享一个实操心得:如果条件允许,大模型任务直接上BF16,把loss scaling这套麻烦彻底去掉。FP16+TF32这套留着处理老卡和特殊精度需求就够了。不管你先用了哪种,都值得把AMP的机制彻底搞懂,因为模型越做越大时,你迟早要在更复杂的并行策略里面对同样的精度管理和显存优化问题。
