4.5 交叉熵损失函数
我们需要一个损失函数,用来表示:对于观测 x x x ,分类器输出 y ^ = σ ( w ⋅ x + b ) \hat y=\sigma(\mathbf w\cdot\mathbf x+b) y ^ = σ ( w ⋅ x + b ) 与正确输出 y y y (0 或 1)有多接近。y ^ \hat y y ^ 表示分类器赋予“观测 x x x 属于正类 1”的概率。因此,y ^ \hat y y ^ 和 y y y 都表示概率,但 y y y 总是 0 或 1,而 y ^ \hat y y ^ 还可以取两者之间的值。我们将这种损失或距离写作:
L ( y ^ , y ) = y ^ 与真实 y 的差异程度 (4.16) L(\hat y,y)=\text{$\hat y$ 与真实 $y$ 的差异程度}\tag{4.16} L ( y ^ , y ) = y ^ 与真实 y 的差异程度 ( 4.16 ) 所用的损失函数应让训练样本的正确类别标签获得更高概率。这称为条件最大似然估计 (conditional maximum likelihood estimation):选择参数 w , b \mathbf w,b w , b ,使给定观测 x x x 时,训练数据真实标签 y y y 的对数概率最大。由此得到的损失函数是负对数似然损失 (negative log likelihood loss),通常称为交叉熵损失 (cross-entropy loss)。
下面针对单个观测 x x x 推导该损失函数。我们希望学习能最大化正确标签概率 p ( y ∣ x ) p(y\mid x) p ( y ∣ x ) 的权重。由于只有 1 和 0 两种离散结果,这是一个伯努利分布。分类器对单个观测给出的概率可写为(若 y = 1 y=1 y = 1 ,式 4.17 化为 y ^ \hat y y ^ ;若 y = 0 y=0 y = 0 ,则化为 1 − y ^ 1-\hat y 1 − y ^ ):
p ( y ∣ x ) = y ^ y ( 1 − y ^ ) 1 − y (4.17) p(y\mid x)=\hat y^y(1-\hat y)^{1-y}\tag{4.17} p ( y ∣ x ) = y ^ y ( 1 − y ^ ) 1 − y ( 4.17 ) 对两边取对数。这样在数学上更方便,而且不会改变最优解:使概率最大的参数也会使其对数最大。
log p ( y ∣ x ) = log [ y ^ y ( 1 − y ^ ) 1 − y ] = y log y ^ + ( 1 − y ) log ( 1 − y ^ ) (4.18) \begin{array}{rcl}\log p(y\mid x)&=&\log[\hat y^y(1-\hat y)^{1-y}]\\&=&y\log\hat y+(1-y)\log(1-\hat y)\end{array}\tag{4.18} log p ( y ∣ x ) = = log [ y ^ y ( 1 − y ^ ) 1 − y ] y log y ^ + ( 1 − y ) log ( 1 − y ^ ) ( 4.18 ) 式 4.18 是需要最大化的对数似然。要把它变为需要最小化的损失函数,只需将其取负。于是得到交叉熵损失 L C E L_{\mathrm{CE}} L CE :
L C E ( y ^ , y ) = − log p ( y ∣ x ) = − [ y log y ^ + ( 1 − y ) log ( 1 − y ^ ) ] (4.19) L_{\mathrm{CE}}(\hat y,y)=-\log p(y\mid x)=-[y\log\hat y+(1-y)\log(1-\hat y)]\tag{4.19} L CE ( y ^ , y ) = − log p ( y ∣ x ) = − [ y log y ^ + ( 1 − y ) log ( 1 − y ^ )] ( 4.19 ) 最后代入 y ^ = σ ( w ⋅ x + b ) \hat y=\sigma(\mathbf w\cdot\mathbf x+b) y ^ = σ ( w ⋅ x + b ) :
L C E ( y ^ , y ) = − [ y log σ ( w ⋅ x + b ) + ( 1 − y ) log ( 1 − σ ( w ⋅ x + b ) ) ] (4.20) L_{\mathrm{CE}}(\hat y,y)=-[y\log\sigma(\mathbf w\cdot\mathbf x+b)+(1-y)\log(1-\sigma(\mathbf w\cdot\mathbf x+b))]\tag{4.20} L CE ( y ^ , y ) = − [ y log σ ( w ⋅ x + b ) + ( 1 − y ) log ( 1 − σ ( w ⋅ x + b ))] ( 4.20 ) 看看这个损失函数在图 4.2 的例子上是否符合预期。模型估计越接近正确答案,损失应越小;模型越困惑,损失应越大。先假设该情感样本的金标准标签为正面,即 y = 1 y=1 y = 1 。模型表现不错,因为式 4.8 赋予正面的概率 .70 高于负面的 .30。将 σ ( w ⋅ x + b ) = . 70 \sigma(\mathbf w\cdot\mathbf x+b)=.70 σ ( w ⋅ x + b ) = .70 和 y = 1 y=1 y = 1 代入式 4.20,右侧第二项消失(未注明底数时,log \log log 表示自然对数):
L C E ( y ^ , y ) = − [ y log σ ( w ⋅ x + b ) + ( 1 − y ) log ( 1 − σ ( w ⋅ x + b ) ) ] = − log σ ( w ⋅ x + b ) = − log ( . 70 ) = . 36 \begin{array}{rl}L_{\mathrm{CE}}(\hat y,y)&=-[y\log\sigma(\mathbf w\cdot\mathbf x+b)+(1-y)\log(1-\sigma(\mathbf w\cdot\mathbf x+b))]\\&=-\log\sigma(\mathbf w\cdot\mathbf x+b)\\&=-\log(.70)\\&=.36\end{array} L CE ( y ^ , y ) = − [ y log σ ( w ⋅ x + b ) + ( 1 − y ) log ( 1 − σ ( w ⋅ x + b ))] = − log σ ( w ⋅ x + b ) = − log ( .70 ) = .36 相反,假设图 4.2 中的样本其实是负面的,即 y = 0 y=0 y = 0 (也许评论者接着说:“But bottom line, the movie is terrible! I beg you not to see it!”)。这时模型判断错误,我们希望损失更高。将 y = 0 y=0 y = 0 和式 4.8 的 1 − σ ( w ⋅ x + b ) = . 30 1-\sigma(\mathbf w\cdot\mathbf x+b)=.30 1 − σ ( w ⋅ x + b ) = .30 代入式 4.20,第一项消失:
L C E ( y ^ , y ) = − [ y log σ ( w ⋅ x + b ) + ( 1 − y ) log ( 1 − σ ( w ⋅ x + b ) ) ] = − log ( 1 − σ ( w ⋅ x + b ) ) = − log ( . 30 ) = 1.2 \begin{array}{rl}L_{\mathrm{CE}}(\hat y,y)&=-[y\log\sigma(\mathbf w\cdot\mathbf x+b)+(1-y)\log(1-\sigma(\mathbf w\cdot\mathbf x+b))]\\&=-\log(1-\sigma(\mathbf w\cdot\mathbf x+b))\\&=-\log(.30)\\&=1.2\end{array} L CE ( y ^ , y ) = − [ y log σ ( w ⋅ x + b ) + ( 1 − y ) log ( 1 − σ ( w ⋅ x + b ))] = − log ( 1 − σ ( w ⋅ x + b )) = − log ( .30 ) = 1.2 果然,预测正确标签时的损失 .36 小于预测错误标签时的损失 1.2。
为什么最小化这个负对数概率能达到目标?完美分类器会给正确结果(y = 1 y=1 y = 1 或 y = 0 y=0 y = 0 )分配概率 1,给错误结果分配概率 0。若 y = 1 y=1 y = 1 ,y ^ \hat y y ^ 越高、越接近 1,分类器就越好;y ^ \hat y y ^ 越低、越接近 0,分类器就越差。若 y = 0 y=0 y = 0 ,则 1 − y ^ 1-\hat y 1 − y ^ 越高、越接近 1,分类器越好。对 y ^ \hat y y ^ (真实 y = 1 y=1 y = 1 时)或 1 − y ^ 1-\hat y 1 − y ^ (真实 y = 0 y=0 y = 0 时)取负对数,是一种方便的损失度量:它从 0(− log 1 -\log1 − log 1 ,无损失)延伸到无穷大(− log 0 -\log0 − log 0 ,无限损失)。
该函数还确保:正确答案的概率被最大化时,错误答案的概率同时被最小化。因为二者之和为 1,正确答案概率的任何增长都以错误答案概率的下降为代价。之所以称为交叉熵损失,是因为式 4.19 也正是真实概率分布 y y y 与估计分布 y ^ \hat y y ^ 之间的交叉熵公式。
现在我们已经知道要最小化什么;下一节将说明如何找到这个最小值。