只需一行代码 tf.keras.mixed_precision.set_global_policy('mixed_float16') 即可启用混合精度,但需满足 tensorflow ≥ 2.4、nvidia volta 及以上 gpu;该策略默认使层计算用 fp16、权重主副本保持 fp32,并自动处理损失缩放与梯度更新,须在模型构建前调用,且仅对 cuda 设备生效。

怎么一行代码启用混合精度?
只要 TensorFlow ≥ 2.4,GPU 是 NVIDIA Volta(如 V100)或更新架构(T4/A100/A800/H100),tf.keras.mixed_precision.set_global_policy('mixed_float16') 这一行就足够启动整套混合精度流程。
它不是“开关”,而是全局策略注册:后续所有 tf.keras.layers 构建的层,只要没显式指定 dtype,默认用 FP16 做前向/反向计算;权重主副本自动维护为 FP32;损失缩放、梯度类型转换、master weights 更新全由框架静默接管。
- 必须在模型构建前调用,否则已创建的层不会受策略影响
- 不兼容 CPU 训练——该策略仅对 CUDA 设备生效,CPU 上会静默回退为 FP32,且无警告
- 若使用自定义训练循环(非
model.fit()),需额外确保optimizer也适配,推荐直接用tf.keras.optimizers.Adam(...)等原生优化器,它们已内置 AMP 支持
为什么输出层和 Embedding 层必须手动设为 float32?
因为 mixed_float16 策略不会“智能判断”哪些层敏感——它只按规则转发 dtype,而 softmax 和词表查表这两类操作极易因 FP16 动态范围窄出问题:
-
softmax输入若含较大负值(如 -100),FP16 下exp(-100)直接下溢为 0,导致概率归一失败,输出全 NaN -
Embedding层本质是索引 + 查表,FP16 的舍入误差可能让index=12345.999被截断成12345.0,错位取错向量 - 所以必须显式声明:
tf.keras.layers.Dense(10, dtype='float32')、tf.keras.layers.Embedding(vocab_size, dim, dtype='float32')
训练突然爆 NaN?先检查这三处
混合精度最典型的失败现象不是变慢,而是某轮后 loss 突然跳成 nan 或 inf,接着梯度爆炸、权重发散。这不是代码写错了,而是数值链路某环脱节:
-
loss函数是否支持 FP16 输入?例如自定义 loss 若用了tf.math.log未加防零,FP16 下log(0)直接返回-inf,再乘缩放因子就崩了 - 数据预处理是否残留
NaN?FP32 下可能被忽略,但 FP16 对无效值更敏感,建议训练前用tf.debugging.check_numerics插桩检测输入张量 - 学习率是否还沿用 FP32 时的值?混合精度收敛更快,同样学习率容易震荡,尤其 Adam 的
epsilon默认是1e-7,在 FP16 下接近下限,建议微调为1e-4或交由tf.keras.optimizers.Adam(learning_rate=...)自动适配
显存没降一半?可能是 batch_size 没调够
混合精度理论显存减半,实际常只省 30–40%,因为显存大头未必是权重,而是中间激活值(activation)和梯度缓存。而这些恰恰依赖 batch size —— batch 越大,激活张量越胖,FP16 的节省才越明显。
实测中,把 batch size 从 64 提到 256,A100 上显存占用下降从 35% 升至 48%,训练吞吐提升 2.1×。但注意:
- 不能盲目堆 batch size,要同步调高 learning rate(线性缩放规则),否则收敛变差
-
tf.data.Dataset的prefetch和cache必须开启,否则 GPU 等待数据的时间占比上升,掩盖了计算加速收益 - 若用多卡
tf.distribute.MirroredStrategy,batch size 是每卡值,总 batch = 卡数 × 单卡 batch,别算错
真正卡住多数人的,从来不是策略怎么开,而是忘记输出层 dtype、或者以为开了策略就万事大吉,结果在 loss 函数里悄悄写了不兼容 FP16 的操作。










