d2l 学习笔记:softmax 工程实现
这一小节(D2L 3.7.2 重新审视 Softmax 的实现)其实是在解决一个工程实现问题,不是在引入新的机器学习思想。
前面 3.4 节你已经理解了 Softmax 的数学意义:
y^j=∑keokeoj
其中:
理论上没有问题,但是计算机实际计算这个公式时,会遇到数值稳定性问题。
这一节主要讲两个问题:
-
为什么直接计算 Softmax 会出现上溢和下溢;
-
为什么实际框架会对 Softmax 和交叉熵进行特殊处理。
1. 先回顾 Softmax 的计算过程
假设一个图片分类模型最后输出:
o=[10,5,1]
Softmax需要计算:
e10,e5,e1
得到:
[22026,148,2.7]
然后:
y^=22026+148+2.7[22026,148,2.7]
得到:
[0.993,0.0067,0.0001]
这个过程没有问题。
但是如果模型输出非常大:
o=[1000,900,800]
Softmax需要:
e1000
问题来了。
2. 什么是上溢(overflow)
计算机存储数字是有限范围的。
例如:
float32:
最大大约:
3.4×1038
但是:
e1000
约等于:
10434
远远超过计算机能表示的范围。
于是:
e1000→inf
也就是:
无限大。
于是Softmax:
inf+inf+infinf
计算机不知道结果是多少。
可能得到:
这就是上溢。
3. 为什么Softmax特别容易出现这个问题?
因为指数函数增长太快。
看一下:
e1=2.7
e10=22026
e100=2.6×1043
只增加100倍,结果增加几十个数量级。
所以:
Softmax数学上正确,但是直接计算不稳定。
4. 解决方法:减去最大值
书中提出:
在计算Softmax之前:
oj′=oj−max(o)
例如:
原始:
o=[1000,900,800]
最大值:
max(o)=1000
减去:
o′=[0,−100,−200]
然后计算:
eo′
得到:
[1,e−100,e−200]
现在指数范围非常安全。
5. 为什么减最大值不会改变Softmax结果?
这是这一节最核心的数学点。
原公式:
y^j=∑keokeoj
现在每个元素减去:
m=max(o)
得到:
∑keok−meoj−m
根据指数性质:
6. $$
e^{o_j-m}
\frac{e^{o_j}}{e^m}
所以:分子:
\frac{e^{o_j}}{e^m}
分母:
\sum_k\frac{e^{o_k}}{e^m}
分子分母同时除以:
e^m
因此:
## 7. $$
\frac{\frac{e^{o_j}}{e^m}}
{\sum_k\frac{e^{o_k}}{e^m}}
\frac{e^{o_j}}
{\sum_ke^{o_k}}
结果完全一样。
所以:
减最大值只是改变计算过程,不改变数学结果。
8. 那为什么又出现下溢?
解决上溢之后:
oj−max(o)
可能出现很大的负数。
例如:
[0,−1000,−2000]
计算:
e−1000
非常接近:
0
计算机可能直接存成:
0
这叫:
下溢(underflow)。
9. 为什么下溢会影响训练?
Softmax之后:
y^j=0
然后交叉熵:
L=−log(y^j)
如果:
y^j=0
那么:
log(0)
数学上:
=−∞
于是:
−log(0)=+∞
然后反向传播:
梯度里面出现:
inf
最终:
参数:
W,b
可能变成:
nan
训练直接崩溃。
10. 那为什么书里说把Softmax和交叉熵放一起?
这是最容易误解的地方。
普通实现:
先:
o
↓
Softmax:
y^
↓
log:
log(y^)
↓
交叉熵。
但是问题:
Softmax已经产生了极小概率。
例如:
y^=10−300
再:
log(10−300)
容易出现精度问题。
现代框架一般直接计算:
CrossEntropy(o,y)
而不是:
−log(Softmax(o))
因为框架内部会使用:
LogSoftmax
也就是:
直接计算:
log∑keokeoj
利用数学变换:
11. $$
\log Softmax(o_j)
o_j-\log\sum_ke^{o_k}
避免真正计算:
e^{o_j}
从根源避免指数爆炸。
---
## 12. 用一句话总结这一节
3.7.2 不是讲新的Softmax理论,而是在讲:
> Softmax数学公式虽然简单,但是直接计算指数和概率会产生数值稳定性问题。因此实际深度学习框架会通过减去最大logit、使用LogSoftmax以及融合交叉熵计算的方法,让模型在计算机有限精度下仍然稳定训练。
你前面学习的流程可以串起来:
X
↓线性层:
o=XW+b
↓得到类别评分:
logits
↓Softmax:
\hat y=P(y|x)
↓交叉熵:
L=-\sum y_j\log\hat y_j
↓反向传播:
\frac{\partial L}{\partial W}
↓更新参数。而这一节解决的是中间这一步:
logits\rightarrow Softmax\rightarrow Loss
在真实计算机环境中如何稳定运行。你前面已经理解了Softmax的概率意义,这一节其实是在补充“为什么深度学习框架里面的Softmax和课本公式看起来不完全一样”。−−−>[!info]关联文档>−[[d2l学习笔记:引言]]:D2L学习入口。>−[[d2l学习笔记:线性回归的损失函数]]:线性回归与损失函数基础。