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

含f(x)=1/(1-exp(-x))-1/x的自定义损失函数训练NaN问题求解

解决自定义损失函数中f(x)导致的NaN问题

问题背景

自定义损失函数依赖函数f(x)+f(x+50),其中f(x)的数学定义为:

  • 当x≠0时,f(x) = 1/(1-exp(-x)) - 1/x
  • 当x=0时,f(x)=0.5
    该函数在全体实数域连续可微,取值范围0~1。原实现及多种尝试均导致训练过程中损失值变为NaN,仅替换为sigmoid函数可解决,但业务场景需要保留原函数。

排查方向

  1. 极端值的数值与梯度稳定性

    • 当x→-∞时,exp(-x)会溢出为inf,直接计算原公式的梯度exp(-x)/(1-exp(-x))²会出现inf/inf的NaN情况;
    • tf.where的硬分支切换会导致函数梯度在边界处不连续,易引发梯度突变或NaN。
  2. x=0附近的近似精度
    原实现仅用一阶泰勒近似0.5+x/12覆盖|x|<0.1的区域,与原函数的高阶项偏差可能导致梯度不连续,进而引发训练不稳定。

  3. f(x+50)的边界情况
    当x=-50时,x+50=0,若此处的近似或梯度处理不当,也会引入NaN风险。

具体解决方案建议

1. 平滑分段近似+连续梯度

用平滑过渡的分段函数替代硬分支切换,保证函数和梯度在全区间连续:

import tensorflow as tf

def f(x):
    # x接近0时,用三阶泰勒展开近似(更贴合原函数的梯度)
    taylor = 0.5 + x/12 + tf.pow(x, 3)/720
    # x>10时,近似为1 - 1/x(原函数极限)
    large_pos = 1 - 1/x
    # x<-10时,近似为 -1/x(原函数极限)
    large_neg = -1/x
    # 中间区域用原公式
    original = 1/(1 - tf.exp(-x)) - 1/x

    # 平滑切换掩码:用sigmoid实现软过渡,避免硬切换的梯度突变
    # x<-10与中间区域的过渡
    mask_neg = tf.sigmoid((x + 10)/0.5)
    # x>10与中间区域的过渡
    mask_pos = tf.sigmoid((10 - x)/0.5)
    # |x|<0.1与中间区域的过渡
    mask_zero = tf.sigmoid((0.1 - tf.abs(x))/0.01)

    # 组合各区域结果
    res = mask_neg * large_neg + (1 - mask_neg) * original
    res = mask_pos * res + (1 - mask_pos) * large_pos
    res = mask_zero * taylor + (1 - mask_zero) * res

    return res

2. 添加梯度裁剪

在优化器中设置梯度裁剪,限制梯度的最大范围,避免梯度爆炸引发NaN:

optimizer = tf.keras.optimizers.Adam(learning_rate=1e-4, clipnorm=1.0)

3. 调试时打印中间值定位问题

在自定义损失函数中添加打印逻辑,监控x、f(x)、f(x+50)及其梯度的极值,定位NaN出现的具体场景:

def custom_loss(y_true, y_pred):
    # 替换为你的x计算逻辑
    x = ... 
    fx = f(x)
    fx50 = f(x + 50)

    # 打印关键数值(仅调试阶段使用)
    tf.print("x 极值:", tf.reduce_min(x), tf.reduce_max(x))
    tf.print("f(x) 极值:", tf.reduce_min(fx), tf.reduce_max(fx))
    tf.print("f(x+50) 极值:", tf.reduce_min(fx50), tf.reduce_max(fx50))

    # 监控f(x)的梯度
    with tf.GradientTape() as tape:
        tape.watch(x)
        fx_debug = f(x)
    grad_fx = tape.gradient(fx_debug, x)
    tf.print("f(x)梯度极值:", tf.reduce_min(grad_fx), tf.reduce_max(grad_fx))

    # 替换为你的损失计算逻辑
    loss = ... 
    return loss

4. 验证浮点精度

保留tf.keras.backend.set_floatx('float64')的设置,结合上述分段近似,进一步降低浮点溢出的概率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 15:37:04