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

TensorFlow使用Gradient Tape训练时梯度出现NaN问题求助

梯度NaN问题排查与解决:自定义Balanced Logarithmic Loss场景

问题背景

使用TensorFlow Gradient Tape构建三层神经网络,采用自定义加权对数损失函数bal_log_loss训练时,出现梯度NaN的问题,已通过断点捕获终止训练,但需彻底定位原因并解决。

训练代码

flag = False
W1, b1 = initialize_parameters_deep(n, layers_dims[0])
W2, b2 = initialize_parameters_deep(layers_dims[0], layers_dims[1])
W3, b3 = initialize_parameters_deep(layers_dims[1], layers_dims[2])

for i in range(5000):
  with tf.GradientTape() as tape:
    Z1 = linear_activation_forward(X, W1, b1, 'leaky_relu')
    Z2 = linear_activation_forward(Z1, W2, W2, b2, 'leaky_relu')

    y_before = linear_activation_forward(Z2, W3, b3, 'sigmoid')

    if tf.reduce_sum(tf.cast(tf.math.is_nan(y_before), dtype=tf.int32)) > 0:
      break

    y_predict = linear_activation_forward(Z2, W3, b3, 'sigmoid')

    loss = bal_log_loss(y_true=y_true, y_pred =y_predict)

  if ((i+1) % 300) == 0:
      print(f'Iteration: {i}')
      print(f'Loss: {loss}')
      print('_-'*30)


  [dW1, dW2, dW3, db1, db2, db3] = tape.gradient(loss, [W1, W2, W3, b1, b2, b3])

  gradients = [dW1, dW2, dW3, db1, db2, db3]

  for g in gradients:
    if tf.reduce_sum(tf.cast(tf.math.is_nan(g), dtype=tf.int32)) > 0:
      print(f'Iteration: {i}')
      print(f'loss: {loss}')
      flag = True
      break
  

  if flag:
    break
  W1.assign_sub(dW1 * learning_rate)
  W2.assign_sub(dW2 * learning_rate)
  W3.assign_sub(dW3 * learning_rate)

  b1.assign_sub(db1 * learning_rate)
  b2.assign_sub(db2 * learning_rate)
  b3.assign_sub(db3 * learning_rate)

自定义损失函数代码

def bal_log_loss(y_true, y_pred):
  epsilon = 1e-10
  m = y_true.shape[-1]

  y_pred = tf.where(y_pred>(1-epsilon), 1-epsilon, y_pred)
  y_pred = tf.where(y_pred<epsilon, epsilon, y_pred)

  n1 = tf.reduce_sum(y_true)
  n0 = tf.reduce_sum(1-y_true)

  w1 = -1/n1
  w0 = -1/n0

  log = tf.where(y_true==1, w1 * tf.math.log(y_pred), w0 * tf.math.log(1-y_pred))
  loss = (tf.reduce_sum(log) / 2) 
  return loss

问题成因分析

  1. 类别样本数为0的除以零风险
    当训练集中某一类样本完全缺失时(n1=0或n0=0),计算w1=-1/n1或w0=-1/n0会直接产生inf或-inf,后续与log值相乘后会输出NaN,反向传播时梯度自然也会变成NaN。

  2. 梯度爆炸引发的NaN
    虽然对y_pred做了截断处理,但反向传播时:

    • log(y_pred)的梯度为1/y_pred,当y_pred接近epsilon时,梯度会变得极大(比如epsilon=1e-10时,梯度为1e10)
    • log(1-y_pred)的梯度为-1/(1-y_pred),当y_pred接近1-epsilon时,梯度同样会极端放大
      再加上w1/w0的加权作用,很容易触发梯度爆炸,最终变成NaN。
  3. 冗余计算的误差累积
    训练循环中两次调用linear_activation_forward计算sigmoid输出,虽然逻辑上不影响,但额外的计算可能增加数值误差累积的概率,间接导致NaN出现。


解决办法

1. 修复类别样本数为0的情况

在计算权重时添加最小分母保护,避免除以零:

def bal_log_loss(y_true, y_pred):
  epsilon = 1e-10
  m = y_true.shape[-1]

  y_pred = tf.where(y_pred>(1-epsilon), 1-epsilon, y_pred)
  y_pred = tf.where(y_pred<epsilon, epsilon, y_pred)

  n1 = tf.reduce_sum(y_true)
  n0 = tf.reduce_sum(1-y_true)

  # 用tf.maximum设置最小分母,防止除以零
  w1 = -1 / tf.maximum(n1, 1e-6)
  w0 = -1 / tf.maximum(n0, 1e-6)

  log = tf.where(y_true==1, w1 * tf.math.log(y_pred), w0 * tf.math.log(1-y_pred))
  loss = (tf.reduce_sum(log) / 2) 
  return loss

2. 优化损失函数的梯度稳定性

改用更稳定的方式计算加权交叉熵,同时调整epsilon值(过小的epsilon会导致梯度极端放大):

def bal_log_loss(y_true, y_pred):
    epsilon = 1e-7  # 调整为更合理的截断阈值
    y_pred = tf.clip_by_value(y_pred, epsilon, 1 - epsilon)
    
    n1 = tf.reduce_sum(y_true)
    n0 = tf.reduce_sum(1 - y_true)
    total = tf.maximum(n1 + n0, 1e-6)
    
    # 采用更合理的权重计算方式(样本数反比)
    w1 = n0 / total
    w0 = n1 / total
    
    # 拆分计算,避免where带来的梯度异常
    loss_pos = -w1 * y_true * tf.math.log(y_pred)
    loss_neg = -w0 * (1 - y_true) * tf.math.log(1 - y_pred)
    loss = tf.reduce_mean(loss_pos + loss_neg)
    return loss

3. 添加梯度裁剪

在参数更新前对梯度进行裁剪,限制最大范数,避免梯度爆炸:

# 计算梯度后添加裁剪逻辑
clip_norm = 1.0  # 根据模型规模调整
[dW1, dW2, dW3, db1, db2, db3] = [tf.clip_by_norm(g, clip_norm) for g in gradients]

4. 简化训练循环逻辑

去掉冗余的sigmoid计算,合并NaN检查:

for i in range(5000):
  with tf.GradientTape() as tape:
    Z1 = linear_activation_forward(X, W1, b1, 'leaky_relu')
    Z2 = linear_activation_forward(Z1, W2, b2, 'leaky_relu')
    y_predict = linear_activation_forward(Z2, W3, b3, 'sigmoid')
    
    # 简化NaN检查逻辑
    if tf.reduce_any(tf.math.is_nan(y_predict)):
        print(f'NaN detected in prediction at iteration {i}')
        break
        
    loss = bal_log_loss(y_true=y_true, y_pred=y_predict)

  # 后续打印、梯度计算、参数更新逻辑不变

5. 调整初始化与学习率

  • 检查参数初始化是否过大:过大的初始权重会导致激活值饱和,sigmoid输出接近0或1,触发极端梯度
  • 降低学习率:若当前学习率过高,大梯度更新会导致参数跳变,进一步引发NaN,可尝试从0.01开始逐步调试

内容的提问来源于stack exchange,提问作者Juan Carlos Ramírez Tinoco

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 05:08:11