TensorFlow 1中损失函数随机发散至NaN的问题求助
训练DNN时损失随机发散至NaN的问题排查与解决
问题背景
刚接触DNN,基于TensorFlow 1实现某论文基准模型(尝试TF2复现未达原性能)。训练中MSE损失初期正常下降,但会随机在若干轮次后发散至NaN;若使用发散前的最后权重重启训练,损失会继续下降直至再次发散。
模型存在特殊逻辑:需多次将网络输出重新计算为输入,循环迭代后再计算损失,简化代码如下:
input = ... for t in range(max_iter): x1 = tf.nn.relu(input@A1+b1) x1 = BatchNormalization()(x1) x2 = tf.nn.relu(x1@A2+b2) x2 = BatchNormalization()(x2) x3 = tf.nn.relu(x2@A3+b3) x3 = BatchNormalization()(x3) output = x3@A4+b4 input = recompute_input(output) compute_loss(input, true_values)
已尝试标准化输入、调整学习率、修改批次大小等方案,当前使用Adam Optimizer(学习率0.0001)仍存在随机发散问题。
解决建议
1. 检查并约束recompute_input的数值稳定性
该函数是循环迭代的核心,极可能引入极端值(如除法溢出、开根号取负、指数爆炸):
- 给除法操作添加小epsilon(如
1e-8)避免除以0; - 用
tf.clip_by_value限制函数输出范围,例如:input = tf.clip_by_value(recompute_input(output), -10.0, 10.0)
2. 调整BatchNorm的使用顺序
当前是ReLU后接BatchNorm,常规稳定做法是先BatchNorm再ReLU(输出层除外),可避免ReLU截断负数后导致BatchNorm统计量偏差,更好稳定数值分布:
x1 = BatchNormalization()(input@A1+b1) x1 = tf.nn.relu(x1) x2 = BatchNormalization()(x1@A2+b2) x2 = tf.nn.relu(x2) # 后续层同理
3. 添加权重正则化与梯度裁剪
- 给权重添加L2正则,抑制权重过大导致的矩阵乘法数值爆炸,在损失中加入:
l2_loss = tf.nn.l2_loss(A1) + tf.nn.l2_loss(A2) + tf.nn.l2_loss(A3) + tf.nn.l2_loss(A4) total_loss = compute_loss(input, true_values) + 1e-4 * l2_loss - 对优化器的梯度做裁剪,避免梯度爆炸导致权重更新幅度过大:
optimizer = tf.train.AdamOptimizer(0.0001) grads_and_vars = optimizer.compute_gradients(total_loss) clipped_grads = [(tf.clip_by_norm(g, 5.0), v) for g, v in grads_and_vars] train_op = optimizer.apply_gradients(clipped_grads)
4. 优化循环迭代逻辑
- 减少
max_iter的次数,避免数值误差累积; - 每几次循环后对
input重新做标准化,拉回数值分布范围。
调试方法:训练中查看张量值
1. 实时打印张量极值
在训练流程中加入打印操作,监控关键张量的数值范围:
# 在循环后或损失计算前,打印input的极值 print_op = tf.print("Input min/max:", tf.reduce_min(input), tf.reduce_max(input), output_stream=tf.logging.info) with tf.control_dependencies([print_op]): total_loss = compute_loss(input, true_values) + 1e-4 * l2_loss
运行时会输出每次迭代的input极值,可定位何时出现极端值。
2. TensorBoard监控数值趋势
将关键张量的统计量写入TensorBoard,直观观察数值变化:
tf.summary.scalar("input_min", tf.reduce_min(input)) tf.summary.scalar("input_max", tf.reduce_max(input)) tf.summary.scalar("input_mean", tf.reduce_mean(input)) merged_summary = tf.summary.merge_all() # 训练时同时运行summary并写入文件 summary, _ = sess.run([merged_summary, train_op], feed_dict=...) writer.add_summary(summary, global_step)
3. 触发NaN/Inf检查
用tf.debugging.check_numerics在关键节点添加检查,出现异常时直接报错定位:
input = tf.debugging.check_numerics(input, "Input contains NaN/Inf") output = tf.debugging.check_numerics(output, "Output contains NaN/Inf") total_loss = tf.debugging.check_numerics(total_loss, "Loss contains NaN/Inf")
内容的提问来源于stack exchange,提问作者miguelrc
相关产品推荐
相关产品推荐

