创见博客
交叉熵损失函数(Cross-Entropy Loss)
七崽爱吃小饼干2025/12/18阅读 3

交叉熵损失函数常用于多分类任务,本质是衡量「模型预测的概率分布」与「真实标签的概率分布」之间的差异,差异越大损失值越高,反之则越低。

在多分类任务中,我们会给每个类别打上一个互斥的标签。让模型经过推理得到,当前的输入样本和每个类别之间的分数也就是logits,分数越大说明越相似。

在训练过程中,模型输出的logits会被用来计算交叉熵损失,从而拉近模型预测的概率分布与真实标签的概率分布。

一、通用数学形式(理论基础)

交叉熵的本质是衡量两个离散概率分布 PPP(真实分布)和 QQQ(预测分布)的差异,公式为:

H(P,Q)=−∑i=1nP(i)⋅log⁡(Q(i))H(P,Q) = -\sum_{i=1}^n P(i) \cdot \log(Q(i))H(P,Q)=−i=1∑n​P(i)⋅log(Q(i))
  • P(i)P(i)P(i):真实分布中第 iii 个类别的概率;
  • Q(i)Q(i)Q(i):模型预测分布中第 iii 个类别的概率;
  • nnn:类别总数;
  • 对数底数通常取 eee(自然对数,对应 torch.log),也可取 2(信息论中比特单位),不影响优化趋势。

二、多分类任务简化公式(最常用)

在分类任务中,真实标签是 one-hot 分布(只有真实类别概率为 1,其余为 0),因此公式会大幅简化。

1. 单样本损失

假设样本的真实类别为 ccc,模型对该样本的预测概率分布为 Q=[q1,q2,...,qn]Q = [q_1, q_2, ..., q_n]Q=[q1​,q2​,...,qn​](由 logits 经 Softmax 得到),则单样本交叉熵损失为:

L=−log⁡(qc)L = -\log(q_c)L=−log(qc​)
  • 解释:one-hot 分布下只有 P(c)=1P(c)=1P(c)=1,其余 P(i)=0P(i)=0P(i)=0,求和项只剩 P(c)⋅log⁡(qc)P(c) \cdot \log(q_c)P(c)⋅log(qc​)。

为什么要取对数

我们对损失函数的要求应该是,当概率qcq_cqc​接近1,就说明模型推理越准确,损失应该更低、更接近0。当概率qcq_cqc​接近0,就说明模型推理越不准确,损失应该更大且最大值没有上限,接近正无穷。

我们知道概率qcq_cqc​的取值范围是[0,1],而log⁡\loglog在(0,1]取间单调递增,且取值范围是(-∞,0]。

  • 所以当qcq_cqc​的取值接近1,LLL的取值就接近0。
  • 而当qcq_cqc​的取值接近0,LLL的取值就越大,且逼近正无穷。
这也就满足了我们对损失函数的要求。

2. 批量样本平均损失

训练时通常计算一个 batch 内所有样本的平均损失,公式为:

Lbatch=−1N∑k=1Nlog⁡(qk,ck)L_{batch} = -\frac{1}{N} \sum_{k=1}^N \log(q_{k,c_k})Lbatch​=−N1​k=1∑N​log(qk,ck​​)
  • NNN:batch 内样本数量;
  • ckc_kck​:第 kkk 个样本的真实类别;
  • qk,ckq_{k,c_k}qk,ck​​:第 kkk 个样本预测为真实类别 ckc_kck​ 的概率。

3. PyTorch 对应实现

nn.CrossEntropyLoss 内置了 Softmax 激活,输入是模型输出的 logits(无需提前转概率),其等价计算过程为:

LCrossEntropy=−1N∑k=1N[zk,ck−log⁡(∑i=1nezk,i)]L_{CrossEntropy} = -\frac{1}{N} \sum_{k=1}^N \left[ z_{k,c_k} - \log\left( \sum_{i=1}^n e^{z_{k,i}} \right) \right]LCrossEntropy​=−N1​k=1∑N​[zk,ck​​−log(i=1∑n​ezk,i​)]
  • zk,iz_{k,i}zk,i​:第 kkk 个样本对应第 iii 类别的 logits;
  • 该公式直接通过 logits 计算,避免了先算 Softmax 导致的数值溢出问题。

三、带类别权重的扩展公式(解决类别不平衡)

当数据存在类别不平衡时,可给每个类别分配权重 wiw_iwi​,多分类损失公式扩展为:

Lweighted=−1N∑k=1Nwck⋅log⁡(qk,ck)L_{weighted} = -\frac{1}{N} \sum_{k=1}^N w_{c_k} \cdot \log(q_{k,c_k})Lweighted​=−N1​k=1∑N​wck​​⋅log(qk,ck​​)
  • wckw_{c_k}wck​​:第 kkk 个样本真实类别 ckc_kck​ 对应的权重;
  • PyTorch 中通过 nn.CrossEntropyLoss(weight=class_weights) 实现。

核心总结

场景公式形式关键特点
通用概率分布H(P,Q)=−∑P(i)log⁡Q(i)H(P,Q) = -\sum P(i)\log Q(i)H(P,Q)=−∑P(i)logQ(i)适用于任意两个概率分布
多分类(单样本)L=−log⁡(qc)L = -\log(q_c)L=−log(qc​)真实标签为 one-hot 分布
二分类(单样本)L=−[ylog⁡q+(1−y)log⁡(1−q)]L = -[y\log q + (1-y)\log(1-q)]L=−[ylogq+(1−y)log(1−q)]仅两个类别,标签为 0/1
带权重多分类L=−1N∑wcklog⁡qk,ckL = -\frac{1}{N}\sum w_{c_k}\log q_{k,c_k}L=−N1​∑wck​​logqk,ck​​解决类别不平衡问题
评论
0/100