1. 从一次真实的报错说起:为什么我的loss.backward()炸了?
那天下午,我正在调试一个图像分类模型,代码跑得好好的,突然就给我甩了个脸子,终端里蹦出来一行刺眼的红字:RuntimeError: grad can be implicitly created only for scalar outputs。相信不少用过PyTorch的朋友都见过这个“老朋友”,尤其是在你刚接触CrossEntropyLoss,并且手贱(或者说,出于好奇)把reduction参数设成了'none'的时候。
这个报错翻译过来就是:“梯度只能为标量输出隐式地创建”。听起来有点绕,对吧?我用大白话给你解释一下:loss.backward()这个方法,它默认的“工作对象”是一个单一的数值,也就是一个标量。比如,你算出来一个损失值是0.85,那没问题,backward()知道该怎么做,它会沿着计算图一路回溯,算出每个参与运算的参数的梯度。但如果你给它的不是一个数,而是一坨数,比如一个形状为[32, 10]的张量(假设你batch size是32,有10个类别),backward()就懵了:“大哥,你这一下子给我32个样本各自的10个类别的损失,我该先算哪个的梯度?我又该把梯度累加到哪儿去?”
我当时的情况就是如此。我写了一句 loss = nn.CrossEntropyLoss(reduction='none'),本意是想看看每个样本的损失具体是多少,方便我做些样本级别的分析或者加权。结果在loss.backward()的时候,程序就直接“罢工”了。这个报错的核心,其实就是reduction这个参数在背后“捣鬼”。它决定了CrossEntropyLoss计算完损失之后,以什么样的形式交到你手上。是给你一个总览性的单一数字(mean或sum),还是把每个样本的“成绩单”都原封不动地给你(none)。而backward()这个函数,它有个“小脾气”:它只喜欢处理那份总览性的“成绩汇总”,不喜欢处理那一沓厚厚的、每个人的原始试卷。
所以,理解reduction,不仅仅是知道它有三个选项,更是理解PyTorch自动微分(autograd)机制如何工作的一把钥匙。选错了,轻则报错,程序跑不起来;重则可能 silently 地引入bug,让你的模型训练朝着错误的方向狂奔而不自知。接下来,我们就彻底掰开揉碎,看看这个reduction参数到底怎么玩,以及怎么避开它挖的坑。
2. 深入CrossEntropyLoss:reduction参数的三种面孔
要弄明白reduction,我们得先回到CrossEntropyLoss本身是干什么的。简单说,它常用于分类任务,把模型输出的原始分数(logits)和真实的类别标签,换算成一个衡量模型“犯错程度”的数值。这个换算过程本身,比如softmax加负对数似然,我们先不深究。关键是在这个换算之后,reduction参数登场了,它来决定如何“汇总”一个批次(batch)里所有样本的损失。
PyTorch官方给了我们三个选择:‘none’、‘mean’和‘sum’。我们一个一个来看,并用代码和例子把它们讲清楚。
2.1 reduction='none':最原始的“成绩单”
当你设置reduction='none'时,损失函数会变得非常“实诚”。它不会做任何额外的聚合操作,直接返回每个样本的损失值。
import torch
import torch.nn as nn
# 模拟一个batch:4个样本,3个类别
logits = torch.randn(4, 3) # 模型输出的原始分数,shape [4, 3]
labels = torch.tensor([0, 2, 1, 0]) # 真实标签,shape [4]
# 使用 reduction='none'
criterion_none = nn.CrossEntropyLoss(reduction='none')
loss_none = criterion_none(logits, labels)
print(f"logits shape: {logits.shape}")
print(f"loss_none value: {loss_none}")
print(f"loss_none shape: {loss_none.shape}")
运行这段代码,你会看到loss_none是一个形状为[4]的一维张量。它包含了4个独立的数值,分别对应batch中第1个样本的损失、第2个样本的损失……以此类推。
这有什么用呢? 场景其实很多。比如,你想做难例挖掘(Hard Example Mining),需要根据每个样本的损失大小来调整其权重,损失大的样本在下一轮训练中给予更多关注。或者,你在做某些自监督学习、对比学习任务时,需要构造样本对之间的特定损失关系。再比如,你只是想单纯地监控一下每个样本的损失分布,看看模型在哪些数据上表现不稳定。‘none’模式给了你最大的灵活性,让你能接触到最原始的损失信息。
但是,最大的“坑”也在这里:正如开篇报错所示,你无法直接将这个形状为[4]的张量扔给loss.backward()。因为backward()需要的是一个标量来启动整个链式求导过程。直接调用loss_none.backward(),百分百会触发RuntimeError。
2.2 reduction='mean':最常用的“平均分”
这是CrossEntropyLoss的


379

被折叠的 条评论
为什么被折叠?



