梯度爆炸典型现象是loss突变为inf/nan或权重绝对值>1e6;快速确认需在反向传播后、优化器更新前检查grads是否含inf/nan,或用tf.debugging.check_numerics定位异常中间张量。

梯度爆炸的典型现象和快速确认方法
训练时 loss 突然变成 inf 或 nan,或者某轮迭代后 model.trainable_variables 中某些权重突然变得极大(如绝对值 > 1e6),基本可以判定是梯度爆炸。更隐蔽的情况是训练 loss 不下降、验证指标震荡剧烈,也值得怀疑。
最直接的排查方式是在训练循环中加一句检查:
if tf.math.is_inf(grads).numpy().any() or tf.math.is_nan(grads).numpy().any():
print("Gradient contains inf/nan!")
注意:必须在 tf.GradientTape 记录后、optimizer.apply_gradients() 前检查 grads,否则 tape 已释放,无法访问。
用 tf.clip_by_global_norm 快速缓解
这不是根治方案,但能立刻让训练跑下去、帮你定位问题模块。它对所有梯度张量做全局裁剪,避免单个大梯度拖垮整个更新步长。
立即学习“Python免费学习笔记(深入)”;
实操建议:
- 在
optimizer.apply_gradients()前插入:grads, _ = tf.clip_by_global_norm(grads, clip_norm=1.0) -
clip_norm从1.0开始试;如果训练仍不稳定,可降到0.5;若 loss 下降太慢,再尝试2.0 - 不要长期依赖裁剪——它掩盖了模型/数据/初始化的问题,只用于诊断过渡
注意:tf.clip_by_global_norm 对 RNN 类结构(如 LSTM)尤其有效,因为其时间展开易累积梯度。
排查初始化、激活函数和损失函数三处关键点
梯度爆炸常源于前向传播阶段就埋下的数值不稳定性,而非反向传播本身。
常见出问题的地方:
-
kernel_initializer用"glorot_uniform"或"he_normal",避免全零或过大的初始权重;RNN 的recurrent_initializer尤其要设为"orthogonal" - 深层网络慎用
tf.nn.relu后接大尺度输入——考虑换tf.nn.leaky_relu或加BatchNormalization层 - 分类任务用
tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),确保 logits 不经 softmax 再算 loss,否则softmax + log在极端 logits 下会数值溢出
一个容易被忽略的坑:from_logits=False(默认)搭配未归一化的输出,会导致 loss 计算中出现 log(0) → -inf → 反向传播出 nan 梯度。
用 tf.debugging.check_numerics 定位具体张量
当裁剪后仍有 nan,说明问题不在梯度大小,而在计算过程中的中间结果溢出(比如某层输出爆掉)。
在模型前向调用中插入检查:
x = layer(x)
tf.debugging.check_numerics(x, message=f"NaN/Inf in {layer.name}")
注意几点:
- 该操作仅在 eager mode 下生效,Graph mode 需配合
@tf.function+tf.config.run_functions_eagerly(False)调试时临时关闭图执行 - 不要全网铺开——先聚焦在 RNN 输出、最后几层 Dense、或自定义 loss 计算处
- 报错信息里的
op name很关键,比如MatMul出问题,大概率是输入矩阵含inf;Softmax出问题,则看上游 logits 是否过大
真正难缠的梯度爆炸,往往卡在某个看似无害的 tf.math.exp 或 tf.pow 操作里——这些地方不会被自动裁剪,必须靠 check_numerics 抓出来。


















