news 2026/8/27 11:09:54

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
从信息熵到交叉熵损失: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。

为什么是概率的对数和负号?我们来拆解一下:

  1. 单个事件的信息量:一个概率为p的事件发生,它所携带的信息量定义为I(p) = -log(p)。概率p越小(事件越罕见),发生时带来的“惊喜度”或信息量-log(p)就越大。比如中彩票(概率极小)的信息量远大于吃饭(概率极大)。
  2. 熵是信息量的期望:熵不是针对一次具体事件,而是针对整个概率分布所有可能事件的信息量,按照其概率加权平均(求期望)。所以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])。显然,pq不一样。我们如何量化这个“不一样”的程度呢?这就是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)]

这个公式可以解读为:

  1. p_i * log(p_i):基于真实分布p,描述事件i所需的平均信息量(即信息熵H(p)的一部分)。
  2. p_i * log(q_i):基于我们的近似分布q,去描述(或编码)真实发生的事件i所需的平均信息量。
  3. 两者相减(用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(f"KL(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(f"KL(p_true || q_perfect) = {kl_perfect:.4f}") # 应该为0(或一个极小的数,由于数值精度)

你会发现,当预测完全正确时,KL散度为0。预测越不准,KL散度值越大。

注意:KL散度有一个重要特性——当p_i > 0q_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的输入有两部分:

  1. input:模型的原始输出(raw scores, logits)。注意,是Softmax之前的logits!它的形状通常是(N, C),其中N是批次大小(batch size),C是类别数。对于更高维的数据(如图像分割),可能是(N, C, H, W)
  2. target:真实标签。它的形状是(N,),每个元素是类别索引(范围在[0, C-1])。对于图像分割等高维任务,形状是(N, H, W)

关键点CrossEntropyLoss内部已经组合了LogSoftmaxNLLLoss(负对数似然损失)。所以你不需要在模型最后一层手动添加Softmax激活函数。如果你加了,反而可能因为数值计算链(Softmax后接LogSoftmax)导致数值不稳定或梯度问题。

一个标准的流程应该是:

import torch import torch.nn as nn # 假设一个简单的分类模型 class SimpleClassifier(nn.Module): def __init__(self, input_dim=784, num_classes=10): 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, dim=1, keepdim=True) # 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_grad=True) 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(f"PyTorch损失: {loss_torch.item():.6f}") print(f"两者是否接近: {torch.allclose(loss_manual, loss_torch)}")

通过手动实现,你可以清晰地看到,损失函数的核心就是取出真实类别对应的LogSoftmax值,然后取负号。这也解释了为什么它叫“交叉熵”,虽然形式上只关注了一个点,但其数学本质与完整的交叉熵定义在one-hot标签下是等价的。

4.3 常见陷阱与调试技巧

  1. 标签越界(Index out of range):这是最常见的运行时错误。确保你的labels张量中的每一个值都严格在[0, C-1]范围内,其中C是你的logits的第二个维度(类别数)。如果你的数据集标签是从1开始的,需要先转换为从0开始。

    # 错误示例:标签为10,但类别数C=10(有效索引0-9) labels = torch.tensor([0, 5, 10, 3]) # 索引10会报错 # 修正:检查并转换标签范围 assert labels.max() < num_classes and labels.min() >= 0, "标签越界!"
  2. 数值稳定性与半精度训练:当使用混合精度训练(AMP)时,logits可能是float16类型。Softmax/LogSoftmax在float16下对于极大或极小的输入值更容易溢出或下溢。PyTorch的CrossEntropyLoss内部已经做了一些数值稳定化处理,但如果你在自定义损失函数时需要自己计算,务必使用F.log_softmaxtorch.log_softmax而不是先softmaxlog

    提示F.log_softmax在实现上使用了“Log-Sum-Exp Trick”,这是一种数值稳定的计算方法,可以有效避免exp(x)过大导致的溢出。其核心思想是:log(∑exp(x_i)) = max(x) + log(∑exp(x_i - max(x)))。通过减去最大值,将指数函数的参数范围控制住。

  3. 忽略的ignore_index参数:在处理像自然语言处理(NLP)中填充符(Padding)时,或者分割任务中需要忽略的特定类别(如背景),CrossEntropyLoss提供了一个非常实用的ignore_index参数。

    criterion = nn.CrossEntropyLoss(ignore_index=-100) # 假设 labels 中值为 -100 的位置是需要忽略的 loss = criterion(logits, labels) # 这些位置不参与损失计算和梯度回传

    这比在计算损失前手动过滤要方便和高效得多。

  4. 类别不平衡与权重设置:当你的训练数据中各类别样本数差异巨大时,直接使用标准交叉熵损失会导致模型偏向于样本多的类别。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(weight=weights)

    设置权重后,损失函数会为每个样本的损失乘以对应类别的权重,从而让模型更关注样本少的类别。

4.4 在多标签分类与二分类任务中的应用

标准的CrossEntropyLoss适用于单标签多分类(每个样本只属于一个类别)。对于其他任务:

  • 二分类任务:虽然可以使用CrossEntropyLoss(此时C=2),但更常见、更数值稳定的是使用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_size=32, shuffle=True) # 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(), lr=0.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(f"Epoch [{epoch+1}/{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, dim=1) print(f"预测类别索引: {predicted_class.item()}") # 方法2:如果需要概率值,则手动应用Softmax probabilities = torch.softmax(logits, dim=1) print(f"各类别概率: {probabilities}") print(f"概率最大的类别: {torch.argmax(probabilities, dim=1).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, smoothing=0.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(smoothing=0.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, alpha=1, gamma=2, reduction='mean'): super().__init__() self.alpha = alpha self.gamma = gamma self.reduction = reduction def forward(self, logits, targets): ce_loss = nn.functional.cross_entropy(logits, targets, reduction='none') 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_loss

6.3 与负对数似然损失(NLLLoss)的关系

我们之前提到,CrossEntropyLoss=LogSoftmax+NLLLossNLLLoss的输入是对数概率(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(dim=1)(logits) loss_nll = nn.NLLLoss()(log_softmax, labels) print(torch.allclose(loss_ce, loss_nll)) # 输出应为 True

理解这种等价关系有助于你在需要自定义概率变换时(例如使用温度缩放进行模型校准),能灵活地组合不同的模块。

7. 总结与核心要点回顾

我们从信息论最基本的概念——信息熵出发,一步步推导出KL散度和交叉熵,最终落地到PyTorch中无处不在的CrossEntropyLoss。这个过程不是枯燥的数学之旅,而是一连串解决实际问题的思想结晶。

核心逻辑链再梳理

  1. 信息熵 H(p):描述一个分布自身的不确定性,是编码该分布所需信息量的下界。
  2. KL散度 KL(p||q):衡量用分布q去近似真实分布p时,产生的额外信息损失。它不对称,不是距离,但能有效衡量差异。
  3. 交叉熵 H(p, q):等于H(p) + KL(p||q)。由于H(p)是常数,最小化交叉熵就等价于最小化KL散度
  4. 在分类任务中:真实分布p是one-hot编码,交叉熵简化为-log(q_t),即模型对真实类别预测概率的负对数。
  5. PyTorch实现nn.CrossEntropyLoss接收logits和类别索引标签,内部高效、稳定地完成了LogSoftmax + NLLLoss的计算。

最重要的几个实战心得

  • 永远记住:传给CrossEntropyLoss的是logits(Softmax前的原始分数),不是概率。不要在模型最后一层加Softmax。
  • 仔细检查形状logits(N, C)labels(N,)。标签值必须在[0, C-1]范围内。
  • 理解你的任务:单标签分类用CrossEntropyLoss,多标签分类或二分类用BCEWithLogitsLoss
  • 善用高级参数:面对类别不平衡,使用weight参数;需要忽略某些标签(如padding),使用ignore_index参数。
  • 进阶优化:在追求更高性能时,可以考虑标签平滑来提升泛化能力,或在极度不平衡的任务中尝试Focal Loss。

最后,损失函数不仅仅是代码里的一行,它是连接模型输出与学习目标的桥梁,是将抽象的“学好”这一目标,转化为具体的、可优化的数学语言的关键。理解交叉熵背后的信息论原理,能让你在调试模型、设计损失函数、甚至理解模型行为时,拥有更深刻的洞察力,而不是仅仅把它当作一个黑盒调用。下次当你写下criterion = nn.CrossEntropyLoss()时,希望你能会心一笑,知道这行简洁代码背后,承载着从香农开始,关于信息、不确定性和学习本质的深刻思考。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/27 11:07:08

PHP报名系统源码,自带支付功能,这也太牛了吧

有许多的表单系统, 其中PHP万能表单系统, 因具备灵活性以及可扩展性, 从而备受开发者的青睐。PHP万能表单系统是一款基于PHP语言的表单生成器, 它能够助力开发者迅速生成各类表单, 像是注册表单、登录表单、留言表单等。接下来给大家分享一款php万能表单系统源码, 它支持自定义…

作者头像 李华
网站建设 2026/8/27 11:04:38

企业获客软件免安装能力横评:网页端和小程序端实测数据

从工具部署门槛的角度&#xff0c;对9款企业获客软件的网页端和小程序端做了一次实测对比。前段时间换电脑&#xff0c;之前装的几个拓客软件客户端全得重新下载安装&#xff0c;光是找安装包、装驱动、等审批权限就折腾了大半天&#xff0c;还差点因为公司电脑限制装软件的政策…

作者头像 李华
网站建设 2026/8/27 11:03:04

基于Matlab的气候变化影响评估:从数学建模到风险预测实战

1. 项目概述&#xff1a;当数学建模遇上气候变化 最近几年&#xff0c;不管是看新闻还是身边朋友聊天&#xff0c;气候变化这个话题出现的频率越来越高。从极端高温、暴雨洪涝&#xff0c;到冰川消融、海平面上升&#xff0c;这些现象不再是遥远的科学报告&#xff0c;而是真切…

作者头像 李华
网站建设 2026/8/27 11:02:47

商业数据分析学习路线:从Excel到Python的完整实战指南

商业数据分析这几年已经从“加分项”变成了很多岗位的“基础项”。不管是产品、运营、销售&#xff0c;还是财务、人力、审计&#xff0c;日常工作里都逃不开看数据、拉报表、找问题、给建议。很多人也买过课、存过资料&#xff0c;但真正能从头到尾把分析思路和工具链打通的人…

作者头像 李华
网站建设 2026/8/27 11:00:45

Buck电路EMI难搞?从物理本质到PCB布局的SWIFT实战解析

1. 为什么 Buck 的 EMI 这么难搞搞过电源设计的人应该都有体会&#xff0c;Buck 降压电路的原理看起来简单到可以用一句话讲完——开关管先导通储能、再关断续流&#xff0c;配合 LC 滤波把方波变成直流——但一到 EMC 测试实验室&#xff0c;三米法暗室一进去&#xff0c;频谱…

作者头像 李华