在深度学习模型规模不断膨胀的今天,训练一个大型模型动辄需要数十甚至上百 GPU·天。当显存成为瓶颈、训练时间成为成本大头时,混合精度训练(Mixed Precision Training)就成了必须掌握的优化手段。它可以在几乎不损失模型精度的前提下,把显存占用降低 30%–50%,并把训练吞吐量提升 1.5–3 倍。

本文面向已经熟悉 TensorFlow 2.x 基础 API、希望在生产环境榨干 GPU 性能的工程师。我们会从 NVIDIA Tensor Core 的硬件原理讲起,再到
1 | tf.keras.mixed_precision |
的两种策略、损失缩放(Loss Scaling)的必要性、自定义训练循环的写法、调试技巧和踩坑案例,最后给出一份可直接套用的生产配置清单。
一、为什么混合精度能提速:从 Tensor Core 说起
混合精度训练的核心思想是:在保证数值精度的关键步骤(参数更新、累加梯度)使用 FP32,而在矩阵乘法这种计算密集、对精度容忍度高的步骤使用 FP16 或 BF16。NVIDIA Volta 及之后的 GPU(V100、A100、H100 等)内置了 Tensor Core,专门为低精度矩阵运算做加速。
简而言之,一个 Tensor Core 在一个时钟周期内可以完成一个 4×4×4 的 FP16 矩阵乘法并累加到 FP32 上,而传统 CUDA Core 需要多个周期。这带来两个直接收益:
- 计算吞吐提升:FP16 矩阵乘法的 FLOPS 通常是 FP32 的 2–8 倍(H100 上甚至支持 FP8)。
- 显存占用降低:FP16 占 2 字节,FP32 占 4 字节。权重、激活、梯度都减半,显存带宽压力同步下降。
需要注意的是,并不是所有 GPU 都受益。Pascal 架构(GTX 10 系列)虽然支持 FP16 计算,但没有 Tensor Core,提速有限甚至更慢;Maxwell 及更早架构完全不支持。下表是主流 GPU 的支持情况:
| 架构 | 代表型号 | Tensor Core | 推荐策略 |
|---|---|---|---|
| Volta | V100 | 是(FP16) | mixed_float16 |
| Turing | RTX 20 系列 | 是(FP16) | mixed_float16 |
| Ampere | A100、RTX 30 系列 | 是(FP16/BF16) | mixed_bfloat16 |
| Hopper | H100 | 是(FP16/BF16/FP8) | mixed_bfloat16 |
| Pascal 及更早 | GTX 10 系列 | 否 | 不推荐 |
二、mixed_float16 与 mixed_bfloat16:两种策略怎么选
TensorFlow 通过
1 | tf.keras.mixed_precision.Policy |
统一管理精度策略。开启全局策略后,所有之后创建的 Keras 层都会按策略自动选择 dtype:
1
2
3
4
5
6
7 from tensorflow.keras.mixed_precision import set_global_policy
# 方式一:FP16 混合精度
set_global_policy('mixed_float16')
# 方式二:BF16 混合精度(Ampere 及之后推荐)
set_global_policy('mixed_bfloat16')
两者的关键区别在于 动态范围,而非精度位数:
| 属性 | FP16 (float16) | BF16 (bfloat16) |
|---|---|---|
| 指数位 | 5 位 | 8 位(同 FP32) |
| 尾数位 | 10 位 | 7 位 |
| 最大值 | ≈65504 | ≈3.4e38(同 FP32) |
| 最小正数 | ≈5.96e-8 | ≈1.18e-38 |
| 需要损失缩放 | 是 | 否 |
结论:如果你的卡支持 BF16(A100 / H100 / 30 系以上),优先选
1 | mixed_bfloat16 |
,省去损失缩放的麻烦;V100 / 20 系只能用
1 | mixed_float16 |
,并且必须配合损失缩放。
三、损失缩放:FP16 训练为什么必须
FP16 的最小正规格化数约为 6e-5,而深度网络的反向梯度常常只有 1e-4 甚至更小。如果不处理,大量梯度会下溢为 0,导致参数不更新、训练停滞。解决方案是损失缩放:在前向计算后把 loss 乘以一个大常数 S(典型值 2^15=32768),梯度随之放大;反向结束后再用梯度除以 S 还原。
TensorFlow 提供了两种缩放器:
- 动态损失缩放(
1LossScaleOptimizer
默认):自动试探合适的 S,遇到 NaN/Inf 自动减半,稳定后逐步加倍。生产环境默认选这个。
- 固定损失缩放:手动指定
1dynamic=False, initial_scale=1024
。适用于已知梯度分布稳定的场景,省去试探开销。
开启
1 | mixed_float16 |
后,必须用
1 | LossScaleOptimizer |
包裹原优化器,否则训练大概率发散:
1
2
3
4
5
6
7 from tensorflow.keras.mixed_precision import LossScaleOptimizer
optimizer = tf.keras.optimizers.Adam(learning_rate=1e-3)
optimizer = LossScaleOptimizer(optimizer) # 动态缩放
# 固定缩放(可选)
# optimizer = LossScaleOptimizer(optimizer, dynamic=False, initial_scale=1024)
注意:
1 | mixed_bfloat16 |
因为动态范围和 FP32 一致,不需要损失缩放,直接用原优化器即可。这是 BF16 在易用性上的最大优势。
四、Keras Model.fit 中的混合精度实战
最简单的用法是在
1 | model.compile |
前开启全局策略,其余代码完全不变。下面是一个完整的图像分类训练示例:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36 import tensorflow as tf
from tensorflow.keras.mixed_precision import set_global_policy, LossScaleOptimizer
set_global_policy('mixed_float16')
# 数据管线(省略详细预处理)
(ds_train, ds_test), info = tfds.load('cifar10', split=['train', 'test'], as_supervised=True, with_info=True)
AUTOTUNE = tf.data.AUTOTUNE
def preprocess(image, label):
image = tf.cast(image, tf.float32) / 255.0
image = tf.image.resize(image, (32, 32))
return image, label
ds_train = ds_train.map(preprocess).shuffle(2048).batch(256).prefetch(AUTOTUNE)
ds_test = ds_test.map(preprocess).batch(256).prefetch(AUTOTUNE)
# 模型:最后一层保持 FP32 保证数值稳定
inputs = tf.keras.Input(shape=(32, 32, 3))
x = tf.keras.layers.Conv2D(64, 3, padding='same', activation='relu')(inputs)
x = tf.keras.layers.Conv2D(128, 3, padding='same', activation='relu')(x)
x = tf.keras.layers.GlobalAveragePooling2D()(x)
x = tf.keras.layers.Dense(256, activation='relu')(x)
# 关键:输出层强制 dtype='float32',避免 softmax 在 FP16 下溢出
outputs = tf.keras.layers.Dense(10, activation='softmax', dtype='float32')(x)
model = tf.keras.Model(inputs, outputs)
optimizer = tf.keras.optimizers.Adam(1e-3)
optimizer = LossScaleOptimizer(optimizer)
model.compile(optimizer=optimizer,
loss='sparse_categorical_crossentropy',
metrics=['accuracy'])
model.fit(ds_train, epochs=30, validation_data=ds_test)
这里有两个必须注意的细节:
- 输出层强制
:softmax、logsoftmax 这类涉及指数运算的层,在 FP16 下极易溢出。即使开了全局 mixed_float16,也要在最后这层显式覆盖 dtype。1dtype='float32'
- 输入数据保持 FP32:数据进入网络后再由 Keras 自动转换。不要在
1preprocess
里提前 cast 成 FP16,否则会丢失精度。
五、自定义训练循环写法
对于需要自定义训练循环的场景(梯度累积、多 GPU 手动同步、复杂多任务),混合精度的写法略有不同。关键是
1 | tf.GradientTape |
内部要正确处理缩放后的梯度:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23 set_global_policy('mixed_float16')
optimizer = LossScaleOptimizer(tf.keras.optimizers.Adam(1e-3))
@tf.function
def train_step(inputs, labels):
with tf.GradientTape() as tape:
predictions = model(inputs, training=True)
# loss 在 FP16 下计算,但要 cast 到 FP32 做缩放
loss = tf.keras.losses.sparse_categorical_crossentropy(labels, predictions)
loss = tf.cast(loss, tf.float32)
scaled_loss = optimizer.get_scaled_loss(loss)
# 计算的是缩放后的梯度
scaled_grads = tape.gradient(scaled_loss, model.trainable_variables)
# 用 get_unscaled_gradients 还原
grads = optimizer.get_unscaled_gradients(scaled_grads)
optimizer.apply_gradients(zip(grads, model.trainable_variables))
return loss
for epoch in range(30):
for x, y in ds_train:
loss = train_step(x, y)
print(f'epoch {epoch} loss={loss.numpy().mean():.4f}')
两对配套 API 一定要成对使用:
1 | get_scaled_loss |
↔
1 | get_unscaled_gradients |
。如果忘了 unscale,梯度会比真实值大几千倍,优化器会瞬间把权重吹爆。
六、性能数据:实测能快多少
下面是在 ResNet-50 + ImageNet 子集上的实测对比(V100 32GB,batch_size 已调到显存上限):
| 配置 | batch_size | step 时间 | 显存峰值 | Top-1 精度 |
|---|---|---|---|---|
| FP32 基线 | 128 | 245ms | 28 GB | 76.2% |
| mixed_float16 | 256 | 132ms | 19 GB | 76.3% |
| mixed_bfloat16 (A100) | 512 | 78ms | 22 GB | 76.1% |

可以看到几个趋势:
- batch_size 翻倍:显存减少直接体现在能塞进更大的 batch,配合更大的学习率进一步加速收敛。
- step 时间接近减半:Tensor Core 在大 batch 下利用率更高,提升更明显。
- 精度几乎无损:动态损失缩放 + 输出层 FP32 保证数值稳定,精度波动在 0.2% 以内属正常范围。
七、常见踩坑与排查清单
混合精度虽然 API 简单,但调试起来有几个高频坑点。按出现频率从高到低列出排查清单:
1. 训练一开始就出现 NaN
最常见的原因是忘了用 LossScaleOptimizer,或者忘了对输出层强制 FP32。检查方式:在
1 | train_step |
末尾加
1 | tf.debugging.check_numerics(loss, 'loss has nan') |
,定位第一个出现 NaN 的层。
2. 损失缩放一直在降不恢复
1 | LossScaleOptimizer |
遇到 Inf 会把 scale 减半。如果发现 scale 一直往下掉、loss 一直抖动,多半是学习率过大或数据中存在异常值(如未归一化的图像像素)。先在 FP32 下复现,确认数据无误后再切回混合精度。
3. 验证集指标异常
推理时如果用了
1 | mixed_float16 |
策略但没 cast 输入,会出现精度抖动。推荐做法:训练用混合精度,推理用 FP32,或者用
1 | tf.keras.models.clone_model |
+
1 | set_weights |
复制一份 FP32 模型专门做 eval。
4. 多 GPU 训练损失缩放不同步
1 | MirroredStrategy |
下每个副本各自算梯度,
1 | LossScaleOptimizer |
会自动 all-reduce 缩放后的梯度,再 unscale。如果你手写分布式训练,必须确保先 all-reduce 缩放梯度,再 unscale,顺序反了会让数值不稳定。
5. TFLite 部署精度掉点
训练完的混合精度模型导出 SavedModel 时仍是 FP32(Keras 默认存 FP32 权重)。但如果你手动保存了 FP16 权重,再用
1 | float16 |
量化转 TFLite,会出现“双重量化”问题。建议训练和导出保持单一精度路径,导出时统一用 FP32。
八、生产配置清单
把上述要点浓缩成一份可复制的生产配置,覆盖 Keras、自定义循环、分布式三种场景:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30 # 一、导入与策略(脚本最顶部)
import tensorflow as tf
from tensorflow.keras.mixed_precision import set_global_policy, LossScaleOptimizer
# 根据 GPU 架构选策略
GPU_ARCH = 'ampere' # volta / ampere / hopper
POLICY = 'mixed_bfloat16' if GPU_ARCH in ('ampere', 'hopper') else 'mixed_float16'
set_global_policy(POLICY)
# 二、模型定义:输出层强制 FP32
outputs = tf.keras.layers.Dense(num_classes, dtype='float32', name='logits')(x)
# 三、优化器:FP16 才需要包裹
base_opt = tf.keras.optimizers.Adam(learning_rate=1e-3 * batch_size / 256)
optimizer = LossScaleOptimizer(base_opt) if POLICY == 'mixed_float16' else base_opt
# 四、分布式策略内使用
strategy = tf.distribute.MirroredStrategy()
with strategy.scope():
model = build_model()
model.compile(optimizer=optimizer, loss=loss_fn, metrics=['accuracy'])
# 五、训练
model.fit(train_ds, epochs=epochs, validation_data=val_ds,
callbacks=[tf.keras.callbacks.ReduceLROnPlateau(patience=3)])
# 六、推理时切回 FP32(可选,精度更稳)
set_global_policy('float32')
eval_model = tf.keras.models.clone_model(model)
eval_model.set_weights(model.get_weights())
总结
混合精度训练是当前深度学习工程化中性价比最高的优化之一:几行代码就能换来近一倍的训练速度和显著下降的显存占用。掌握要点其实就四条:
- 支持 BF16 的卡优先选
1mixed_bfloat16
,省心;只能 FP16 的卡用
1mixed_float16必须配
1LossScaleOptimizer。
- 输出层(softmax / log_softmax)强制
1dtype='float32'
,避免数值溢出。
- 自定义训练循环中
1get_scaled_loss
与
1get_unscaled_gradients必须成对使用。
- 训练用混合精度、推理切回 FP32,规避部署时的精度抖动。
把这些细节吃透,下次再遇到训练慢、显存不够的场景,你就能从容地把它提速到极限。建议在迁移到新硬件或新模型时,先跑一次混合精度 vs FP32 的对照实验,用真实数据校准 batch_size 和学习率,再投入大规模训练——这往往比无脑堆 GPU 更划算。
汤不热吧