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

TensorFlow训练中使用tf.math.is_nan自定义损失出现NaN loss问题求助

核心问题原因
  • TensorFlow的tf.where在前向传播时只会返回符合条件的分支结果,但两个分支的计算逻辑都会执行,且反向传播时两个分支的梯度都会被计算。
  • 你的旧损失代码中,当true_labels存在NaN时,tf.square(tf.subtract(true_labels, predicted_labels))这一步会先计算出NaN结果,即使tf.where最终选择了0的分支,反向传播时未被选中分支的NaN梯度会回传到模型参数,导致参数被更新为NaN,后续前向传播的predicted_labels自然全部为NaN。
  • 非训练模式下不需要计算反向梯度,所以人工测试时不会出现NaN问题。
  • 你替换为-1e4的方案能生效,本质是避免了NaN参与算术运算,未被选中的分支计算结果和梯度都是正常值,不会污染参数。
正确实现方案

不需要修改标签预处理逻辑,直接调整损失函数的计算顺序,先把NaN从算术运算中移除再计算损失,代码如下:

def custom_mse_loss(true_labels, predicted_labels):
    # 构造掩码:非NaN位置为1,NaN位置为0
    mask = tf.cast(tf.math.logical_not(tf.math.is_nan(true_labels)), tf.float32)
    # 把标签中的NaN替换为任意数值(此处填0即可,后续会被掩码消掉)
    true_labels_filled = tf.where(tf.math.is_nan(true_labels), 0.0, true_labels)
    # 计算平方差后乘以掩码,消掉NaN位置的损失
    squared_error = tf.square(true_labels_filled - predicted_labels) * mask
    return tf.reduce_mean(squared_error)

因为你的标签是整行全NaN/全非NaN,也可以按行构造掩码提升计算效率,上述按元素的实现也能正常运行。

可选优化建议

如果你的5个回归目标取值范围差异较大(比如示例中的106、189和2.64、19等),可以对回归标签做标准化处理,能进一步提升训练稳定性,避免梯度爆炸问题。

内容的提问来源于stack exchange,提问作者Varun

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 09:24:08