交叉熵损失的形式
对于回归任务,特别是简单的回归问题,我们通常使用 MSE(Mean Squared Error,均方误差)作为损失函数,少数情况下也使用 MAE(Mean Absolute Error)。MSE 有诸多优良性质,比如光滑、便于求导且导数仍光滑、在某些理想条件下具有凸性且等价于 MLE(Maximum Likelihood Estimation,极大似然估计)。此外,MSE 本身能被分解为偏差的平方、方差与噪声之和,这也为 MSE 提供了一种合理的解释。关于这部分内容,可见于另一篇我在本科时为《回归分析》课程所撰写的笔记:回归分析。
不过,今天我想从统计学与信息论两个角度分别讨论一个在分类任务中被广泛应用的损失函数——CE(Cross Entropy,交叉熵)损失。
CE 的形式十分简单,在 $M$ 分类问题中,对于输入样本 $x$ ,记模型对 $x$ 关于第 $m$ 个分类的输出预测概率为 $q_m(x)$ 或 $q(y_m\mid x)$,而真实的概率为 $p_m(x)$ 或 $p(y_m\mid x)$,则样本 $x$ 的 CE 为
$$ CE_x=-\sum^M_{i=1}p_i(x)\log_2q_i(x) $$对于事先给出准确标签的有监督学习,在训练过程中往往有数据的真实标签,因此真实概率 $p$ 通常是 One-Hot 编码的,此时 CE 可简化为
$$ CE_x=-\log_2q_{y}(x) $$若考虑全部的 $N$ 个样本,则取各样本的平均 CE 作为损失函数,即
$$ CE=-\frac{1}{N}\sum^N_{i=1}\sum^M_{j=1}p_j(x_i)\log_2q_j(x_i) $$对于 One-Hot 编码,可进一步简化为
$$ CE=-\frac1{N}\sum_{i=1}^{N}\log_2q_{y_i}(x_i) $$这应当是人尽皆知的故事,几乎是我们的常识。可是,为什么 CE 能够在(多)分类任务损失函数家族中占据统治地位呢?本文将从统计学与信息论两个视角分别讨论这一点,说明为什么 CE 适合分类任务。
统计学视角
从上文给出的 CE 表达式中我们可以知道,最小化 CE 完全等价于最大化似然,因为二者具有完全一致的单调性。
对于固定输入 $x$ 的分类任务,标签 $Y$ 服从参数为 $q(y\mid x)$ 的分类分布(Categorical Distribution)。若忽略输入差异,将所有样本视为来自同一个类别分布,则退化为多项分布。假设样本是独立同分布的,则该分类分布的似然为
$$ L(\theta;x,y)=\prod^N_{i=1}q_{\theta}(y_i\mid x_i) $$如果所有样本均来自同一个类别分布,则该似然的表达式与多项分布似然的形式 $\displaystyle\frac{N!}{k_1!k_2!\cdots k_M!}\prod^M_{i=1}q_i^{k_i}$ 只相差常数倍。
接着对上式取对数。取对数有两点原因:
- 对似然函数取对数是计算 MLE 的常见手法,目的是将连乘转为连加,方便计算。无论是笔算利用 Lagrange 乘子法求 MLE 的解析解,还是利用计算机通过沿梯度的反向传播与梯度下降优化参数,都能够大幅降低计算量。
- 数值上更稳定。如果不取对数,由于 $q_i\in[0,1]$,因此当样本量充分大时,不可避免地 $\displaystyle\prod^N_{i=1}q_{\theta}(y_i\mid x_i)$ 将接近 $0$。当在计算机上进行计算时,这也是做对数运算的一个重要理由。
取对数后再除样本总数 $N$ 以使数值范围可控,有
$$ \frac1{N}\log L(\theta;x,y)=\frac1{N}\sum_{i=1}^{N}\log q_{\theta}(y_i\mid x_i) $$该形式与 CE 仅相差一个负号,因此最大化似然与最小化 CE 是完全等价的。
综上所述,如果将 CE 作为损失函数,则在分类问题中所得到的估计值正是 MLE。
上文所推导出的负对数似然是没有 $p_j(x_i)$ 项的,因为在最初分类分布的似然中就隐含了按真实标签的 One-Hot 编码 $p(x_i)$ 的结果。
信息论视角
随机变量的编码
对于信源编码(Source Coding),我们希望找到一种编码方式,使得平均编码长度足够小。
举一个例子,假设我们作为消息的发送方,将等概率地发送 A、B、C 与 D 四种消息,那么我们可以按下述方式编码:
- A:00
- B:01
- C:10
- D:11
这样,所有消息的码长均为 2,由于每种消息都具有相同的概率被选中,因此平均编码长度也为 2。易见这是最优编码策略。
可假设我们发送 A 的概率是 $\frac12$、发送 B 的概率是 $\frac14$,而发送 C 与 D 的概率均为 $\frac18$ 呢?如果依然按上述方式编码,则平均编码长度为 $\frac12\times2+\frac14\times 2+\frac18\times2+\frac18\times2=2$。然而,这并不是该情况下的最优编码方案。
最优编码为:
- A:0
- B:10
- C:110
- D:111
该编码方案的平均码长为 $\frac12\times1+\frac14\times 2+\frac18\times3+\frac18\times3=1.75$。这是因为,我们发送 $A$ 的概率足够大,因此应当为 $A$ 分配最短的编码。
熵与最优编码
构建即时码(Instantaneous Code)的基本要求是任何一个码字都不应是另一个码字的前缀——这是系统无需前瞻便可即时唯一解码的充要条件。在该要求下,我们希望得到最优编码。
事实上,信源编码定理(Source Coding Theorem)指出:对于某个信源,存在一种编码方式使其平均码长任意接近于该信源的熵(Entropy),但不可能低于熵。
这意味着最优编码不会低于熵。下面我们可以简要地证明这一结论。
根据 Gibbs 不等式,对于概率分布 $p$ 与 $q$,有
$$ -\sum p_i\log_2 p_i\leqslant-\sum p_i\log_2 q_i $$构造辅助概率分布 $q_i=2^{-l_i}$ 并代入,得
$$ -\sum p_i\log_2 p_i\leqslant-\sum p_i\log_2 q_i=\sum p_il_i $$当且仅当 $p=q$ 时 Gibbs 不等式取等,即 $p_i=q_i=2^{-l_i}$,此时 $l_i=-\log_2 p_i$。由于 $\sum p_il_i$ 为平均码长,因此 $-\sum p_i\log_2 p_i$ 便是平均码长的下界。这也意味着最优码长理论上满足 $l_i=-\log_2 p_i$。
Kraft 不等式表明,唯一可译码的码字长度必须满足充要条件 $\sum D^{-l_i}\leqslant 1$,其中 $D$ 为字母表所包含的元素数量,$l_i$ 为第 $i$ 个码字长度。对于二进制编码,即 $\sum2^{-l_i}\leqslant1$。因此,我们构造的辅助概率分布 $q_i=2^{-l_i}$ 是合法的,和值不会大于 $1$。
至于和值不严格等于 1 的问题,只需要令 $q_0=1-\sum 2^{-l_i}$ 作为哑元即可。
下界值 $-\sum p_i\log_2 p_i$ 正是熵(Entropy),记为 $H(p)$。
霍夫曼编码(Huffman Coding)就是一种最优的信源编码,但这不是本文的重点。总而言之,熵刻画了一个概率分布本身所蕴含的信息量,也等价于对该信源进行无损编码时所能达到的理论最小平均码长。
如何理解熵刻画了概率分布的信息量?简而言之:概率越大,码长应约短。
由于小概率事件出现时更加出人意料,因此我们需要更多比特描述它,这也意味着它携带了更多信息。熵正是这种信息量在整个概率分布上的平均值。
交叉熵
在实际问题中,我们往往并不知道真实分布 $p$,而只能通过建立模型 $q$ 近似 $p$。即,按照模型输出的分布设计最优码本,却用于编码来自真实分布的数据。CE 恰好给出了此时所得到的平均码长。
对于平均码长 $\sum p_il_i$,当我们使用 $q$ 编码时,最优码长理论上有 $l_i=-\log_2 q_i$,代入平均码长公式,得到
$$ -\sum p_i\log_2 q_i $$该值正是 CE,通常记为 $H(p,q)$。
因此,Gibbs 不等式也可以被写为下述形式
$$ H(p)\leqslant H(p,q) $$而不等式的差值 $H(p,q)-H(p)=\sum p_i\log_2\frac{p_i}{q_i}$ 正是 KL 散度(Kullback-Leibler Divergence)$D_{\text{KL}}(p\Vert q)$。KL 散度衡量了两个分布的近似程度。
所以,当我们最小化交叉熵,便最小化了 KL 散度,使得模型的预测分布最接近真实分布,此时用 $q$ 编码 $p$ 的平均码长逼近其理论下界——熵。
补充讨论
上文分别从统计学和信息论的视角论证了在分类任务中使用 CE 作为损失的合理性。除此之外,CE 在深度学习中的广泛应用还有一个重要原因:当它与 softmax 配合使用时,梯度具有十分简洁的形式。
记 logits 为 $z$、softmax 输出为 $q_i=\frac{e^{z_i}}{\sum_j e^{z_j}}$,则 CE 关于 $z_i$ 的梯度为
$$ \frac{\partial CE}{\partial z_i}=q_i-p_i $$可以看到,该梯度与线性回归中 MSE 关于预测值的梯度 $\hat{y}-y$ 在形式上完全一致,均为「预测值减去真实值」。
相比之下,若使用 MSE 配合 sigmoid 或 softmax 输出,根据链式法则,在梯度中还需要乘上 sigmoid(或 softmax)的导数。当模型给出一个置信度极高但实际错误的预测时,由于此时输出已进入饱和区域,故导数趋近于零,从而使整体梯度显著减小,导致参数更新缓慢。
但 CE 与 softmax 在联合使用时,该导数项恰好在求导时被抵消,因此即使模型做出了高置信度的错误预测,仍能够产生较大的梯度,这有利于模型优化。
因此,CE 不仅在统计学与信息论的视角下是合理的、有意义的,在数值上也是利于训练的,兼具理论与工程价值。
对于二分类任务,CE 退化为
$$ CE = -\frac{1}{N}\sum_{i=1}^N \bigl[ y_i \log q_i + (1-y_i) \log(1-q_i) \bigr] $$这恰好是 Logistic 回归所使用的负对数似然损失。