拓冰建站拓冰建站
首页 / 资讯中心 / 正文

PyTorch F.cross_entropy 深度解析:原理、数值稳定性与踩坑指南

先问大家一个问题如果你的分类模型里有一行loss F.cross_entropy(logits, target)你能清楚说出它和torch.softmax之后再取负对数有什么区别吗我刚用 PyTorch 的前两年对这个问题并不在意因为两行代码跑出来的 Loss 数字差不多。直到有一次我手写一个特殊损失函数在一个特别大的 logits 输入下得到了 nan才意识到F.cross_entropy()这个“一行调用”背后其实藏了不少设计决策。这篇文章就围绕F.cross_entropy()本身展开适合对 PyTorch 有一定了解、但一直把它当黑盒用的读者。我会把它的真实计算流程、类别维度过约定、数值稳定性逻辑、常用参数、容易踩的坑以及一些实战里的 Debug 习惯一次讲清楚。读完你应该能明白遇到形状报错或 Loss 异常时该往哪个方向检查。1. F.cross_entropy 到底替你做了什么1.1 先别在脑子里生成 softmax很多教程把 cross_entropy 解释成“softmax 交叉熵”这个说法不算错但它是从数学公式角度说的不是从实际实现角度说的。F.cross_entropy的真实执行逻辑是先对输入做log_softmax再做nll_loss负对数似然损失。也就是下面这个式子loss(x, c) -log_softmax(x)[c] -log( exp(x_c) / sum_j exp(x_j) ) -x_c log( sum_j exp(x_j) )这里有个非常反直觉的点函数内部根本不单独输出 softmax 概率也不要求你传概率进来。它要求传进来的是 logits也就是模型最后一层线性层nn.Linear的原始输出。为什么因为 softmax 里的指数运算在数值上特别容易爆掉具体细节我在第 3 节展开。从工程实现角度看PyTorch 把log_softmax和nll_loss合并成一个函数还有一个额外收益只需要一次前向计算就能同时得到 log 概率以及 target 位置对应的 loss。如果拆开写你要先保存一份完整概率矩阵再取 log再索引目标位置中间多出来的内存和访存开销在 batch 很大的时候会明显放大。训练阶段用F.cross_entropy推理阶段如果只想看概率分布再单独调用torch.softmax这是最常见的组合方式。所以结论是这个函数是给“训练时算损失”准备的不是给“模型输出概率”准备的。1.2 数据类型的硬性约定输入输出的约定上logits的 dtype 必须是浮点类型通常是 float32target的 dtype 必须是整数类型通常是 int64 / long并且 target 的值表示类别索引范围是 0 到 C-1C 是类别总数。很多人第一次跑分类模型就报Expected object of scalar type Long but got scalar type Float就是因为 target 是 float 格式。第二个容易犯的错是把 target 传成 one-hot 向量F.cross_entropy完全不认这种格式它把 target 的每个元素当作一个类别索引来处理。如果你手里恰好是 one-hot 标签正确做法是先target.argmax(dim1)再传进去。一个典型的标准调用长这样import torch import torch.nn.functional as F logits torch.tensor([[2.0, 0.5, 0.1], [1.0, 3.0, 0.2]]) target torch.tensor([0, 1]) # 必须是 LongTensor loss F.cross_entropy(logits, target) print(loss)这里 logits 的最后一维长度是 3对应 3 个类别target 里的每个值都不能超过 2超过就会报IndexError: Target 3 is out of bounds。2. 类别维度dim1 的默认约定是最容易翻车的地方2.1 N,C,d1,d2 的布局不只是二维输入才适用PyTorch 对张量维度的组织习惯是“通道维在前”分类任务的标准形状是(N, C)N 是 batch sizeC 是类别数。F.cross_entropy内部默认dim1作为类别维但这个约定很多时候被忽略了。对于图像分割这类任务模型的输出形状是(N, C, H, W)这时候F.cross_entropy依然有效它会把 C 维当成类别维对 H 和 W 上的每个像素分别计算交叉熵再取平均。看下面这个例子logits torch.randn(2, 3, 224, 224) # N, C, H, W target torch.randint(0, 3, (2, 224, 224)) # N, H, W loss F.cross_entropy(logits, target) print(loss)注意 target 的形状是(N, H, W)它不是把 C 维去掉而是直接把 C 维“折叠”掉了。也就是说PyTorch 要求 target 的形状和 logits 去掉 C 维之后保持一致。再往大了说如果模型输出是(N, C, d1, d2, ...)target 就是(N, d1, d2, ...)。这个规则对很多人来说需要反应一下。我看到过不少人在分割任务里先对 logits 做argmax再算 loss或者把 target 手工 flatten 成(N*H*W,)这些都不是必须的。直接按原始形状传进去F.cross_entropy自己知道怎么对齐。2.2 最容易出现的三种 shape 相关报错第一种是Expected target size [2, 3], got [2]。这类报错常见于你有一个形状(N, C, H, W)的 logits却只传了一个(N,)的 target函数期望的是(N, H, W)。第二种是把 one-hot 标签直接当 target 传比如(N, C)的 target 配上(N, C)的 logits函数会尝试把每个 one-hot 向量当成一个类别索引结果当然对不上。第三种是把 logits 写成了(N,)而 target 是(N,)这是维度数量本身就不够。我的建议是碰到 shape 报错先别急着改代码先打印两行print(logits.shape) print(target.shape)然后人工确认 logits 的倒数形状是否满足(..., C)target 是否把 C 排除在外。这一步能省掉一半以上的 debug 时间。2.3 为什么类别维永远是 dim1而不是最后一维这个问题经常有人问。F.cross_entropy对输入形状的约定是在函数内部写死的它把第 1 维0 开始计数当作类别维。如果你的数据形状是(N, H, W, C)比如某些图像库读出来的是通道在最后不经过permute直接丢给F.cross_entropy它会把 H 当成类别维。如果 H 恰好大于类别数函数不会报错但 loss 会变成一个让你完全摸不着头脑的数值而且模型训练起来异常缓慢。遇到这种情况第一步检查的就是通道维到底在第几位。3. 手写交叉熵和内置 F.cross_entropy 的差距当朴素实现输出 nan 时我在想什么3.1 朴素实现与内置实现的公式等价性用一个简单的例子来验证手动实现“softmax 取负对数”和内置F.cross_entropy在小数值输入下结果应该完全一致。def manual_cross_entropy(logits, target): exp_x torch.exp(logits) # (N, C) probs exp_x / exp_x.sum(dim1, keepdimTrue) # (N, C) return -torch.log(probs[range(len(target)), target]).mean() logits torch.tensor([[1.0, 2.0, 3.0], [2.0, 1.0, 0.5]]) target torch.tensor([2, 0]) print(manual_cross_entropy(logits, target)) print(F.cross_entropy(logits, target))正常情况下两个结果一致精确到小数点后好几位。这份等价性也说明了一个问题只要数值不溢出拆开写和合起来写是一样的。但一旦 logits 里的某个数字比较大差异立刻显现。3.2 大 logits 下的数值灾难我遇到过一次很明显的情况。模型在某个迭代里输出了一批 logits最大值超过了 1000手动实现直接得到 inf再取 mean 变成 nan训练直接崩了。而同样的 logits 传给F.cross_entropyloss 是一个正常的有限数值模型在下一个 step 还能继续跑。原因在于内置实现计算log_softmax时不是直接做exp(logits)而是先减去最大值再算。我刚才推导过log_softmax(x)_i x_i - max(x) - log( sum_j exp(x_j - max(x)) )x_j - max(x)最大是 0剩下的都小于 0这样exp的结果不会超过 1就不会上溢。而手动实现里exp(1000)直接就是无穷大后面全部失效。这个“减最大值”的操作是 softmax 家族函数防止数值溢出的标准做法F.cross_entropy内部帮你做了但拆开写的时候你得自己做。3.3 把已经 softmax 的结果再传给 cross_entropy 会发生什么还有一个非常常见的误用有人先算了prob torch.softmax(logits, dim1)然后把prob当成 logits 传给F.cross_entropy。表面上看形状没问题target 也对loss 也能算出来但结果完全错误。因为F.cross_entropy内部会对输入再做一次log_softmax等于说你对一个已经归一化到 0 到 1 之间的概率分布又做了一次对数归一化。这不是双重惩罚那么简单它让梯度信息被严重扭曲。更迷惑的是这种错误写法算出来的 loss 往往更小容易给初学者一种“效果不错”的错觉——实际上模型根本没有按照正确的目标在优化。判断自己有没有犯这个错最简单的办法是看模型最后一层有没有接 softmax。如果nn.Linear的输出直接传给F.cross_entropy那就是对的如果中间插了nn.Softmax或torch.softmax那就要小心了。4. 那些容易被忽略但关键时刻能救命的参数4.1 weight类别不平衡时的正确姿势和副作用F.cross_entropy的weight参数可以对每个类别单独加权。处理类别不平衡时一个简单做法是统计训练集里每个类别的样本数然后让少数类的权重更大。比如二分类正负样本比是 1:9可以用weight torch.tensor([9.0, 1.0]) # 给第 0 类更大的权重 loss F.cross_entropy(logits, target, weightweight)注意 weight 的取值含义它是乘在该类别 loss 上的系数少数类的 loss 被放大之后梯度更新力度也会变大。副作用是使用 weight 后 loss 的绝对数值明显变大如果你之前已经调好了一组学习率加了 weight 之后可能会发现训练变得不稳定。我的经验是weight 不要设置得过于极端一个从 1 到 10 之间的比例通常够用如果类别非常不均衡还可以考虑对 loss 做归一化或者换用 Focal Loss 这类更平滑的加权方式。4.2 ignore_index处理 padding 片段时的暗坑NLP 或序列任务里经常要对 padding 位置做 maskignore_index就是干这个的loss F.cross_entropy(logits, target, ignore_index-100)这里-100是一个特殊 flag表示“这个位置的 loss 不算”。PyTorch 内部会把这个位置的梯度也去掉这看似干净利落但有一个细节很多人不知道ignore_index 仅仅是让这个位置不贡献 loss不代表它完全不参与计算。log_softmax的归一化分母仍然包含这个被忽略位置的指数值。如果某个被 ignore 的 logit 特别大它照样会影响其他类别的概率估计。举个实际例子训练序列模型时你把 padding 位置全部设成ignore_index0然而你的真实类别里恰好也有 0 这个类别那么所有真实标签为 0 的样本都会被当成 padding 忽视掉模型永远学不会这个类。我用过一次之后就把一个分类任务里的第 0 类彻底学废了后来检查类别分布才发现问题。解决办法要么把 ignore_index 设成一个类别数之外的数比如 -100要么确保 padding 用的特殊 id 和真实类别完全不重叠。4.3 label_smoothing 和 reduction 的配合从 PyTorch 1.10 开始F.cross_entropy原生支持label_smoothing参数。设置label_smoothing0.1时target 不再是一个 one-hot 硬标签而是把一部分概率平均分给所有类别。这种做法在很多分类任务里能提升泛化性但它会带来一个现象训练 loss 的下限不再是 0而是围绕某个正数波动。如果你习惯了看 loss 掉到非常低才认为收敛用 label smoothing 之后要注意改变判断标准。reduction参数控制 loss 的聚合方式默认是mean返回 batch 内所有样本的平均值sum返回总和none返回每个样本单独的 loss。训练时通常用默认的mean就可以了但如果你在做一些特殊的 loss 加权、比如要把不同 batch 的 loss 拼接起来分析reductionnone会更有用。还有一个细节用了ignore_index之后mean的分母不是整个 batch 的样本数而是扣除被 ignore 的样本数这在某些统计需求下要留意。5. 二分类、多标签和 one-hot跨出“多分类单标签”舒适区后的常见混淆5.1 为什么二分类也要用两列 logits很多人会想二分类嘛模型输出一个神经元用 sigmoid 变成 0 到 1再和标签 0/1 算交叉熵不就行了这个想法没错但在F.cross_entropy的框架里它要求输入是一个(N, C)形状的 logits。所以哪怕只有两类你也需要输出两列比如(N, 2)target 取 0 或 1。如果你把最后一层nn.Linear的输出维度设成 1再直接把输出丢给F.cross_entropy当 target 里有 1 这个类别时就会报IndexError: Target 1 is out of bounds因为输入最后一维长度只有 1索引范围是 0 到 0。很多人一开始都踩过这个坑其实不是函数的问题是类别维度的设计不符合函数约定。如果坚持要用单神经元输出更合适的函数是F.binary_cross_entropy_with_logits它接受(N, 1)的 logits 和(N,)的 0/1 标签内部自动完成 sigmoid 和 BCE 的计算数值稳定性也是处理好的。两者选哪个取决于你模型最后一层的设计而不是随便换一个。5.2 多标签任务为什么不能直接用 F.cross_entropy多标签任务里一个样本可能同时属于多个类别比如一张图里同时有猫和狗。F.cross_entropy的数学假设是“每个样本只能属于一个类别”它把所有类别的概率归一化成一个分布天然不适合多标签场景。多标签的正确做法是为每个类别单独做一个二分类然后用F.binary_cross_entropy_with_logits对每个标签计算 loss再取平均。这也解释了为什么常见的目标检测和图像多标签分类框架里head 通常会输出(N, num_classes)但配合的是 BCE Loss 而不是 CrossEntropy Loss。两者看起来形状一样但损失函数对“类别之间关系”的假设完全不同。5.3 one-hot target 的处理argmax 可能是最少出错的方式如果你从某个数据集或预处理流程里拿到的就是 one-hot 标签形状是(N, C)传给F.cross_entropy会直接报错或产生诡异结果。最直接的转换是target target.argmax(dim1)这句代码会把(N, C)变成(N,)每个元素是 one-hot 里值为 1 的位置正好是类别索引。看起来很无脑但确实是最不容易出错的方式。不过要注意如果你是想做知识蒸馏或自定义软标签比如 target 不是 one-hot而是 0.7/0.3 这种概率分布F.cross_entropy原生并不支持因为它要求 target 是整数索引。PyTorch 官方提供的方案是label_smoothing但它只能做均匀平滑不能自定义任意分布。如果真有这个需求通常需要手动实现或者改用nn.KLDivLoss配合log_softmax来算。6. 高频运行时错误与 NaN 问题的排查清单6.1 运行时错误定位顺序为了在实际项目里快速定位F.cross_entropy相关的报错我总结了一个排查顺序先看 dtype再看 shape最后看数值范围。报错信息原因解决方案Expected object of scalar type Long but got scalar type Floattarget 不是整数类型target.long()IndexError: Target X is out of boundstarget 中的类别值 类别总数 C检查标签是否从 0 开始检查 C 是否设置正确Expected target size [...]target 形状与 logits 去掉类别维后的形状不一致打印两者 shape 人工对齐loss 为 nan 但不报错logits 中存在 inf/nan或学习率过高检查上游输入、梯度、学习率这个表看起来简单但实际项目里最容易漏掉的是最后一行。F.cross_entropy本身不会对你的 logits 做数值合法性校验如果上游某个地方已经产生了 nan它只会安静地把 nan 传播到 loss 里而不是抛异常。所以你看到 loss 是 nan 的时候第一反应不应该是怀疑这个函数而是去看 logits 是从哪一层开始变成 nan 的。6.2 loss 出现 nan 的排查思路F.cross_entropy返回 nan大致可以分成三类原因。第一类是 logits 本身有 nan 或 inf。这通常是由于模型内部出现了数值异常比如某些层的权重更新过大、矩阵乘里出现了除零。排查手段是在 loss 计算前加一个钩子或直接 print观察logits是否存在torch.isnan(logits).any()。如果 logits 正常再往后看。第二类是学习率过大。分类模型训练初期如果 loss 直接变成 nan多半是学习率太高优化一步就把权重推到数值不稳定的区间。这个和F.cross_entropy本身无关需要调小学习率或者加梯度裁剪。第三类是混合精度训练下的偶发 nan。AMP 会动态调整 loss 缩放因子如果缩放过程出现问题也可能导致反向传播后出现 nan。这个时候通常不是F.cross_entropy的锅而是整个训练管线里的梯度缩放策略需要调整。6.3 冻结层 损失函数的一个容易被忽略的问题还有一个场景迁移学习时经常冻结 backbone 的梯度只训练分类头。这时候如果发现某个 step 的 loss 突然不变或者变成 0新手容易怀疑F.cross_entropy写错了。实际上冻结层本身不会影响损失函数但如果模型输出在冻结层之后变成常量loss 自然也会变成常量。调试时可以先单独检查logits是否对不同输入有不同输出再决定是不是要解冻某些层。7. 一些实操心得7.1 给常用训练代码封装一层我习惯在项目里把F.cross_entropy的使用封装成一个简单函数统一好weight、ignore_index等参数避免在多个训练脚本里重复写容易出错的逻辑。比如我会定义这样一个 wrapperdef ce_loss(logits, target, ignore_index-100, weightNone, label_smoothing0.0): if logits.dim() 2: # 高维输入比如分割任务保持原形状即可 pass return F.cross_entropy( logits, target.long(), weightweight, ignore_indexignore_index, label_smoothinglabel_smoothing, )不要小看这一步。实际项目里target经常是各种 dtype或者从 DataLoader 出来就已经是 float 了统一在入口处转一次 long能减少大量重复报错。7.2 训练前先做一次 loss 冒烟测试在真正开始训练前我会先用一个随机 batch 跑一次前向和反向验证 loss 是否能正常下降。具体做法是构造一个固定 batch 的随机 logits 和随机 target跑一步优化看 loss 是不是变小。如果随机数据下 loss 都掉不下去那多半是模型结构或函数用法的问题如果能正常下降再换成真实数据。这个小习惯帮我过滤掉了很多“看起来没问题但一训练就崩”的情况。7.3 不要省那一行 shape 打印最后再分享一点经验无论 logits 是(N, C)还是(N, C, H, W)训练前都先打印一次logits.shape和target.shape。很多人觉得这行 print 无所谓但根据我的经验F.cross_entropy的绝大多数报错和异常行为最后都能追溯到“形状假设不一致”上。先确认类别维在 dim1、target 的值范围在 0 到 C-1 之间再去看所谓的高级问题。这些基础检查看起来不起眼但在实际调试里回报率是最高的。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门