NumPy实现稳定版sigmoid函数仍触发overflow溢出错误,原因是什么?
问题解答
你的核心理解偏差出在np.where的执行逻辑上:NumPy的np.where确实会先完整计算两个分支的所有输入表达式,再根据条件掩码选择对应位置的结果,和TensorFlow 1.x中tf.cond的分支惰性求值逻辑完全不同。
你实现的sigmoid逻辑本身是正确的,最终也输出了正确的[1., 0.]结果,触发溢出警告的原因是:
np.where接收的三个参数都是先完成全量计算后才会传入函数,两个分支的exp运算会覆盖输入数组的所有元素,不会根据条件跳过对应位置的计算- 当输入包含绝对值极大的负数时,else分支的
np.exp(-preds)会先被执行,此时-preds是绝对值极大的正数,直接超出exp的运算上界触发溢出 - 当输入包含绝对值极大的正数时,第一个分支的
np.exp(preds)会先被执行,同样触发溢出
这些溢出的中间结果虽然不会被最终选中,但计算过程已经触发了NumPy的运行时警告。
你写的逐元素for循环没有同类警告的原因是,每个元素只会执行对应分支的运算,不会触发不需要的极端值exp计算,自然不会产生溢出。
如果要保留向量化运算效率的同时规避警告,可以对输入做合理截断再计算,参考实现如下:
import numpy as np def stable_sigmoid(preds): # 限制输入范围在exp不会溢出的区间,不会影响sigmoid的计算精度 preds = np.clip(preds, -500, 500) return np.where(preds < 0, np.exp(preds)/(1+np.exp(preds)), 1/(1+np.exp(-preds)))
也可以直接使用SciPy中封装好的无警告稳定实现scipy.special.expit。
内容的提问来源于stack exchange,提问作者figs_and_nuts
相关产品推荐
相关产品推荐

