TensorFlow 2实现{-1,1}标签逻辑回归损失为nan是否正确
逻辑回归损失NaN问题排查与修复
你的实现逻辑整体思路符合标签取值为{-1,1}时的逻辑回归交叉熵损失公式,决策函数返回logits的写法也和损失计算逻辑匹配,出现NaN的核心原因是数值计算溢出,同时还有其他潜在问题可以一并排查。
核心问题原因
当y_true * y_pred[:, 0]的取值为绝对值很大的正数时,tf.math.exp(- y_true * y_pred[:, 0])会因为输入值过小(负数绝对值过大)出现上溢,计算结果为无穷大inf,再经过tf.math.log计算后就会得到nan,梯度反传过程中会进一步放大这个问题。
修复方案
方案1:使用TensorFlow内置数值稳定的损失实现(推荐)
直接调用官方封装好的sigmoid_cross_entropy_with_logits接口,官方已经做了数值溢出兼容,只需要把你的{-1,1}标签转为{0,1}格式即可:
def cross_entropy(y_pred, y_true): # 标签从{-1, 1}映射到{0, 1} y_true_01 = (y_true + 1) / 2 ce_loss = tf.reduce_mean(tf.nn.sigmoid_cross_entropy_with_logits( labels=y_true_01, logits=y_pred[:, 0] )) return ce_loss + tf.nn.l2_loss(W)
方案2:手动实现数值稳定的损失计算
如果需要自行实现损失,可以对公式做等价改写,避免直接计算大参数的指数函数:
def cross_entropy(y_pred, y_true): z = y_true * y_pred[:, 0] # 等价于log(1 + exp(-z)),规避数值溢出 stable_ce = tf.maximum(0., -z) + tf.math.log(1. + tf.math.exp(-tf.abs(z))) return tf.reduce_mean(stable_ce) + tf.nn.l2_loss(W)
其他潜在问题排查
- 检查输入特征是否做了归一化,特征取值过大会导致线性层输出数值爆炸,触发溢出
- 检查权重
W和偏置b的初始化逻辑,避免初始权重过大导致初始预测值过高 - 检查学习率是否设置过大,梯度过大会导致权重更新后取值异常,引发数值溢出
内容的提问来源于stack exchange,提问作者Avijit Dasgupta
相关产品推荐
相关产品推荐

