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
问题成因分析
类别样本数为0的除以零风险
当训练集中某一类样本完全缺失时(n1=0或n0=0),计算w1=-1/n1或w0=-1/n0会直接产生inf或-inf,后续与log值相乘后会输出NaN,反向传播时梯度自然也会变成NaN。梯度爆炸引发的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。
冗余计算的误差累积
训练循环中两次调用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
相关产品推荐
相关产品推荐

