为何部分深度学习库未用log1p实现log_softmax?相关技术疑问
关于log_softmax的精度优化与实现探讨
PyTorch现有log_softmax实现的精度局限
PyTorch为提升数值稳定性,将log_softmax(x)实现为x - x.max() - (x - x.max()).exp().sum().log()。但当最大值远大于其余值时(float32下差值约16,float64下约36),最大值位置的log_softmax结果会返回0,而实际上我们可以获得更精确的计算结果。
数值示例如下:
>>> eps = torch.tensor(torch.finfo(torch.float32).eps) >>> -torch.tensor([1-(2*eps).log(), 0]).log_softmax(dim=0) tensor([1.1921e-07, 1.6249e+01]) >>> -torch.tensor([1-eps.log(), 0]).log_softmax(dim=0) tensor([-0.0000, 16.9424])
(float16下该值为7.2e-4)
精度提升的原理与改进实现
log和exp底层采用多项式展开类方式实现。float32在1附近的最小精度单位是1.19e-7,但它可正常表示小至约1.18e-38的正常值(即torch.finfo(torch.float32).smallest_normal),甚至亚正常值可小至约1.18e-45,因此存在精度提升的空间。我们可以借助expm1(x)(计算exp(x)-1且精度更高)和log1p(x)(计算log(1+x)且在x极小时精度更高)这类函数优化实现,改进后的log_softmax代码如下:
maxi = x.argmax() xoffset = x - x[maxi] xoffsetexp = xoffset.exp() # xoffsetexp[maxi] 当前约为1 xoffsetexp[maxi] = 0 xoffsetexp_sum_m1 = xoffsetexp.sum() return xoffset - xoffsetexp_sum_m1.log1p()
这种实现或许能让模型训练不会过早受浮点误差主导。
问题解答
1. 是否有深度学习库采用该方式实现log_softmax?
目前主流深度学习库(如PyTorch、TensorFlow、JAX)的官方log_softmax实现并未采用这种方式。不过部分研究代码或第三方开源实现中可能存在类似的精度优化版本,尤其是在对数值精度要求极高的特定场景(如高精度分类、小样本学习)中。
2. 存在哪些需避免采用该实现的原因?
- 额外计算开销:需要先执行
argmax定位最大值位置,再修改张量元素,相比原实现多了索引和赋值步骤,大规模张量计算场景下会增加少量耗时。 - 并行化效率降低:原实现的所有操作可完全并行,而修改最大值位置元素的步骤会引入串行操作,在GPU等并行计算设备上可能影响整体效率。
- 边缘情况处理复杂:当张量中存在多个相同最大值时,
argmax仅返回第一个最大值的索引,此时修改单个位置为0会导致计算偏差,需要额外处理多最大值场景,增加实现复杂度。 - 收益场景有限:只有当最大值远大于其他元素时才会体现精度优势,大多数常规深度学习训练场景中,原实现的精度已足够满足需求,额外优化带来的收益不明显。
内容的提问来源于stack exchange,提问作者Jason Gross
相关产品推荐
相关产品推荐

