根本原因是数值不稳定:除零、log(0)、sqrt(负数)、exp溢出或梯度爆炸;nan会通过加法、乘法快速污染整个计算图,常见于数据含inf/nan、logits未裁剪、batch size=1时batchnorm除零等场景。

为什么TensorFlow训练中会出现loss为NaN?
loss变成NaN不是偶然,而是模型在某个计算环节遭遇了数学上未定义的操作。最常见的是:除零、对负数取对数、log(0)、exp爆炸后溢出、梯度爆炸导致权重突变为inf再参与后续运算。一旦出现一个NaN,它会像病毒一样通过加法、乘法迅速污染整个计算图。
-
tf.nn.softmax_cross_entropy_with_logits输入未归一化logits时容易因exp过大溢出 - 使用
tf.keras.losses.CategoricalCrossentropy(from_logits=False)但传入了logits(未经过softmax) - 数据含
NaN或inf(比如读取损坏的图像、缺失值未清洗) - 学习率过大,导致某次参数更新后权重剧烈震荡,激活值发散
如何快速定位NaN出现在哪一层或哪一步?
TensorFlow提供tf.add_check_numerics_ops()和tf.debugging.enable_check_numerics()(TF 2.10+),但更实用的是在训练循环中插入检查点:
with tf.GradientTape() as tape:
predictions = model(x, training=True)
loss = loss_fn(y, predictions)
<h1>检查loss是否合法</h1><p>if tf.math.reduce_any(tf.math.is_nan(loss)) or tf.math.reduce_any(tf.math.is_inf(loss)):
print("Loss is NaN/Inf at batch:", i)</p><h1>打印前几项预测值和标签,辅助判断</h1><pre class="brush:python;toolbar:false;">print("pred[:3]:", predictions[:3].numpy())
print("y_true[:3]:", y[:3].numpy())
raise RuntimeError("NaN loss detected")gradients = tape.gradient(loss, model.trainable_variables)
检查梯度
if any(tf.math.reduce_any(tf.math.is_nan(g)) for g in gradients if g is not None): print("NaN gradient found")
- 不要只检查
loss标量,也要检查predictions和中间层输出(如用model.layers[i].output) - 若使用
tf.data.Dataset,在map函数末尾加tf.debugging.check_numerics可捕获数据预处理阶段的问题 -
tf.print比print更安全,能在图模式下输出调试信息
哪些模型结构或损失函数配置容易触发NaN?
某些组合看似合理,实则暗藏风险:
- 用
tf.keras.layers.Softmax+tf.keras.losses.CategoricalCrossentropy(from_logits=False)是安全的;但若误用from_logits=True,而输出层已带Softmax,会导致二次归一化,数值极小→log(0)→NaN - 自定义损失函数中手写
tf.log(y_pred + 1e-8)看似防错,但如果y_pred本身是NaN,加常数无效 - 使用
tf.keras.optimizers.Adam时,epsilon=1e-7(默认)在梯度极小时可能放大数值误差;可尝试设为1e-5缓解 - BatchNorm层在batch size=1时方差为0,导致除零;训练时确保
batch_size >= 4,或改用tf.keras.layers.LayerNormalization
训练前必须做的三件事
这不是“建议”,是踩过坑后确认有效的硬性操作:
-
数据清洗:在送入模型前,用
tf.debugging.check_numerics包装输入Dataset的最后一步,例如dataset.map(lambda x, y: (tf.debugging.check_numerics(x, "x"), tf.debugging.check_numerics(y, "y"))) -
初始化检查:构建模型后,用
model(x_sample)跑一次前向,立刻检查输出是否含NaN;不要等到训练十几轮才发现 -
学习率试探:从
1e-5开始试训2–3个step,确认loss下降且无NaN,再逐步放大;用tf.keras.callbacks.ReduceLROnPlateau不如先手动控住起点
复杂模型里NaN可能延迟数个step才暴露,所以早期验证不能省。真正棘手的情况是:loss正常,但某层梯度突然inf,接着下一轮loss就NaN——这种必须逐层检查梯度范数,而不是等loss报错。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











