Skip to article frontmatterSkip to article content
Site not loading correctly?

This may be due to an incorrect BASE_URL configuration. See the MyST Documentation for reference.

4.5 交叉熵损失函数

我们需要一个损失函数,用来表示:对于观测 xx,分类器输出 y^=σ(wx+b)\hat y=\sigma(\mathbf w\cdot\mathbf x+b) 与正确输出 yy(0 或 1)有多接近。y^\hat y 表示分类器赋予“观测 xx 属于正类 1”的概率。因此,y^\hat yyy 都表示概率,但 yy 总是 0 或 1,而 y^\hat y 还可以取两者之间的值。我们将这种损失或距离写作:

L(y^,y)=y^ 与真实 y 的差异程度(4.16)L(\hat y,y)=\text{$\hat y$ 与真实 $y$ 的差异程度}\tag{4.16}

所用的损失函数应让训练样本的正确类别标签获得更高概率。这称为条件最大似然估计(conditional maximum likelihood estimation):选择参数 w,b\mathbf w,b,使给定观测 xx 时,训练数据真实标签 yy 的对数概率最大。由此得到的损失函数是负对数似然损失(negative log likelihood loss),通常称为交叉熵损失(cross-entropy loss)。

下面针对单个观测 xx 推导该损失函数。我们希望学习能最大化正确标签概率 p(yx)p(y\mid x) 的权重。由于只有 1 和 0 两种离散结果,这是一个伯努利分布。分类器对单个观测给出的概率可写为(若 y=1y=1,式 4.17 化为 y^\hat y;若 y=0y=0,则化为 1y^1-\hat y):

p(yx)=y^y(1y^)1y(4.17)p(y\mid x)=\hat y^y(1-\hat y)^{1-y}\tag{4.17}

对两边取对数。这样在数学上更方便,而且不会改变最优解:使概率最大的参数也会使其对数最大。

logp(yx)=log[y^y(1y^)1y]=ylogy^+(1y)log(1y^)(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}

式 4.18 是需要最大化的对数似然。要把它变为需要最小化的损失函数,只需将其取负。于是得到交叉熵损失 LCEL_{\mathrm{CE}}

LCE(y^,y)=logp(yx)=[ylogy^+(1y)log(1y^)](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}

最后代入 y^=σ(wx+b)\hat y=\sigma(\mathbf w\cdot\mathbf x+b)

LCE(y^,y)=[ylogσ(wx+b)+(1y)log(1σ(wx+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}

看看这个损失函数在图 4.2 的例子上是否符合预期。模型估计越接近正确答案,损失应越小;模型越困惑,损失应越大。先假设该情感样本的金标准标签为正面,即 y=1y=1。模型表现不错,因为式 4.8 赋予正面的概率 .70 高于负面的 .30。将 σ(wx+b)=.70\sigma(\mathbf w\cdot\mathbf x+b)=.70y=1y=1 代入式 4.20,右侧第二项消失(未注明底数时,log\log 表示自然对数):

LCE(y^,y)=[ylogσ(wx+b)+(1y)log(1σ(wx+b))]=logσ(wx+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}

相反,假设图 4.2 中的样本其实是负面的,即 y=0y=0(也许评论者接着说:“But bottom line, the movie is terrible! I beg you not to see it!”)。这时模型判断错误,我们希望损失更高。将 y=0y=0 和式 4.8 的 1σ(wx+b)=.301-\sigma(\mathbf w\cdot\mathbf x+b)=.30 代入式 4.20,第一项消失:

LCE(y^,y)=[ylogσ(wx+b)+(1y)log(1σ(wx+b))]=log(1σ(wx+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}

果然,预测正确标签时的损失 .36 小于预测错误标签时的损失 1.2。

为什么最小化这个负对数概率能达到目标?完美分类器会给正确结果(y=1y=1y=0y=0)分配概率 1,给错误结果分配概率 0。若 y=1y=1y^\hat y 越高、越接近 1,分类器就越好;y^\hat y 越低、越接近 0,分类器就越差。若 y=0y=0,则 1y^1-\hat y 越高、越接近 1,分类器越好。对 y^\hat y(真实 y=1y=1 时)或 1y^1-\hat y(真实 y=0y=0 时)取负对数,是一种方便的损失度量:它从 0(log1-\log1,无损失)延伸到无穷大(log0-\log0,无限损失)。

该函数还确保:正确答案的概率被最大化时,错误答案的概率同时被最小化。因为二者之和为 1,正确答案概率的任何增长都以错误答案概率的下降为代价。之所以称为交叉熵损失,是因为式 4.19 也正是真实概率分布 yy 与估计分布 y^\hat y 之间的交叉熵公式。

现在我们已经知道要最小化什么;下一节将说明如何找到这个最小值。