TensorFlow自定义Softmax出现NaN问题求助
问题分析:手动实现Softmax导致NaN的原因及解决办法
我之前也踩过这个坑,其实核心问题就是数值溢出,咱们慢慢理清楚:
为什么手动实现会出NaN?
你自己写的tf.exp(s) / tf.reduce_sum(tf.exp(s), axis=1)没做数值稳定处理:
- 当
s里的元素值偏大时,tf.exp(s)会直接溢出变成inf(无穷大); - 分子分母都是
inf时,计算结果就成了NaN; - 这个NaN会顺着计算图扩散,不仅
p是NaN,反向传播时梯度也会变成NaN,进而导致后续迭代中W、b更新异常,最终连s也跟着变成NaN。
而tf.nn.softmax内部已经帮你做了数值稳定优化——它会先对s做减最大值的处理:s_stable = s - tf.reduce_max(s, axis=1, keepdims=True),把exp的输入范围限制在合理区间,从根源避免溢出。
修复后的代码
只需要给你的手动实现加上数值稳定步骤就行,其余代码完全不用改:
tempX = x tempW = W tempMult = tf.matmul(tempX, W) s = tempMult + b # ---------------------------- # p = tf.nn.softmax(s) # 新增数值稳定处理 s_stable = s - tf.reduce_max(s, axis=1, keepdims=True) p = tf.exp(s_stable) / tf.reduce_sum(tf.exp(s_stable), axis=1, keepdims=True) # ------------------------------ myTemp = y*tf.log(p) cost = tf.reduce_mean(-tf.reduce_sum(myTemp, axis=1))
这里要注意给reduce_sum加上keepdims=True,保证分子分母的维度匹配,避免广播出错。
额外小提醒
如果加了数值稳定还是出现NaN,那大概率是学习率设置太高了——过大的学习率会让参数更新幅度过大,导致s的数值再次变得极端。可以试试调小学习率,比如从0.1降到0.01或者更小。
内容的提问来源于stack exchange,提问作者V.Vocor
相关产品推荐
相关产品推荐

