从信息熵到交叉熵损失:PyTorch分类任务核心原理与实战

📅 2026/8/27 11:10:03
从信息熵到交叉熵损失:PyTorch分类任务核心原理与实战
1. 从“不确定性”到“信息熵”一个直觉化的理解在机器学习和深度学习的领域里我们经常听到“交叉熵损失函数”这个词尤其是在使用PyTorch、TensorFlow这类框架做分类任务时它几乎是标配。但如果你只是简单地调用nn.CrossEntropyLoss()然后看着损失值下降可能并没有真正理解这个损失函数背后蕴含的深刻思想。它不是一个凭空捏造的数学公式而是建立在信息论坚实基石上的一个优雅工具。今天我们不谈空洞的公式推导就从最朴素的直觉出发聊聊信息熵、KL散度、交叉熵这三兄弟以及它们是如何最终化身为我们代码里那行简洁的loss criterion(outputs, labels)的。想象一下你是一个天气预报员。对于明天的天气如果你在一个四季如春、极少下雨的城市比如传说中的某个“春城”你几乎可以百分百确定地预报“晴天”。这个预报带来的“信息量”大吗不大因为结果几乎是确定的你的预报没有消除什么不确定性。但如果你在一个天气变化莫测的沿海城市你艰难地给出了“50%概率晴天50%概率暴雨”的预报。当第二天实际是暴雨时这个预报带来的“信息量”感觉上就更大一些因为它帮你理解了一个原本很不确定的事件。信息熵本质上就是衡量这个“不确定性”的数学工具。一个系统比如天气越不确定、越混乱它的信息熵就越高。那么如何量化这种直觉呢克劳德·香农给了我们一个惊艳的定义。对于一个离散随机变量X它有n种可能的状态每个状态i发生的概率是p_i。那么它的信息熵H(X)定义为H(X) - Σ (p_i * log(p_i)) 其中求和 i 从1到n。为什么是概率的对数和负号我们来拆解一下单个事件的信息量一个概率为p的事件发生它所携带的信息量定义为I(p) -log(p)。概率p越小事件越罕见发生时带来的“惊喜度”或信息量-log(p)就越大。比如中彩票概率极小的信息量远大于吃饭概率极大。熵是信息量的期望熵不是针对一次具体事件而是针对整个概率分布所有可能事件的信息量按照其概率加权平均求期望。所以H(X) Σ p_i * I(p_i) Σ p_i * (-log(p_i)) - Σ p_i * log(p_i)。所以信息熵H(X)衡量的是基于真实概率分布p去描述或编码随机变量X所需的平均信息量比特数。它只依赖于分布p本身。当所有事件等概率发生时最混乱最不确定熵最大当某个事件概率为1完全确定熵为0。在PyTorch里虽然没有直接计算一个分布信息熵的单一函数但我们可以轻松实现import torch def entropy(p): # p是一个概率分布张量例如torch.tensor([0.1, 0.2, 0.7]) # 确保概率和为1且避免log(0)的情况 p p 1e-12 # 添加一个极小值防止数值问题 return -torch.sum(p * torch.log(p)) # 示例一个三分类的预测概率 p_dist torch.tensor([0.8, 0.1, 0.1]) print(f分布 {p_dist} 的信息熵为: {entropy(p_dist):.4f}) # 输出较低因为分布较确定 p_uniform torch.tensor([0.333, 0.333, 0.334]) print(f均匀分布 {p_uniform} 的信息熵为: {entropy(p_uniform):.4f}) # 输出较高接近最大值理解信息熵是第一步它为我们设定了一个基准描述事物本身的不确定性需要多少“成本”。2. KL散度衡量两个概率分布间的“距离”或差异现在场景升级了。你作为天气预报员经过多年观察心里有一个关于本地天气的“真实”概率分布p比如[晴:0.6, 雨:0.3, 阴:0.1]。但为了简化模型或者因为数据来源不同你实际使用的、对外发布的预报分布是q比如[晴:0.7, 雨:0.2, 阴:0.1]。显然p和q不一样。我们如何量化这个“不一样”的程度呢这就是KL散度Kullback-Leibler Divergence要解决的问题。KL散度的定义直接揭示了它的含义KL(p || q) Σ p_i * log(p_i / q_i) Σ [p_i * log(p_i) - p_i * log(q_i)]这个公式可以解读为p_i * log(p_i)基于真实分布p描述事件i所需的平均信息量即信息熵H(p)的一部分。p_i * log(q_i)基于我们的近似分布q去描述或编码真实发生的事件i所需的平均信息量。两者相减(用q编码p的成本) - (用p编码p的理想成本)再对所有事件i求期望用p加权就得到了因为使用近似分布q而不是真实分布p所导致的额外信息成本或效率损失。所以KL(p||q)衡量的是当真实分布为p时用分布q去近似p所产生的不必要的额外信息损失。它不是一个真正的“距离”因为不对称KL(p||q) ≠ KL(q||p)但它完美地刻画了两个分布的差异。在机器学习的语境下p通常是数据的真实分布例如一个样本的真实标签是狗那么它的分布就是[狗:1, 猫:0, 鸟:0]即one-hot编码而q是我们的模型预测出的概率分布例如[狗:0.8, 猫:0.15, 鸟:0.05]。我们的目标就是让模型的预测分布q无限接近真实分布p也就是最小化它们之间的KL散度。在PyTorch中我们可以手动计算KL散度来加深理解def kl_divergence(p, q): # p, q 是两个概率分布张量 # 注意p和q需要满足概率分布的性质和为1非负 p p 1e-12 q q 1e-12 return torch.sum(p * torch.log(p / q)) # 示例真实分布p (one-hot) 和 模型预测q p_true torch.tensor([1.0, 0.0, 0.0]) # 真实标签是第0类 q_pred torch.tensor([0.7, 0.2, 0.1]) # 模型预测 kl kl_divergence(p_true, q_pred) print(fKL(p_true || q_pred) {kl:.4f}) # 如果预测完全正确 q_perfect torch.tensor([1.0, 0.0, 0.0]) kl_perfect kl_divergence(p_true, q_perfect) print(fKL(p_true || q_perfect) {kl_perfect:.4f}) # 应该为0或一个极小的数由于数值精度你会发现当预测完全正确时KL散度为0。预测越不准KL散度值越大。注意KL散度有一个重要特性——当p_i 0而q_i 0时log(p_i / q_i)会趋于无穷大。这意味着如果你的模型给真实类别分配了0概率即完全不相信正确答案那么损失会变得无穷大这在训练中会导致梯度爆炸。这解释了为什么在分类问题中我们通常使用Softmax函数将模型输出转换为概率并且要避免极端的概率值例如通过标签平滑技术。3. 交叉熵KL散度的“亲兄弟”与损失函数的直接形态我们回到KL散度的公式KL(p || q) Σ p_i * log(p_i) - Σ p_i * log(q_i)。仔细观察等号右边第一项Σ p_i * log(p_i)是什么这正是我们第一部分讲到的真实分布p的信息熵 H(p)。它是一个只与真实分布有关的常数与我们的模型q无关。等号右边第二项- Σ p_i * log(q_i) 被单独拿出来定义就是交叉熵Cross-Entropy记作H(p, q)。所以我们有这样一个关键等式KL(p || q) H(p, q) - H(p)移项得到H(p, q) H(p) KL(p || q)这个等式意义重大交叉熵 H(p, q) 等于真实分布的信息熵 H(p) 加上两个分布的KL散度。由于H(p)是固定常数那么最小化交叉熵 H(p, q)就等价于最小化KL散度 KL(p || q)这就是为什么交叉熵能作为损失函数的核心原因——它直接驱动模型分布q去逼近真实分布p并且计算形式比KL散度更简洁少了一项。现在我们把场景具体到分类任务。对于单个样本真实分布p通常是one-hot编码。例如对于一个三分类问题真实标签是第2类则p [0, 0, 1]。模型预测分布q是模型最后一层通常是线性层经过Softmax函数后的输出例如q [0.1, 0.2, 0.7]。那么这个样本的交叉熵损失为H(p, q) - Σ p_i * log(q_i)由于p是one-hot的只有真实类别索引t对应的p_t 1其他都为0。因此求和公式瞬间简化H(p, q) - 1 * log(q_t) -log(q_t)看这就是我们最熟悉的那个形式交叉熵损失就是模型对真实类别所预测概率的负对数模型对真实类别的预测概率q_t越高越接近1-log(q_t)就越小因为log(1)0损失就越低。反之如果模型预测真实类别的概率很低-log(q_t)就会很大给予模型很大的惩罚。这个形式极其优雅且实用因为它避免了计算整个KL散度或完整的交叉熵求和只需要关注真实类别对应的预测概率即可。4. PyTorch中的CrossEntropyLoss细节、陷阱与最佳实践理解了理论我们来看实践。PyTorch中的torch.nn.CrossEntropyLoss是使用最广泛的损失函数之一但它有一些“沉默的约定”和容易踩坑的地方。4.1 输入与输出的形状约定这是新手最容易出错的地方。nn.CrossEntropyLoss的输入有两部分input模型的原始输出raw scores, logits。注意是Softmax之前的logits它的形状通常是(N, C)其中N是批次大小batch sizeC是类别数。对于更高维的数据如图像分割可能是(N, C, H, W)。target真实标签。它的形状是(N,)每个元素是类别索引范围在[0, C-1]。对于图像分割等高维任务形状是(N, H, W)。关键点CrossEntropyLoss内部已经组合了LogSoftmax和NLLLoss负对数似然损失。所以你不需要在模型最后一层手动添加Softmax激活函数。如果你加了反而可能因为数值计算链Softmax后接LogSoftmax导致数值不稳定或梯度问题。一个标准的流程应该是import torch import torch.nn as nn # 假设一个简单的分类模型 class SimpleClassifier(nn.Module): def __init__(self, input_dim784, num_classes10): super().__init__() self.linear nn.Linear(input_dim, num_classes) # 输出logits没有Softmax def forward(self, x): return self.linear(x) # 直接返回logits model SimpleClassifier() criterion nn.CrossEntropyLoss() # 定义损失函数 # 模拟一个batch的数据 batch_size 4 num_classes 10 logits model(torch.randn(batch_size, 784)) # logits形状: (4, 10) labels torch.randint(0, num_classes, (batch_size,)) # 标签形状: (4,) loss criterion(logits, labels) # 正确用法 print(loss)4.2 内部计算过程拆解为了更透彻地理解我们可以手动拆解CrossEntropyLoss的计算步骤并与直接调用进行对比# 手动计算交叉熵损失以验证和理解 def manual_cross_entropy(logits, labels): logits: (N, C) labels: (N,) # Step 1: 对logits应用LogSoftmax # LogSoftmax(x_i) x_i - log(∑exp(x_j)) 更数值稳定 log_softmax logits - torch.logsumexp(logits, dim1, keepdimTrue) # Step 2: 根据labels选取对应类别的log概率 # 等价于 NLLLoss nll_loss -log_softmax[range(len(labels)), labels] # Step 3: 对batch求平均 return torch.mean(nll_loss) # 对比 logits torch.randn(4, 10, requires_gradTrue) labels torch.tensor([2, 5, 1, 9]) loss_manual manual_cross_entropy(logits, labels) loss_torch nn.CrossEntropyLoss()(logits, labels) print(f手动计算损失: {loss_manual.item():.6f}) print(fPyTorch损失: {loss_torch.item():.6f}) print(f两者是否接近: {torch.allclose(loss_manual, loss_torch)})通过手动实现你可以清晰地看到损失函数的核心就是取出真实类别对应的LogSoftmax值然后取负号。这也解释了为什么它叫“交叉熵”虽然形式上只关注了一个点但其数学本质与完整的交叉熵定义在one-hot标签下是等价的。4.3 常见陷阱与调试技巧标签越界Index out of range这是最常见的运行时错误。确保你的labels张量中的每一个值都严格在[0, C-1]范围内其中C是你的logits的第二个维度类别数。如果你的数据集标签是从1开始的需要先转换为从0开始。# 错误示例标签为10但类别数C10有效索引0-9 labels torch.tensor([0, 5, 10, 3]) # 索引10会报错 # 修正检查并转换标签范围 assert labels.max() num_classes and labels.min() 0, 标签越界数值稳定性与半精度训练当使用混合精度训练AMP时logits可能是float16类型。Softmax/LogSoftmax在float16下对于极大或极小的输入值更容易溢出或下溢。PyTorch的CrossEntropyLoss内部已经做了一些数值稳定化处理但如果你在自定义损失函数时需要自己计算务必使用F.log_softmax或torch.log_softmax而不是先softmax再log。提示F.log_softmax在实现上使用了“Log-Sum-Exp Trick”这是一种数值稳定的计算方法可以有效避免exp(x)过大导致的溢出。其核心思想是log(∑exp(x_i)) max(x) log(∑exp(x_i - max(x)))。通过减去最大值将指数函数的参数范围控制住。忽略的ignore_index参数在处理像自然语言处理NLP中填充符Padding时或者分割任务中需要忽略的特定类别如背景CrossEntropyLoss提供了一个非常实用的ignore_index参数。criterion nn.CrossEntropyLoss(ignore_index-100) # 假设 labels 中值为 -100 的位置是需要忽略的 loss criterion(logits, labels) # 这些位置不参与损失计算和梯度回传这比在计算损失前手动过滤要方便和高效得多。类别不平衡与权重设置当你的训练数据中各类别样本数差异巨大时直接使用标准交叉熵损失会导致模型偏向于样本多的类别。CrossEntropyLoss提供了weight参数来解决这个问题。# 假设我们有3个类别样本数比例为 class0: 50%, class1: 30%, class2: 20% # 一种常见的权重设置是类别的倒数或者更常用的“逆频率” class_counts torch.tensor([500, 300, 200]) # 各类别样本数 weights 1.0 / class_counts # 逆频率 weights weights / weights.sum() # 可选归一化 # 或者使用 median frequency balancing: weight median_freq / class_freq criterion nn.CrossEntropyLoss(weightweights)设置权重后损失函数会为每个样本的损失乘以对应类别的权重从而让模型更关注样本少的类别。4.4 在多标签分类与二分类任务中的应用标准的CrossEntropyLoss适用于单标签多分类每个样本只属于一个类别。对于其他任务二分类任务虽然可以使用CrossEntropyLoss此时C2但更常见、更数值稳定的是使用nn.BCEWithLogitsLoss二元交叉熵损失。它将二分类视为两个独立的伯努利分布模型的最终输出是一个在0到1之间的概率通常用Sigmoid激活。BCEWithLogitsLoss内部集成了Sigmoid和BCE损失避免了数值问题。# 二分类任务 bce_criterion nn.BCEWithLogitsLoss() # 模型输出一个值 (N, 1) 或 (N,) logits model(x) labels labels.float() # 标签需要是float类型值为0或1 loss bce_criterion(logits.squeeze(), labels)多标签分类任务一个样本可以同时属于多个类别。此时应使用nn.BCEWithLogitsLoss并将模型的输出通道数设置为类别数C对每个通道独立地应用Sigmoid和二元交叉熵损失。# 多标签分类C个类别 bce_criterion nn.BCEWithLogitsLoss() # 模型输出形状 (N, C) logits model(x) # 标签形状 (N, C)每个位置是0或1 loss bce_criterion(logits, labels)理解这些区别至关重要错误地使用损失函数会导致模型无法正常学习到有效的特征。5. 从理论到实战一个完整的图像分类训练循环剖析让我们将所有知识点串联起来通过一个简化的图像分类训练循环看看交叉熵损失是如何在PyTorch中实际运作的。import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader, TensorDataset # 1. 准备模拟数据 num_samples 1000 num_features 784 # 例如28x28图像展平 num_classes 10 # 模拟logits和标签 X torch.randn(num_samples, num_features) # 生成模拟标签这里我们简单随机生成真实场景中来自数据集 y torch.randint(0, num_classes, (num_samples,)) dataset TensorDataset(X, y) dataloader DataLoader(dataset, batch_size32, shuffleTrue) # 2. 定义模型省略了复杂的网络结构仅用线性层示意 class TinyModel(nn.Module): def __init__(self): super().__init__() self.fc nn.Linear(num_features, num_classes) # 输出logits def forward(self, x): return self.fc(x) model TinyModel() # 3. 定义损失函数和优化器 criterion nn.CrossEntropyLoss() # 核心角色登场 optimizer optim.SGD(model.parameters(), lr0.01) # 4. 训练循环 num_epochs 5 for epoch in range(num_epochs): running_loss 0.0 for batch_idx, (inputs, labels) in enumerate(dataloader): # 清零梯度 optimizer.zero_grad() # 前向传播模型输出logits outputs model(inputs) # outputs形状: (batch_size, 10) # 计算损失CrossEntropyLoss内部进行LogSoftmax NLLLoss loss criterion(outputs, labels) # 反向传播 loss.backward() # 参数更新 optimizer.step() running_loss loss.item() avg_loss running_loss / len(dataloader) print(fEpoch [{epoch1}/{num_epochs}], Average Loss: {avg_loss:.4f}) # 5. 推理时的处理 model.eval() with torch.no_grad(): test_input torch.randn(1, num_features) logits model(test_input) # 方法1直接取最大logits的索引作为预测类别因为argmax在logits和softmax后结果一致 predicted_class torch.argmax(logits, dim1) print(f预测类别索引: {predicted_class.item()}) # 方法2如果需要概率值则手动应用Softmax probabilities torch.softmax(logits, dim1) print(f各类别概率: {probabilities}) print(f概率最大的类别: {torch.argmax(probabilities, dim1).item()})在这个循环中交叉熵损失函数扮演了“教练”的角色。它接收模型“猜”出的分数logits和标准答案labels计算出一个标量损失值。这个损失值衡量了当前模型预测的“糟糕”程度。通过loss.backward()这个“糟糕程度”被转化为每个模型参数的梯度即每个参数应该向哪个方向、以多大的幅度调整才能降低损失。优化器optimizer.step()则根据这些梯度实际更新参数。一个重要的实战细节在推理Inference阶段我们通常不需要计算损失也不需要显式调用Softmax来获得概率除非你需要概率值进行后续分析如计算置信度、模型校准或集成。因为对于分类任务我们只关心最大概率对应的类别而argmax(logits)和argmax(softmax(logits))的结果是完全相同的Softmax是单调函数不改变大小顺序。因此在部署模型时为了提升效率可以省去Softmax层。6. 超越基础交叉熵的变体与相关损失函数理解了标准的交叉熵损失后你会发现它在很多场景下被调整和扩展以适应更复杂的需求。6.1 标签平滑Label Smoothing标准交叉熵损失使用one-hot标签这会导致模型对真实类别的预测概率过度自信趋向于1可能会降低模型的泛化能力并使其对对抗样本更敏感。标签平滑通过将真实标签的1“分摊”一点给其他类别来缓解这个问题。PyTorch的CrossEntropyLoss本身不支持标签平滑但可以很容易地通过自定义损失函数或修改标签来实现class LabelSmoothingCrossEntropy(nn.Module): def __init__(self, smoothing0.1): super().__init__() self.smoothing smoothing def forward(self, logits, targets): num_classes logits.size(-1) # 将one-hot标签转换为平滑后的分布 with torch.no_grad(): targets torch.zeros_like(logits).scatter_(1, targets.unsqueeze(1), 1) targets targets * (1 - self.smoothing) self.smoothing / num_classes # 计算交叉熵 log_probs torch.log_softmax(logits, dim-1) loss -torch.sum(targets * log_probs, dim-1).mean() return loss # 使用方式 criterion LabelSmoothingCrossEntropy(smoothing0.1) loss criterion(logits, labels)标签平滑相当于在训练中加入了正则化告诉模型“正确答案很可能就是这个但其他答案也有一点点可能”这通常能带来轻微但稳定的性能提升尤其是在防止过拟合方面。6.2 Focal Loss在目标检测等领域前景和背景类别极度不平衡一张图中背景像素远多于目标像素。标准交叉熵损失会被大量简单的负样本背景主导导致模型难以学习难分的样本前景或模糊的目标。Focal Loss通过降低简单样本对损失的贡献让模型更关注难分样本。其核心是在标准交叉熵损失上增加了一个调制因子(1 - p_t)^γFL(p_t) -α_t * (1 - p_t)^γ * log(p_t)其中p_t是模型对真实类别的预测概率γ是聚焦参数通常0α_t是类别平衡权重。当样本被正确分类且p_t很大时(1 - p_t)^γ很小该样本的损失被大幅降低。当样本被错分或p_t很小时调制因子接近1损失基本不受影响。这样训练就聚焦在了那些难分的样本上。class FocalLoss(nn.Module): def __init__(self, alpha1, gamma2, reductionmean): super().__init__() self.alpha alpha self.gamma gamma self.reduction reduction def forward(self, logits, targets): ce_loss nn.functional.cross_entropy(logits, targets, reductionnone) p_t torch.exp(-ce_loss) # p_t exp(-CE) 预测概率 focal_loss self.alpha * (1 - p_t) ** self.gamma * ce_loss if self.reduction mean: return focal_loss.mean() elif self.reduction sum: return focal_loss.sum() else: return focal_loss6.3 与负对数似然损失NLLLoss的关系我们之前提到CrossEntropyLossLogSoftmaxNLLLoss。NLLLoss的输入是对数概率log-probabilities而不是logits。它的计算很简单loss -log_prob[target]的平均值。如果你在模型最后一层使用了nn.LogSoftmax那么就需要配合nn.NLLLoss使用。CrossEntropyLoss只是将这两步合并提供了更方便的接口。# 等价关系演示 logits torch.randn(4, 10) labels torch.tensor([1, 3, 5, 7]) # 方式1使用CrossEntropyLoss推荐 loss_ce nn.CrossEntropyLoss()(logits, labels) # 方式2手动分解为 LogSoftmax NLLLoss log_softmax nn.LogSoftmax(dim1)(logits) loss_nll nn.NLLLoss()(log_softmax, labels) print(torch.allclose(loss_ce, loss_nll)) # 输出应为 True理解这种等价关系有助于你在需要自定义概率变换时例如使用温度缩放进行模型校准能灵活地组合不同的模块。7. 总结与核心要点回顾我们从信息论最基本的概念——信息熵出发一步步推导出KL散度和交叉熵最终落地到PyTorch中无处不在的CrossEntropyLoss。这个过程不是枯燥的数学之旅而是一连串解决实际问题的思想结晶。核心逻辑链再梳理信息熵 H(p)描述一个分布自身的不确定性是编码该分布所需信息量的下界。KL散度 KL(p||q)衡量用分布q去近似真实分布p时产生的额外信息损失。它不对称不是距离但能有效衡量差异。交叉熵 H(p, q)等于H(p) KL(p||q)。由于H(p)是常数最小化交叉熵就等价于最小化KL散度。在分类任务中真实分布p是one-hot编码交叉熵简化为-log(q_t)即模型对真实类别预测概率的负对数。PyTorch实现nn.CrossEntropyLoss接收logits和类别索引标签内部高效、稳定地完成了LogSoftmax NLLLoss的计算。最重要的几个实战心得永远记住传给CrossEntropyLoss的是logitsSoftmax前的原始分数不是概率。不要在模型最后一层加Softmax。仔细检查形状logits是(N, C)labels是(N,)。标签值必须在[0, C-1]范围内。理解你的任务单标签分类用CrossEntropyLoss多标签分类或二分类用BCEWithLogitsLoss。善用高级参数面对类别不平衡使用weight参数需要忽略某些标签如padding使用ignore_index参数。进阶优化在追求更高性能时可以考虑标签平滑来提升泛化能力或在极度不平衡的任务中尝试Focal Loss。最后损失函数不仅仅是代码里的一行它是连接模型输出与学习目标的桥梁是将抽象的“学好”这一目标转化为具体的、可优化的数学语言的关键。理解交叉熵背后的信息论原理能让你在调试模型、设计损失函数、甚至理解模型行为时拥有更深刻的洞察力而不是仅仅把它当作一个黑盒调用。下次当你写下criterion nn.CrossEntropyLoss()时希望你能会心一笑知道这行简洁代码背后承载着从香农开始关于信息、不确定性和学习本质的深刻思考。