
开头这几年带过不少做视觉、NLP方向的同学大家第一天接触PyTorch分类任务时几乎都会在同一个地方犯迷糊F.cross_entropy()到底要不要先接一个Softmax很多人上网一搜有说“cross_entropy里面已经自带softmax了你直接传logits就行”也有帖子说“我明明传了softmax之后的结果怎么loss是负的”。两种说法看似矛盾其实背后全是细节。这个函数是整个PyTorch里最常用、也最容易踩坑的损失函数之一搞懂它等于搞懂了分类任务的半条命脉。这篇文章不打算照抄官方文档我用一个偏实战的视角把这个函数从数学原理到代码实现、从参数含义到常见报错一层层拆干净。适合刚入门PyTorch的初学者也适合那些已经跑通代码但一直对“为什么这么设计”云里雾里的同学。1. 函数机制拆解它到底在算哪一步1.1 Logits 与 Softmax 输出的本质区别先明确一个概念F.cross_entropy()的官方定义是LogSoftmax NLLLoss的组合。这句话我建议你打印出来贴在显示器旁边因为它解释了90%的问题。当你拿一个分类模型做推理最后一层全连接输出的往往是一组未经过归一化的实数也就是logits。比如三分类任务模型输出[2.0, 1.0, 0.5]这组数本身没有概率含义它只表示各个类别的“得分”。分数可以是负数可以大于1没有任何限制。要把 logits 变成“概率”就需要 softmax[ p_i \frac{e^{z_i}}{\sum_{j1}^{C} e^{z_j}} ]softmax 做的事情是先对每个 logit 取指数再归一化。它的好处是把输出压缩到(0, 1)区间所有类别概率之和为1并且保留了原始分数的相对大小关系。而F.cross_entropy()接受的输入正是这个logits而不是 softmax 之后的结果。如果你手动在外面套了一层torch.softmax()再传给F.cross_entropy()就相当于把 logits 先归一化、取对数、再计算 NLLLoss整条链路被重复执行了最终得到的 loss 不仅数值不对方向都有可能反过来。我见过最典型的一个错误案例是某同学在模型 forward 里已经写了return torch.softmax(x, dim1)训练时又用F.cross_entropy(output, target)去算loss结果 loss 一路飙升到 2.3 以上而且怎么调学习率都降不下来。后来把 forward 里的 softmax 去掉loss 立刻正常了。1.2 公式里的“负号”和“log”到底从哪来F.cross_entropy()的计算分两步第一步是LogSoftmax[ \log(p_i) z_i - \log\left(\sum_{j1}^{C} e^{z_j}\right) ]第二步是 NLLLoss负对数似然。假设真实类别是c那么 NLLLoss 就是取LogSoftmax结果中第c个位置的值的负数[ loss -\log(p_c) ]两个步骤合在一起就变成了我们熟悉的交叉熵公式[ loss -\log\left(\frac{e^{z_c}}{\sum_{j1}^{C} e^{z_j}}\right) ]为什么非要加个log因为 softmax 输出的概率值在(0, 1)区间概率越接近1说明模型对这个样本的预测越有把握。取log之后值域变成(-∞, 0]再用负号一翻就变成了一个“越接近0越好的损失值”。这其实和日常经验是一致的模型预测正确的概率是0.9损失就是-log(0.9) ≈ 0.105如果预测正确概率是0.99损失约0.010如果概率是0.5损失约0.693。模型越自信loss 越小梯度更新时也就越温和。1.3 数值稳定性的设计哲学既然 logits 需要经过取指数这一步就不得不考虑数值溢出问题。假设某个 logit 等于1000直接算e^{1000}在 float32 下直接就是inf后面所有计算都会变成 NaN。PyTorch 在实现F.cross_entropy()时做了一个非常经典的数值稳定处理对每个 logits 向量先减去最大值再做 softmax[ p_i \frac{e^{z_i - \max(z)}}{\sum_{j1}^{C} e^{z_j - \max(z)}} ]减去最大值之后指数部分的输入都小于等于0e^0 1是上限不管 logits 多大都不会溢出。因为这种操作是线性的缩放最终 softmax 的概率分布不变所以数值上安全数学上等价。这也是为什么 PyTorch 官方推荐直接把 logits 传给F.cross_entropy()而不是自己先手动做 softmax、再取 log——自己写代码时很容易漏掉减最大值这一步导致训练过程中突然出现 NaN。2. 核心参数使用要点reduction、weight 与 ignore_index2.1 reduction 三种模式的真实差异F.cross_entropy()默认参数reductionmean也就是对 batch 内所有样本的 loss 取平均。这个默认值对大多数场景都够用但我建议你心里清楚另外两种模式什么时候值得切换。第一种是reductionsum。它把所有样本的 loss 直接相加不理解梯度下降细节的同学很容易觉得“sum 和 mean 不就是差一个倍数吗反正梯度方向一样”。但实际上如果你换了sum学习率必须重新调。同样一组超参数mean模式下的学习率是1e-3换成sum之后 loss 尺度大了 batch size 倍梯度也同步放大很容易直接把参数冲到 NaN。更典型的场景是 batch size 比较小但你想维持和平时一样的学习率表现这时候可以手动把 loss 除以 batch size等效于mean。没有必要为了这点事去改reduction。我在实际使用中sum只出现在特殊情况比如你自己实现一个计算图需要把不同的 loss 分量按权重叠加这时每个分量的 scale 需要精确控制。第二种是reductionnone。这个模式返回的是一个和输入样本数相同的一维 Tensor每个位置是那个样本的 loss。它的价值在于你可以自己对每个样本的 loss 做加权或者观察 hard example 的分布。我之前做过一个难例挖掘hard example mining的功能就是先算reductionnone的 loss然后取 top-k 个 loss 最大的样本只对这部分样本做梯度回传。这个模式在目标检测的 OHEM 方法里很常见其他场景用得少一些但知道它能干什么遇到需要时就不会卡壳。2.2 weight 参数的传参细节与陷阱weight参数用于类别不平衡场景。比如二分类问题中正样本只有5%如果你一视同仁地计算 loss模型会倾向于把所有样本都预测成负类因为这样整体 loss 最小。但这不是我们想要的行为。weight的用法很简单传一个长度为类别数的 Tensor每个位置代表该类别的权重。公式变成[ loss -w_c \cdot \log(p_c) ]其中w_c是真实类别c对应的权重。实操中需要注意几个细节第一weight的数据类型要和输入 logits 保持一致否则会报数据类型不匹配的错误。比如 logits 是float32weight也得是float32。第二如果weight是 Python 列表需要先转成 Tensor 再传入。我见过有人直接传weight[1.0, 5.0]PyTorch 新版本有些接口会帮你转换但老版本会直接报错稳妥起见torch.tensor([1.0, 5.0], device...)先转好。第三权重的绝对大小没有意义重要的是相对比例。[1, 5]和[10, 50]的效果完全一样因为梯度会被等比例放大或缩小最终影响的是学习率的选择。所以一般建议权重归一化让最大权重为1方便和其他超参数对齐。2.3 ignore_index 的语义和使用场景ignore_index的作用是忽略某些类别在 loss 中的贡献。最典型的应用场景是语义分割任务图片中有大量标注为255的 ignore 区域这些像素没有实际类别信息如果计入 loss模型会被迫去学习这些垃圾标注。使用ignore_index-1这种写法在 PyTorch 中也是合法的表示忽略所有target-1的样本。需要提醒的是ignore_index指定的类别仍然会参与 softmax 的归一化计算只是在 NLLLoss 阶段把对应位置的 loss 置零。也就是说模型依然会在这些位置上产生概率输出只是不用于反向传播。实际使用中还有一个隐藏坑当 batch 内某个样本的所有 target 都被 ignore 时reductionmean的除数是剩余有效样本数不是 batch size。这意味着如果一批样本中大量被 ignore分母会变小正常样本的 loss 会被放大。视觉分割任务里如果出现这种问题训练曲线会有明显的抖动。排查方式是打印reductionnone的结果看有效样本到底有多少。3. 实操过程从调用到验证的完整流程3.1 一个最小可运行的分类训练片段为了把前面的理论串起来我写一个最精简的三分类训练循环你直接复制就能跑。模型用一个极小的 MLP数据用随机生成的 Tensor优化器用 SGD不加任何花哨的东西。import torch import torch.nn as nn import torch.nn.functional as F import torch.optim as optim # 随机造一批数据20个样本每个样本10维特征 x torch.randn(20, 10) # 随机生成类别标签0、1、2 y torch.randint(0, 3, (20,)) # 定义一个两层的MLP最后一层输出3个logits model nn.Sequential( nn.Linear(10, 32), nn.ReLU(), nn.Linear(32, 3) ) optimizer optim.SGD(model.parameters(), lr0.01) for step in range(100): optimizer.zero_grad() logits model(x) # shape: (20, 3) loss F.cross_entropy(logits, y) # shape: () loss.backward() optimizer.step() if step % 20 0: print(fstep {step:3d}, loss: {loss.item():.4f})这个代码里有几个值得注意的点logits的 shape 是(20, 3)y的 shape 是(20,)F.cross_entropy()会自动把每个样本的第y[i]个位置作为真实类别来计算 loss。整个函数内部处理了 softmax 和取 log 的操作所以模型最后一层不需要再额外加 softmax。这也是初学者最常搞错的地方。F.cross_entropy()既可以用torch.nn.functional.cross_entropy这种函数式写法也可以用torch.nn.CrossEntropyLoss()这个类。两者的计算逻辑完全一致区别只是类版本可以把它作为nn.Module的一个子层挂在模型里方便管理 state_dict。函数式版本则更灵活适合在复杂前向逻辑中直接调用。3.2 logits 维度与 target 形状不匹配时怎么办F.cross_entropy()对输入维度的要求是logits 的第二维必须是类别数。比如(N, C)是最常见的二分类/多分类格式(N, C, d1, d2, ...)是分割任务里常见的多维格式。target 的形状则必须和 logits 去掉C维之后的形状一致。我经常遇到两个维度方面的报错。第一个是二分类任务中模型输出维度写成了(N, 1)而 target 是(N,)。这种情况下F.cross_entropy()会认为类别数是1而 target 的值却是0或1于是索引越界直接报错。正确的做法是二分类也保持输出(N, 2)target 取0/1。如果你坚持要用(N, 1)的输出那就不该用cross_entropy而应该用BCEWithLogitsLoss。第二个是 y 的 dtype 问题。target必须是一个整数类型的 Tensortorch.int64或torch.long不能是浮点类型。有些人从 CSV 或者外部接口读入标签时没做类型转换就直接传入PyTorch 会抛出类似Expected dtype int64 for target的报错。处理方法很简单y y.long()如果熟悉 PyTorch 的人应该都知道这个操作但新手经常在这里卡壳。我见过调试了半小时最后发现就是.long()没加的情况而且是那种非常明显的报错信息不看真的会让人崩溃。3.3 类版本与函数式版本怎么抉择nn.CrossEntropyLoss和F.cross_entropy是一个东西的两个面但从工程角度来看我倾向于在nn.Module中直接定义一个self.criterion nn.CrossEntropyLoss(weight...)然后在训练循环里调用self.criterion(logits, target)。这样做的理由很实际如果你在多个文件里用到了损失函数类版本可以统一管理参数比如 weight、label_smoothing万一想切换参数只需要改一处。函数式版本则适合你在某个模块里临时调用一次不值得单独抽成一个类属性。另外一点是类版本在保存模型 checkpoint 时不会有任何额外参数因为它内部没有可学习变量。但你如果用 DataParallel 或者 DDP 包裹模型时criterion 不需要被包裹它只是一个纯函数式的计算模块没有需要同步的梯度参数。4. 进阶坑位与交叉熵变体当你以为会了代码还是炸4.1 label_smoothing 参数的实际效果PyTorch 1.10 之后nn.CrossEntropyLoss直接支持label_smoothing参数。它的作用是不再把真实标签当作一个 one-hot 向量而是一个平滑过的概率分布。假设label_smoothing0.1三分类任务真实标签是[1, 0, 0]经过平滑后变成[0.9333, 0.0333, 0.0333]。也就是说模型不再被强制要求输出1而是允许保留一小部分“怀疑”。这个参数在大型分类任务里几乎成了标配。原因有两个第一纯 one-hot 目标会引导模型产生极端自信的预测特征空间中同类样本的嵌入会变得非常尖锐泛化性差。第二在数据噪声较多的任务里one-hot 本质上是在拟合一个可能存在标注错误的硬目标。实际使用中label_smoothing一般设0.1左右太大了模型会变得“过于谦虚”训练指标如 accuracy会比不设平滑时略低但验证集和测试集的表现往往会更好。这个 trade-off 需要自己权衡。4.2 softmax 输出二次传入导致的梯度消失我前面提到的“模型 forward 里已经加了 softmax外面又调 cross_entropy”这个场景除了 loss 数值变大之外还有一个更隐蔽的后果梯度消失。原因是cross_entropy里自带的LogSoftmax会对输入再次取 log。如果你传进来的是 softmax 之后的概率值这个概率值一旦非常接近0或1取 log 之后梯度会趋于无穷或消失。直观地说如果 softmax 输出是[0.99, 0.005, 0.005]取 log 后得到[-0.01, -5.3, -5.3]这个时候再用 NLLLoss 去对 logits 求导梯度链路上的值会变得非常不均衡——有的位置几乎不更新有的位置更新幅度异常大。最终结果就是训练不稳定loss 曲线上下剧烈震荡模型输出在几个类别间疯狂跳变。排查方法很简单在 forward 里把输出打印出来看是不是已经通过了 softmax。这一类错误有个共同的“味道”就是模型表现忽好忽坏看起来像学习率太大但调小学习率又收敛得太慢卡在一个很奇怪的状态。4.3 处理极度不平衡数据手动加权比你想的更有效有一个 5% 正样本的二分类任务我一开始试了weighttorch.tensor([1.0, 19.0])训练过程还算稳定但模型预测的召回率始终不达标。后来我把weight调成[1.0, 10.0]并对负样本做了 random undersample效果反而更好了。这个经验想说明的是class weight 不是越极端越好。当正负样本比例差距过大比如 1:100单纯靠 loss 加权会让模型对正样本过拟合产生大量假阳性。合理的做法是先做一个轻量的 undersample 或 oversample把比例控制在 1:5 到 1:10 左右再用weight微调。这样模型既能看到足够的负样本防止误报又不会因为正样本过多而出现过拟合。另外一点weight要记得放到和模型参数同一个 device 上。分布式训练时尤其容易漏掉因为weight不会自动被 DDP 同步你得手动保证它在每一块 GPU 上都是相同的否则会出现不同 rank 计算的 loss 不一致的问题。4.4 与 BCEWithLogitsLoss 的选型对比F.cross_entropy()处理的是多分类F.binary_cross_entropy_with_logits()处理的是二分类或多标签分类。两者核心区别在于cross_entropy 使用 softmax 作为归一化所有类别概率和为1BCEWithLogitsLoss 使用 sigmoid 作为激活函数每个类别独立判断概率之间不强制互斥。多标签任务比如一张图片里可能同时有猫和狗就应该用BCEWithLogitsLoss而不是把问题硬生生转成多分类。我在一个行人属性识别任务中一开始用 cross_entropy 强行训练效果非常差。后来改成 BCEWithLogitsLoss每个属性单独训一个二分类头精度立刻上去了。选择标准就一句话类别之间如果互斥用 cross_entropy类别之间独立存在用 BCEWithLogitsLoss。这是个经验法则虽然简单但能避免很多无谓的调参。5. 常见报错排查一分钟定位问题报错信息可能原因解决方法Expected target size [N, C], got [N]target 形状不对多分类任务传了 one-hot 向量用torch.argmax(target, dim1)或者直接传类别索引Expected dtype int64 for targettarget 是 float 类型.long()或.int()转成整数IndexError: index X is out of bounds for dimension 0 with size Ctarget 里有超出类别范围的数值检查标签是否从0开始是否有漏标或标注错误loss is nanlogits 溢出或者学习率太大检查模型输出是否有 NaN调低学习率确认没有手动 softmax训练时 loss 不下降可能外部套了 softmax或weight设置不合理打印 logits 和 loss确认 forward 逻辑多卡训练时 loss 不一致weight 没有同步到所有 device用weight.to(device)并在每个进程里都执行一次我在面试算法工程师时经常用F.cross_entropy()当开胃题。大多数人能说出“它是 softmax NLLLoss”但问到“如果模型输出是 logits 但 target 是 one-hot你会怎么处理”不少人答不上来。其实这是个很好的分水岭理解函数设计的人不会拘泥于函数表面而是清楚它的数据流走向。关于 one-hot target 多说一句。PyTorch 的 cross_entropy 设计成了类索引形式这是从内存和计算效率角度考虑的。如果你手头的数据是 one-hot需要先用argmax转成索引。不要自己写target_onehot * torch.log(softmax(logits))这类公式不仅数值不稳定还会把你绕进维度匹配的死胡同。6. 关联热词的延伸思考6.1 从 cross_entropy 到自定义损失函数的思路理解了F.cross_entropy()的组成之后很多自己写损失函数的需求就变得清晰了。比如你想实现 focal loss核心逻辑就是在标准交叉熵前面加一个调制因子(1 - p_t)^γ。由于F.cross_entropy()已经把log_softmax和 NLLLoss 合在一起你拿不到每个样本的p_t所以一般的做法是手动拆开计算log_probs F.log_softmax(logits, dim-1) probs torch.exp(log_probs) # 或者直接 softmax loss_per_sample F.nll_loss(log_probs, target, reductionnone) # 取正确类别的概率 pt probs.gather(1, target.unsqueeze(1)).squeeze(1) focal_weight (1 - pt) ** gamma loss (focal_weight * loss_per_sample).mean()gather在这里的作用是从probs矩阵中按照target指定的索引把每个样本真实类别的概率取出来。很多人第一次遇到gather会懵但其实它做的就是“按位置取值”这件事。以上代码可以灵活实现任何自定义变体掌握了从 cross_entropy 拆解到 log_softmax nll 这条链路自定义损失函数就不再是玄学。6.2 与其他类似函数的关系回顾F.nll_loss()接受的是 log 概率F.log_softmax()接受的是 logits。F.cross_entropy()就是这两个函数的缝合怪。单独用F.nll_loss(F.log_softmax(logits, dim-1), target)得到的结果和F.cross_entropy(logits, target)完全一致。唯一差别是 cross_entropy 内部做了一次数值稳定的优化不会出现溢出的风险。这在代码层面解释了为什么“手动套 softmaxlog”是一个 bad idea因为如果你想手动复现最接近的写法是nll_loss(log_softmax(logits), target)而不是自己用torch.log(torch.softmax())。一旦你用后者就失去 PyTorch 内置的数值保护了。6.3 环境搭建与模型训练的常见关联坑说到 PyTorch 损失函数不可避免地会牵扯到环境搭建问题。最近网上 pytorch 安装教程非常多GPU 版本和 CPU 版本选哪个、conda 还是 pip 装、CUDA 版本怎么匹配这些老生常谈的话题我这边就不再展开了。只提醒一句如果你在 Windows 上遇到pip命令未识别、python环境混用等问题多半是环境变量和虚拟环境的问题建议直接上 anaconda 虚拟环境管理避免在系统级 Python 里装了一堆互相冲突的包。顺带说一下很多人问“为什么我计算 loss 的时候GPU 利用率上不去”。如果你已经在用F.cross_entropy()并且确认数据已经.cuda()或.to(device)了那问题大概率出在 batch size 太小或 CPU 数据加载成了瓶颈。loss 本身的计算量很小不会成为性能瓶颈但如果你在 loss 里做了什么奇怪的操作比如对每个样本循环 Python 遍历那 GPU 利用率掉下去是必然的。结尾关于F.cross_entropy()我实际使用中的最大体会是不要把它当成一个“黑盒”去调参。你现在知道了它是LogSoftmax NLLLoss知道了它吃的是 logits 而不是概率知道了 target 要用整数索引而不是 one-hot后面写模型时很多报错和异常行为都能一眼定位。最后再分享一个小技巧。当你怀疑自己某段代码里 cross_entropy 传参有问题时最简单的验证方法是对同一个 logits 和 target分别跑一遍F.cross_entropy(logits, target)和F.nll_loss(F.log_softmax(logits, dim-1), target)如果两者结果不一致那肯定是数值类型、device 或者维度上出了问题。这个对照思路帮我在线上排查时省下过不少时间希望对你有用。