欢迎光临
我们一直在努力

损失函数大汇总(四)(Cross Entropy Loss附公式推导与代码)

损失函数大汇总(四)(Cross Entropy Loss附公式推导与代码)

一、引言

前面三篇写的都是回归任务里常见的损失函数:MSE、MAE 和 Huber Loss。它们衡量的都是预测值和真实值之间的数值偏差,适合连续变量预测。

但在分类任务里,问题就不一样了。分类模型通常不是直接输出一个连续实数,而是输出每个类别对应的概率,最后再根据概率大小决定样本属于哪一类。这个时候,如果还用 MSE 这类回归损失去训练,往往效果并不好。分类任务更常用的,是 Cross Entropy Loss(交叉熵损失)。

交叉熵损失几乎是现代分类模型的默认选择。无论是图像分类、文本分类,还是很多多类别识别任务,基本都会碰到它。本文就详细整理一下 Cross Entropy Loss 的定义、由来、推导、梯度、使用场景和代码实现。

限于笔者水平,文中如有疏漏或错误,欢迎留言交流。

二、Cross Entropy Loss

Cross Entropy Loss,中文一般叫交叉熵损失。它本质上用来衡量:模型预测的概率分布,与真实标签对应分布之间到底差了多少。

在分类问题中,我们通常希望模型给真实类别分配尽可能大的概率。交叉熵损失做的事情,其实就是:

  • 如果模型给真实类别的概率很高,损失就小;
  • 如果模型给真实类别的概率很低,损失就大。

因此,优化交叉熵损失,本质上就是在逼着模型把正确类别的概率往上提。

1. 数学定义

先看最常见的多分类情形。

假设一个样本共有 CCC 个类别,模型输出的预测概率分布为:

y^=[y^1,y^2,…,y^C]
\\hat{\\mathbf{y}} = [\\hat{y}_1, \\hat{y}_2, \\dots, \\hat{y}_C]
y^=[y^1,y^2,,y^C]

其中:

y^c∈(0,1),∑c=1Cy^c=1
\\hat{y}_c \\in (0,1), \\qquad \\sum_{c=1}^{C}\\hat{y}_c = 1
y^c(0,1),c=1Cy^c=1

真实标签通常写成 one-hot 形式:

y=[y1,y2,…,yC]
\\mathbf{y} = [y_1, y_2, \\dots, y_C]
y=[y1,y2,,yC]

其中只有真实类别对应的位置为 1,其余位置为 0。

那么单个样本的交叉熵损失定义为:

LCE=−∑c=1Cyclog⁡y^c
\\mathcal{L}_{\\mathrm{CE}} = -\\sum_{c=1}^{C} y_c \\log \\hat{y}_c
LCE=c=1Cyclogy^c

由于 one-hot 标签中只有真实类别那一项为 1,所以这个式子实际上会简化成:

LCE=−log⁡y^k
\\mathcal{L}_{\\mathrm{CE}} = -\\log \\hat{y}_k
LCE=logy^k

其中 kkk 表示真实类别。

也就是说,交叉熵损失其实就是对“真实类别预测概率”的对数取负号。

对于包含 NNN 个样本的数据集,平均交叉熵损失可写为:

LCE=−1N∑i=1N∑c=1Cyi,clog⁡y^i,c
\\mathcal{L}_{\\mathrm{CE}} = -\\frac{1}{N}\\sum_{i=1}^{N}\\sum_{c=1}^{C} y_{i,c}\\log \\hat{y}_{i,c}
LCE=N1i=1Nc=1Cyi,clogy^i,c

2. Softmax 与交叉熵的关系

在多分类问题中,模型通常不会直接输出概率,而是先输出一组实数分数,称为 logits:

z=[z1,z2,…,zC]
\\mathbf{z} = [z_1, z_2, \\dots, z_C]
z=[z1,z2,,zC]

然后通过 Softmax 把它们变成概率:

y^c=ezc∑j=1Cezj
\\hat{y}_c = \\frac{e^{z_c}}{\\sum_{j=1}^{C} e^{z_j}}
y^c=j=1Cezjezc

所以,多分类训练里常说的交叉熵损失,通常实际指的是:

  • 模型先输出 logits;
  • logits 经过 Softmax 得到概率;
  • 再和真实标签计算交叉熵。

这也是为什么在 PyTorch 里,nn.CrossEntropyLoss 直接接收 logits,而不是接收已经做过 Softmax 的概率。因为它内部已经把 Softmax 和交叉熵一起做了,而且数值稳定性更好。

3. 由来与推导

交叉熵损失的来源,既可以从信息论角度理解,也可以从最大似然估计角度推出来。

(1)从最大似然估计理解

假设对于一个样本,模型预测属于各类别的概率为:

p(y=c∣x)=y^c
p(y=c \\mid x) = \\hat{y}_c
p(y=cx)=y^c

如果真实类别是 kkk,那么这个样本的似然就是:

p(y=k∣x)=y^k
p(y=k \\mid x) = \\hat{y}_k
p(y=kx)=y^k

对于整个数据集,似然函数为:

Llikelihood=∏i=1Ny^i,ki
\\mathcal{L}_{\\mathrm{likelihood}} = \\prod_{i=1}^{N}\\hat{y}_{i,k_i}
Llikelihood=i=1Ny^i,ki

其中 kik_iki 表示第 iii 个样本的真实类别。

为了便于优化,通常取对数,得到对数似然:

log⁡Llikelihood=∑i=1Nlog⁡y^i,ki
\\log \\mathcal{L}_{\\mathrm{likelihood}} = \\sum_{i=1}^{N}\\log \\hat{y}_{i,k_i}
logLlikelihood=i=1Nlogy^i,ki

训练时我们希望最大化对数似然,这等价于最小化负对数似然:

−∑i=1Nlog⁡y^i,ki
-\\sum_{i=1}^{N}\\log \\hat{y}_{i,k_i}
i=1Nlogy^i,ki

再写成 one-hot 形式,就是:

−∑i=1N∑c=1Cyi,clog⁡y^i,c
-\\sum_{i=1}^{N}\\sum_{c=1}^{C} y_{i,c}\\log \\hat{y}_{i,c}
i=1Nc=1Cyi,clogy^i,c

这正是交叉熵损失。

所以,从统计意义上看,交叉熵损失就是分类任务在最大似然估计下得到的负对数似然形式。

(2)从信息论理解

在信息论里,交叉熵用来衡量两个概率分布之间的差异。对于真实分布 y\\mathbf{y}y 和预测分布 y^\\hat{\\mathbf{y}}y^,交叉熵定义为:

H(y,y^)=−∑c=1Cyclog⁡y^c
H(\\mathbf{y}, \\hat{\\mathbf{y}}) = -\\sum_{c=1}^{C} y_c \\log \\hat{y}_c
H(y,y^)=c=1Cyclogy^c

如果模型预测分布和真实分布越接近,交叉熵就越小;如果差得越远,交叉熵就越大。

在 one-hot 标签的分类问题里,真实分布其实就是一个“某一类概率为 1,其余为 0”的分布,因此交叉熵会特别关注真实类别那一项的预测概率。

4. 为什么真实类别概率越低,损失越大

交叉熵最核心的单样本形式是:

LCE=−log⁡y^k
\\mathcal{L}_{\\mathrm{CE}} = -\\log \\hat{y}_k
LCE=logy^k

这里的 y^k\\hat{y}_ky^k 是模型给真实类别分配的概率。

我们来看几个简单的数值:

  • 如果 y^k=0.9\\hat{y}_k = 0.9y^k=0.9,则损失约为 −log⁡0.9≈0.105-\\log 0.9 \\approx 0.105log0.90.105
  • 如果 y^k=0.5\\hat{y}_k = 0.5y^k=0.5,则损失约为 −log⁡0.5≈0.693-\\log 0.5 \\approx 0.693log0.50.693
  • 如果 y^k=0.1\\hat{y}_k = 0.1y^k=0.1,则损失约为 −log⁡0.1≈2.303-\\log 0.1 \\approx 2.303log0.12.303
  • 如果 y^k=0.01\\hat{y}_k = 0.01y^k=0.01,则损失约为 −log⁡0.01≈4.605-\\log 0.01 \\approx 4.605log0.014.605

可以看到:

  • 真实类别概率越接近 1,损失越小;
  • 真实类别概率越接近 0,损失增长非常快。

这正符合分类训练的需求。因为如果模型把正确类概率压得特别低,那就说明它犯了一个很严重的错误,损失理应给出更强的惩罚。

5. 梯度推导

交叉熵损失最常见的推导,是和 Softmax 一起推。

设模型输出 logits 为 z=[z1,…,zC]\\mathbf{z}=[z_1,\\dots,z_C]z=[z1,,zC],经过 Softmax 得到:

y^c=ezc∑j=1Cezj
\\hat{y}_c = \\frac{e^{z_c}}{\\sum_{j=1}^{C}e^{z_j}}
y^c=j=1Cezjezc

单个样本的交叉熵损失为:

LCE=−∑c=1Cyclog⁡y^c
\\mathcal{L}_{\\mathrm{CE}} = -\\sum_{c=1}^{C} y_c \\log \\hat{y}_c
LCE=c=1Cyclogy^c

将 Softmax 和交叉熵联立后,可以得到一个非常经典的结果:

∂LCE∂zc=y^c−yc
\\frac{\\partial \\mathcal{L}_{\\mathrm{CE}}}{\\partial z_c} = \\hat{y}_c – y_c
zcLCE=y^cyc

也就是说,交叉熵对 logits 的梯度,正好等于“预测概率减去真实分布”。

这个结果非常重要,因为它说明:

  • 对真实类别,如果预测概率不够高,梯度会推动它继续升高;
  • 对错误类别,如果预测概率太高,梯度会推动它降低。

对于一个 batch 的平均损失,再除以样本数 NNN 即可:

∂LCE∂zi,c=1N(y^i,c−yi,c)
\\frac{\\partial \\mathcal{L}_{\\mathrm{CE}}}{\\partial z_{i,c}} = \\frac{1}{N}(\\hat{y}_{i,c} – y_{i,c})
zi,cLCE=N1(y^i,cyi,c)

这个梯度形式非常简洁,也是交叉熵在分类问题中优化效果好的一个重要原因。

6. 函数特性

(1)直接面向概率分布

交叉熵不是在比较“类别编号差多少”,而是在比较预测概率分布和真实标签分布之间的差异,因此非常适合分类任务。

(2)对错误高置信预测惩罚很强

如果模型把错误类别预测得特别自信,也就是给真实类别分配了极低概率,那么交叉熵损失会迅速增大。这种性质能有效推动模型纠正严重错误。

(3)和 Softmax 配合非常自然

多分类模型通常都需要输出归一化概率分布,而 Softmax 正好能把 logits 转成概率,和交叉熵天然匹配。

(4)梯度形式简洁

Softmax 和交叉熵组合后的梯度可以简化成 y^c−yc\\hat{y}_c – y_cy^cyc,实现和优化都很方便。

7. 使用场景与局限性

Cross Entropy Loss 主要用于多分类任务,例如:

  • 图像分类
  • 文本分类
  • 语音分类
  • 故障类别识别
  • 损伤类型识别
  • 多类别目标识别

只要任务的目标是“从多个互斥类别中选出一个正确类别”,交叉熵基本都是首选。

优势
  • 非常适合分类概率建模;
  • 有明确的最大似然解释;
  • 和 Softmax 配合自然;
  • 对高置信错误惩罚明显;
  • 梯度形式简洁,训练效果通常稳定。
缺点
  • 对标签噪声比较敏感;
  • 当类别极不平衡时,普通交叉熵可能不够好用;
  • 默认认为类别之间是互斥的,不适合多标签任务;
  • 如果模型过度自信,可能带来概率校准问题。

8. Cross Entropy 与 MSE 的简单对比

有些初学者会问:分类任务能不能直接用 MSE?

理论上不是完全不行,但通常不推荐。原因主要有两点:

(1)目标形式不同

分类任务本质上要学的是概率分布,交叉熵直接衡量分布差异;MSE 更偏向数值误差,不够贴合分类目标。

(2)优化特性不同

在分类任务里,交叉熵和 Softmax 组合后的梯度通常更合理,训练效率和效果往往都比 MSE 更好。尤其在多分类问题中,这一点很明显。

所以一般来说:

  • 连续值预测优先考虑 MSE、MAE、Huber;
  • 分类任务优先考虑 Cross Entropy。

9. 代码实现

下面先用 NumPy 实现一个基础版本的多分类交叉熵损失。这里假设输入已经是 Softmax 之后的概率。

import numpy as np

def cross_entropy_loss(y_true, y_pred, eps=1e-12):
"""
计算多分类交叉熵损失

参数:
y_true — 真实标签,one-hot 形式,形状为 (N, C)
y_pred — 预测概率,形状为 (N, C)
eps — 防止 log(0) 的微小常数

返回:
平均交叉熵损失
"""
y_true = np.asarray(y_true)
y_pred = np.asarray(y_pred)

y_pred = np.clip(y_pred, eps, 1.0)
loss = np.sum(y_true * np.log(y_pred), axis=1)
return np.mean(loss)

如果想自己把 logits 转成概率,可以先写一个 Softmax:

import numpy as np

def softmax(logits):
"""
计算 Softmax 概率
"""

logits = np.asarray(logits)
shifted = logits np.max(logits, axis=1, keepdims=True)
exp_logits = np.exp(shifted)
return exp_logits / np.sum(exp_logits, axis=1, keepdims=True)

示例使用

logits = np.array([
[2.0, 1.0, 0.1],
[0.5, 2.5, 0.3]
])

y_true = np.array([
[1, 0, 0],
[0, 1, 0]
])

y_pred = softmax(logits)
loss = cross_entropy_loss(y_true, y_pred)

print("Predicted Probabilities:")
print(y_pred)
print("Cross Entropy Loss:", loss)

10. PyTorch中的实现

在 PyTorch 中,多分类通常直接使用 nn.CrossEntropyLoss()。

需要特别注意的是:它的输入应该是 logits,而不是 Softmax 之后的概率。

import torch
import torch.nn as nn

criterion = nn.CrossEntropyLoss()

logits = torch.tensor([
[2.0, 1.0, 0.1],
[0.5, 2.5, 0.3]
], dtype=torch.float32)

labels = torch.tensor([0, 1], dtype=torch.long)

loss = criterion(logits, labels)
print("Cross Entropy Loss:", loss.item())

如果放到模型训练中,一般写法如下:

import torch
import torch.nn as nn
import torch.optim as optim

model = nn.Linear(10, 3)

criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(model.parameters(), lr=0.01)

x = torch.randn(4, 10)
y = torch.tensor([0, 2, 1, 0], dtype=torch.long)

logits = model(x)
loss = criterion(logits, y)

optimizer.zero_grad()
loss.backward()
optimizer.step()

print("Loss:", loss.item())

11. 一个常见坑:不要先做 Softmax 再送进 CrossEntropyLoss

这个问题非常常见,单独说一下。

很多人会写成这样:

probs = torch.softmax(logits, dim=1)
loss = criterion(probs, labels)

这种写法通常是不对的。因为 nn.CrossEntropyLoss() 内部已经自动做了 log_softmax,如果你在外面先做一次 Softmax,相当于重复处理了一遍,不仅数值稳定性更差,还可能影响训练效果。

正确写法是:

loss = criterion(logits, labels)

也就是说:

  • 模型最后一层输出 logits;
  • 不手动做 Softmax;
  • 直接送给 CrossEntropyLoss。

12. 一个简单理解

交叉熵损失可以简单理解为:

模型对正确答案越有把握,损失越小;模型对正确答案越没把握,损失越大。

它真正关心的,不是类别编号本身,而是你给正确类别分配了多大的概率。

所以它特别适合分类问题,因为分类任务本质上就是在学一个“把概率尽量分对”的过程。

三、小结

Cross Entropy Loss 是分类任务中最核心、最常用的损失函数之一。它既有明确的概率解释,也有很好的优化性质,在多分类问题中几乎是默认选择。

它的核心思想其实很简单:提高真实类别的预测概率,压低错误类别的预测概率。
再配合 Softmax 使用,整个分类训练过程就会变得很自然。

四、参考文献

  • Bishop, C. M. Pattern Recognition and Machine Learning. Springer, 2006.
  • Goodfellow, I., Bengio, Y., & Courville, A. Deep Learning. MIT Press, 2016.
  • Murphy, K. P. Machine Learning: A Probabilistic Perspective. MIT Press, 2012.
  • PyTorch Documentation. torch.nn.CrossEntropyLoss.
赞(0)
未经允许不得转载:171主机测评 » 损失函数大汇总(四)(Cross Entropy Loss附公式推导与代码)
分享到: 更多 (0)

评论 抢沙发

  • 昵称 (必填)
  • 邮箱 (必填)
  • 网址