PyTorch中代码'-input[range(target.shape[0]),target]'工作原理咨询
解析自定义负对数似然(NLL)损失函数的工作原理
这段代码是**负对数似然损失(Negative Log-Likelihood, NLL)**的手动实现,和PyTorch内置的nn.NLLLoss()(默认参数)功能完全一致,下面一步步拆解逻辑:
1. 确定批量样本数量
target.shape[0]获取当前批次的样本总数,比如target是形状为[4]的张量(对应4个样本),range(target.shape[0])会生成[0,1,2,3]的序列,对应每个样本在批量中的索引。
2. 提取真实类别对应的模型输出
input[range(target.shape[0]), target]是PyTorch的高级索引用法:
- 第一个维度的索引
range(...)指定每个样本的位置 - 第二个维度的索引
target指定每个样本对应的真实类别下标 - 两者组合后,会从input张量中逐个取出每个样本真实类别对应的模型输出值,最终得到一个和target形状相同的一维张量。
举个实际例子:
假设input(已经过log_softmax处理,因为NLL损失基于对数概率)是:
tensor([[-1.5, -0.2, -3.1], # 样本0,3个类别 [-0.4, -2.0, -0.1]]) # 样本1
target是tensor([1, 2]),那么这行代码会取出:
- 样本0的第1个元素:
-0.2 - 样本1的第2个元素:
-0.1
最终得到张量tensor([-0.2, -0.1])。
3. 计算负对数似然
给提取出的张量加负号-,得到tensor([0.2, 0.1])。这一步的逻辑是:模型对真实类别的对数概率越高(越接近0),取负后的值越小,代表模型预测越准确,损失越低,符合损失函数的优化目标。
4. 计算批次平均损失
最后调用.mean()对所有样本的损失值取平均,得到整个批次的平均NLL损失,上面例子的结果就是0.15。
补充说明:这里的input必须是经过log_softmax处理后的张量,因为NLL损失是直接基于对数概率计算的——如果用原始的模型输出(未经过log_softmax),应该用交叉熵损失nn.CrossEntropyLoss()(它内置了log_softmax和NLLLoss的计算)。
内容的提问来源于stack exchange,提问作者Yan W
相关产品推荐
相关产品推荐

