如果你也和我一样,每天第一件事是打开 nvidia-smi 看显存和功耗,第二件事是翻训练日志里每秒蹦出来的 token 数,那么这篇文章你应该能直接用上。做 LLM 训练和推理的人这几年都有同一个感受:token 就是钱,算力就是命。模型每吞一个 token,背后都是实打实的前向计算、反向传播、显存读写和排队等待;而混合精度训练——FP16 配合 TF32——恰恰是在这些环节里压缩成本最有效的组合拳之一。
这不是什么新概念,但我在实际给团队搭训练管线、帮别人调模型的时候发现,绝大多数人只停留在"有个 FP16 能省显存"的层面,对 TF32 的原理、两者的适用边界、以及为什么开了之后训练会突然变 NaN,基本是一头雾水。所以这篇不只是列步骤,我会把底层机制、完整配置、踩坑排查和实测数据全部摊开,尽量让看的人少走弯路。适合正在做大模型训练、微调 LoRA、或者被高昂 token 成本压得喘不过气的工程师参考。
1. 先把账算明白:token消耗、算力占用和精度之间的三角关系
很多人在还没搞清楚自己瓶颈在哪的时候,就急着开 FP16,结果要么收益不明显,要么训练直接崩掉。我建议动手之前先把这个三角关系捋清楚:精度影响的是什么,token 消耗和算力占用又是从哪来的。
1.1 一个 token 从进入模型到离开,到底烧掉了什么
拿一个 decoder-only 的大模型举例。训练时每个 token 要经历前向传播和反向传播:前向的计算量大约是 2 倍参数量(单位 FLOPs),反向传播大约是前向的 2 倍,也就是 4 倍参数量,所以训练一个 token 的总计算量约等于 6 倍参数量 FLOPs。
以 7B 模型来算,每个 token 训练大概要烧 42 GFLOPs。听起来不多,但一个 batch 处理 4 万 token,就是 1680 TFLOPs。这时候 GPU 的算力就变成了硬约束:
- A100 的 FP32 算力大概 19.5 TFLOPS,处理这批数据要 86 秒;
- 同样的卡,FP16 配合 Tensor Core 能跑到 312 TFLOPS,只要 5.4 秒。
这就是为什么说"精度模式直接决定每秒钟能消化多少 token"。算力利用率上去了,单位时间内能处理的 token 数量才能上去,单 token 的成本才会降下来。
1.2 算力占用的本质是"每 token 的单位成本"和显存配额
除了计算时间,显存同样重要。模型权重、梯度、优化器状态、中间激活值,每一块都在占显存。显存一旦爆掉,要么降低 batch size,要么用更短的序列长度,这两种操作都会直接拉低 token 吞吐量。
看一组最朴素的对比:7B 模型全参数训练,模型权重本身 FP32 要占 28GB,FP16 只要 14GB。当你只有一张 80GB 的卡,省下来这 14GB 意味着可以塞进更大的 batch、更长的序列,或者为激活值留出更多空间,这些都是实打实的 token 吞吐量。更关键的是,显存减半之后,梯度同步和参数更新的通信量也会跟着下降,这在多卡训练里收益会被进一步放大。
1.3 一张卡上的精度模式差异到底有多大
我给了一张 A100 80GB 上不同精度模式的对照表,这是理解后面所有操作的基础:
| 精度模式 | 7B 模型权重显存 | 主要算力规格 | 特点 |
|---|---|---|---|
| FP32 | 28GB | 19.5 TFLOPS | 动态范围和精度都最稳,但最慢 |
| TF32 | 28GB(权重仍存 32 位) | 156 TFLOPS(Tensor Core) | 只加速计算,不省显存,精度接近 FP32 |
| FP16 | 14GB | 312 TFLOPS(Tensor Core) | 计算和显存双降,需要额外维护数值稳定 |
很多人会把 TF32 误解成"半精度",其实完全不同。TF32 是 Tensor Core 上的一个计算模式,权重在内存里还是 32 位存储,只是做矩阵乘的时候把尾数截断,用硬件加速;FP16 则是从存储到计算全都减半。搞清楚这个区别,你才能知道自己到底该开哪个开关。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. FP16和TF32的底层原理:不是简单"砍一半精度"那么粗暴
我见过太多人把 FP16 理解成"数字变小了所以省内存",这个说法只对了一半。要想用好混合精度,至少得知道这两种格式到底动了哪几位。
2.1 FP16的动态范围和精度:为什么训练时梯度容易"消失"
浮点数由三部分组成:1 位符号位、指数位和尾数位。FP32 是 1 位符号 + 8 位指数 + 23 位尾数;FP16 是 1 位符号 + 5 位指数 + 10 位尾数。指数位从 8 位变成 5 位,意味着能表示的最大值从大约 3.4e38 掉到了 65504,最小正常值从 1.18e-38 涨到了 6.1e-5。
问题就出在这。深度学习中很多梯度值远小于 1,比如 1e-6、1e-7 这种量级。在 FP16 里,小于 6.1e-5 的数值会进入非规格化区间,能表示到的最小值大约是 5.96e-8,再往下直接变 0。一旦梯度在反向传播的过程中下溢成 0,底层那些参数就彻底不更新了,模型表现为训练到一半 loss 不再下降。
这不是玄学,我在跑 LoRA 微调时就遇到过,某个低秩矩阵的梯度一直很小,FP16 训练下几千步之后某些参数彻底冻结。FP32 和 TF32 因为动态范围和 FP32 一致,不会出现这种问题。
2.2 TF32:一个非常"鸡贼"的中间方案
TF32 的设计思路很巧妙,它本质上是从 FP32 里截断出来的 19 位格式:1 位符号 + 8 位指数 + 10 位尾数。注意关键差异,它的指数位和 FP32 完全一致,动态范围没有任何损失;只是把 23 位尾数截成了 10 位,精度略降。
这带来的直接好处是:你在改代码的时候几乎不用操心数值溢出问题,不需要做 Loss Scaling,也不需要担心梯度下溢。它的代价是加速比没有 FP16 那么夸张,而且因为权重存储还是 32 位,显存也省不下来。所以在一些对精度敏感、但不太缺显存的场景,TF32 反而是比 FP16 更省心的选择。
2.3 Tensor Core:所有加速都发生在硬件级别
FP16 和 TF32 之所以能跑得比 FP32 快那么多,核心在 NVIDIA 从 Volta 架构开始引入的 Tensor Core。Tensor Core 不是简单地把 CUDA 核心数字变多,它是独立的矩阵乘单元,一次指令可以算 4x4 甚至更大的矩阵乘法块,吞吐量远超普通 CUDA 核心。
打个比方:普通 CUDA 核心是一辆小货车,每趟拉一件货;Tensor Core 是一辆重卡,直接把一批货打包运走。FP16 比 TF32 在 A100 上快一倍,就是因为 FP16 数据宽度只有 TF32 的一半,同一个 Tensor Core 单位时间能处理更多元素。
这也是为什么代码层面不需要显式调 CUDA 算子,只要精度模式和数据类型选对了,PyTorch 底层的 cuBLAS 会自动把矩阵乘分发到 Tensor Core 上。
3. PyTorch环境下的实战配置:从一行flag到全流程接入
原理说清楚了,下面直接上实操。我用 PyTorch 来演示,这是目前训练 LLM 最主流的框架。
3.1 开启TF32:两行代码的事,但很多人不知道
TF32 在 PyTorch 中默认是关闭的,因为从 1.7 开始官方担心尾数截断会带来精度损失,就默认改为 FP32。开启方式非常简单:
python复制# 全局开启 TF32
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
第一个开关控制的是矩阵乘法,第二个控制的是 cuDNN 里的卷积等算子。如果你只开 matmul 不开 cudnn,一些卷积层还是跑在普通 FP32 下,收益会打折扣。我在微调视觉-语言模型时就踩过这个坑,改了一行代码速度上去了,但后来对比 profile 发现卷积部分一直没被加速。
但要注意,对于大模型训练,如果使用了 autocast 且算子走到了 FP16 路径,TF32 开关实际上不会生效,因为数据类型已经明确变成了 FP16。TF32 最典型的应用场景是:你想保留 FP32 的训练管线(比如不想动优化器状态、不想做 Loss Scaling),只想白嫖 Tensor Core 的计算加速。
3.2 AMP混合精度训练:autocast 与 GradScaler
FP16 训练的标准做法是使用自动混合精度,PyTorch 提供的两个核心组件是 torch.autocast 和 GradScaler。
python复制import torch
model = MyLLM().cuda()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
scaler = torch.amp.GradScaler("cuda")
for batch in dataloader:
optimizer.zero_grad()
with torch.autocast(device_type="cuda", dtype=torch.float16):
output = model(batch["input_ids"], attention_mask=batch["attention_mask"])
loss = criterion(output, batch["labels"])
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
autocast 做的事是自动把模型前向传播里适合低精度的算子(比如线性层、矩阵乘法、卷积)切到 FP16,把不适合低精度的算子(比如 LayerNorm、Softmax、交叉熵)留在 FP32。GradScaler 则是负责在反向传播前把 loss 放大,防止梯度下溢,再在更新参数之前把梯度缩小回去。
抄作业的时候有一个很容易漏掉的点:model 本身的参数可以保持 FP32,不需要手动 .half()。因为 autocast 只影响算子的计算精度,参数存储还是 FP32,梯度更新时数值稳定得多。手动把整个 model 转成 FP16 反而容易出问题。
3.3 分布式训练和FSDP下的精度配置
如果你在用 DDP 或者 FSDP 做多卡训练,混合精度的收益会更明显,但配置也更讲究。FSDP 下开启混合精度的标准做法:
python复制from torch.distributed.fsdp import MixedPrecision
bf16_policy = MixedPrecision(
param_dtype=torch.bfloat16,
reduce_dtype=torch.bfloat16,
buffer_dtype=torch.bfloat16,
)
这里我直接写了 bf16,因为大模型分布式训练里 bf16 其实是比 FP16 更稳的选择——它的指数位和 FP32 一样,动态范围够大,不存在溢出问题,精度只靠多出来的尾数弥补。如果你的卡是 A100 及以上的架构,我真心建议优先考虑 BF16 而不是 FP16。FP16 的优势主要是生态最成熟、推理框架支持最广,而且很多量化工具链围绕它做优化。
如果你是 A100 之前的卡(V100、T4),不支持 BF16,那就老老实实走 FP16 + GradScaler。
3.4 推理场景:一样能靠精度省钱
混合精度不只是训练阶段的事。推理时开启 TF32 同样能提速,特别是纯 FP32 的推理管线,两行代码就能白拿接近 1 倍的加速。而 FP16 推理更关键的价值在于 KV Cache:每个 token 的 KV 如果用 FP16 存储,相比 FP32 直接把缓存占用减半,同样的显存能承载更长的上下文和更大的并发量。
在 vLLM 这类推理框架里,默认就是用 FP16 或者 BF16 来跑,dtype=auto 会自动读取模型权重的存储精度。对线上服务来说,KV Cache 从 FP32 切到 FP16,往往比加张卡还管用,因为上下文长度和并发数是直接决定 token 成本的关键指标。
4. 训练不稳才是最大的坑:溢出、NaN和精度抖动的完整排查链路
到了实战阶段,最折磨人的永远不是代码跑不起来,而是代码跑起来了但 loss 曲线突然崩溃。我把自己遇到过的混合精度训练事故和完整排查思路写在下面,这一节建议收藏。
4.1 一次典型NaN事故的完整定位过程
有一次我做 13B 模型的 LoRA 微调,FP16 AMP 开起来之后,前 500 步一切正常,loss 稳步下降,500 步之后突然开始震荡,到了 600 步直接变成 NaN。
第一次遇到这种事,千万别先去改学习率,那是玄学。我的排查路径是这样的:
- 先看 GradScaler 的 scale 值。日志里
LossScaler从 65536 开始不断翻倍,最后溢出变成 inf。这说明梯度在放大后溢出,scaler.step(optimizer)会跳过参数更新,导致 loss 卡住或跳变; - 打开梯度范数的监控。发现某个 transformer 层的梯度范数明显比其他层大几个数量级,说明梯度爆炸源头是在那一层附近;
- 检查那一层的实现,发现是我手写的一个自定义 RoPE 旋转位置编码模块,里面用了一个自定义 CUDA kernel。这个 kernel 没有注册到 autocast 的白名单里,导致它接收了 FP16 输入又输出 FP32 结果,中间某个中间变量直接溢出;
- 修复方式:给那个模块加上
@torch.cuda.amp.custom_fwd(cast_inputs=torch.float32),强制它用 FP32 计算;同时调整 GradScaler 的init_scale,从默认的 65536 降到 4096,稳住梯度放大幅度。
修复后重新训练,loss 曲线恢复正常。
4.2 Loss Scaling 的本质:把梯度放大到"看得见"的范围
FP16 的下溢问题不是概率事件,是必然事件。深度学习里的梯度分布通常非常不均匀,少量参数梯度大,大量参数梯度小。小的那部分在 FP16 里约等于 0。
GradScaler 干的事就是在反向传播计算 loss 梯度之前,先把 loss 乘上一个大系数(比如 65536),让梯度整体变大,进入 FP16 的表示范围,反向传播完成之后再除以这个系数,恢复出真实的梯度值。
动态调整机制也不复杂:如果连续 N 步都没有出现 inf/NaN 梯度,scale 翻倍;一旦出现 inf/NaN,就跳过这一步的参数更新,并把 scale 减半。这也是为什么 GradScaler 处于高 scale 状态时,某个意外溢出会导致 loss 毛刺——那一整步直接废掉了。
4.3 哪些模块必须留在FP32
用 autocast 的时候,PyTorch 已经自动把数值敏感的算子留在 FP32 了,但有几个地方容易踩雷:
- LayerNorm 和 RMSNorm:虽然 PyTorch 的 autocast 会把它们留在 FP32,但如果你用了某些第三方实现,或者自己手写了 normalization,一定要手动确认;
- Softmax 和交叉熵损失:softmax 的指数运算在 FP16 下非常容易溢出,因为 exp(x) 在 x 很大时会快速涨到 65504 以上;
- Embedding 层:有些实现会希望 embedding 梯度在 FP32 下累积,防止词向量更新太慢。Llama 系列模型经常配置
lm_head不使用混合精度; - 自定义的 loss 函数:如果你自己写了 loss,建议在里面加
torch.cuda.amp.autocast(enabled=False)来强制为 FP32。
4.4 快速验证精度的烟雾测试
改完精度配置后,我不建议直接跑完整训练集,而是做一个快速的烟雾测试:用 8 到 16 条样本,让模型过拟合到 100% 准确率。如果模型连小样本都过拟合不了,不是精度配置有问题就是代码有问题,这种问题越早暴露越好。
另外一个我常用的验证手段是对照实验:固定随机种子,分别在 FP32 和混合精度下跑同一个 batch,对比最后一层输出的 logits 余弦相似度。正常情况下余弦相似度应该在 0.99 以上。如果掉到 0.9 以下,说明某个算子在低精度下的误差被放大了,需要定位是哪一层的问题。同理,梯度范数的分布图如果出现严重的长尾,也说明缩放策略可能需要调整。
5. 收益实测:多个模型上的对比数据与取舍判断
理论再漂亮也是纸上谈兵。下面是我的实测数据,全部来自过去小半年在不同团队和机器上的真实训练任务,硬件以 A100 80GB 为主。
5.1 测试环境与实验设计
我选了三个有代表性的模型:
- 7B 参数量 decoder-only LLM,LoRA 微调;
- 13B 参数量 LLM,全参数继续预训练;
- 11B 多模态模型(视觉编码器 + LLM),全参数微调。
统一控制 batch size 尽量让显存不成为瓶颈(FP32 和 TF32 用同一 batch size;FP16 因为显存减半,可以增大 batch,但为了对比公平,下面的数据还是统一 batch size,显存占用数据单独列出)。
5.2 核心对比数据
| 模型 | 精度模式 | 吞吐量(token/s) | 显存占用峰值 | loss 收敛行为 | 稳定性 |
|---|---|---|---|---|---|
| 7B LoRA | FP32 | 1210 | 52GB | 基准 | 稳定 |
| 7B LoRA | TF32 | 4620 | 52GB | 与 FP32 基本一致 | 稳定 |
| 7B LoRA | FP16 AMP | 8350 | 27GB | 收敛略快 | 需 Loss Scaling |
| 13B 全参 | FP32 | 386 | 74GB | 基准 | 稳定 |
| 13B 全参 | FP16 AMP | 2480 | 40GB | 收敛步数略多但单步快 | 需监控 NaN |
| 11B 多模态 | TF32 | 1810 | 63GB | 与 FP32 接近 | 稳定 |
| 11B 多模态 | FP16 AMP | 3370 | 34GB | 收敛略慢 | 视觉塔需额外注意 |
说明一下,这里的吞吐量受数据加载、模型并行策略和磁盘 IO 影响很大,看相对差异比看绝对值更有意义。我自己的结论是:
- TF32 在 7B LoRA 上收益非常明显,四倍左右的吞吐提升,且完全不需要改优化器和 Loss Scaling 逻辑,基本零成本接入。这是因为 LoRA 训练时梯度本身不算极端,TF32 的动态范围和 FP32 相同,不会出现下溢;
- FP16 AMP 的显存节省是实实在在的,对显存敏感任务的收益最有价值;
- 13B 全参数训练上,FP16 的收益和稳定性成反比。吞吐快了很多,但需要更多的 NaN 监控和 checkpoint 回滚机制。如果不具备这些基础设施,宁可先上 TF32。
5.3 什么时候不应该用FP16/TF32
不是所有场景都适合混合精度。我遇到过几个反例,供你排雷:
- 小规模模型的微调(比如几千万参数的小模型),FP32 本身就很快,混合精度的收益不明显,反而可能因为精度问题引入额外调参成本;
- 对数值精度极其敏感的科研实验,比如某些需要复现论文精确数值的任务,TF32 和 FP32 的尾数截断会带来差异;
- 如果你的代码里有很多自定义 CUDA 算子,又没有为它们配置 FP16 内核,混合精度反而会因为频繁的格式转换变慢。这种情况下先补齐算子精度再说。
5.4 容易被忽略的推理端收益:KV Cache 的 token 预算
训练端之外还有一个很容易被忽略的场景:推理服务的 token 预算。假设你的 GPU 有 80GB 显存,KV Cache 用 FP32 存储和用 FP16 存储,差别是整个服务能同时承载的上下文 token 数量。
上下文越长,占用的 KV Cache 越大。把 KV Cache 从 FP32 换成 FP16,相当于把单请求的显存开销砍半,在线服务能支撑的并发和上下文长度直接翻倍,token 拒绝率下降,单位 token 服务成本也跟着下降。
6. 组合拳进阶:混合精度之外还能怎么压榨token和算力
混合精度是性价比最高的一步,但它的收益是乘法而不是加法。想进一步压缩成本,要把下面几个方向和它叠加使用。
6.1 消除填充token:让每个计算都落在真实数据上
大部分训练框架处理变长序列时会做 padding,把短样本补齐到 batch 内最长序列的长度。这些 padding token 白白消耗计算量,因为模型照样会对它们做 attention。
我实测过一个 13B 模型的数据集,平均序列长度 800,最长 2048,padding 率超过 60%。这意味着整整六成的算力烧在无意义的 token 上。解决方法有两个:
- 按长度排序的 batching,俗称 bucket,让每个 batch 内序列长度尽量接近,减小 padding 浪费;
- 序列打包,把多个短样本拼成一个长序列,并在 attention mask 里隔离不同样本。这个方案对代码侵入性大,但收益非常可观。
配合混合精度之后,训练速度还能再翻一倍,这是最容易忽略的"免费午餐"。
6.2 梯度检查点:用算力换显存,再换token吞吐
混合精度把显存省下来之后,梯度检查点还能再省一波:在反向传播时重新计算前向的中间激活值,而不是全部保存在显存里。它的代价是增加大约 30% 的计算量。
这个交换值不值?我的判断标准是:如果你的显存因为激活值爆掉导致 batch size 被迫减半,而 batch size 减半又让 GPU 利用率掉到 50% 以下,那梯度检查点绝对是划算的。很多时候你以为自己卡在算力上,实际上卡在显存上,batch size 上不去,吞吐量就上不去。
6.3 优化器状态的低精度化
Adam 优化器要维护一阶动量 m 和二阶动量 v,每个参数要多占 8 字节甚至更多。对于一个 70B 模型,优化器状态本身就要几百 GB 显存。
bitsandbytes 提供了 8bit Adam,能将优化器状态压缩到原来的 1/4 左右。配合 FP16 权重和梯度,整个训练管线的显存占用能降到原来 FP32 训练的 1/3 到 1/4。不要一听到 8bit 就觉得会损失精度,实测下来 8bit Adam 的收敛行为与 32bit Adam 非常接近,但对部分模型反而需要调一下 beta2 参数。
6.4 数据与调度层面的token压缩
最后说一个和精度无关、但直接影响 token 消耗的方向:数据质量。我在实际项目中见过太多重复训练数据,模型反复看相同内容的 token,算力大量浪费。在训练之前做一轮数据去重、去噪声,比任何硬件层面的优化都能更快降低 token 消耗,却基本没人主动做。
另一个思路是在长上下文场景利用 token 级注意力剪枝,把不重要的 token 在进入 attention 之前丢弃,或者对历史 token 做压缩。这些技术还在快速演进,但在已经跑通混合精度、梯度检查点和序列打包之后,再研究这些高级手段也不迟。
根据我个人这段时间的经验,混合精度训练和推理其实没有想象中那么玄乎。FP16 和 TF32 的选型核心就是三件事:你缺不缺显存,你怕不怕溢出,你愿不愿意为了稳定调试付出时间。如果你不想在数值稳定上花太多心思,TF32 是零风险起步;如果你需要显存红利来撑起更大的 batch 和更长的序列,FP16 + GradScaler 是必经之路。最后再分享一个小技巧:上线混合精度之后,一定记得在训练脚本里把每个 step 的吞吐量打出来,token/s 这个数字是检验所有优化是否有效的最直观指标,它不会说谎。
