You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.02 01:50:29