You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch CrossEntropyLoss类别权重修改无效问题及实现疑问

PyTorch CrossEntropyLoss 权重计算逻辑详解

你遇到的问题其实是对PyTorch中CrossEntropyLoss权重处理逻辑的误解,核心有两个关键点:加权平均的计算方式和reduction参数的行为,我们一步步拆解:

1. 为什么修改权重(非0)后损失值不变?

看你第一个单样本的例子:

from torch import nn
import torch
logits = torch.FloatTensor([ [0.1, 0.9], ])
label = torch.LongTensor([0])
# 权重[1,1]
criterion = nn.CrossEntropyLoss(weight=torch.FloatTensor([1, 1]))
print(criterion(logits, label).item()) # 1.1711
# 权重[0.1,1]
criterion = nn.CrossEntropyLoss(weight=torch.FloatTensor([0.1, 1]))
print(criterion(logits, label).item()) # 还是1.1711

这是因为PyTorch默认的reduction='mean'在有权重时,不是简单对加权损失取平均(除以batch size),而是将加权损失的总和除以batch中所有样本对应类别的权重之和。

在单样本场景下:

  • 权重[1,1]时:加权损失是1.1711 * 1 = 1.1711,权重之和是1,最终损失是1.1711 / 1 = 1.1711
  • 权重[0.1,1]时:加权损失是1.1711 * 0.1 = 0.11711,权重之和是0.1,最终损失是0.11711 / 0.1 = 1.1711

两者结果自然一致!只有当权重设为0时,该样本的权重之和为0,PyTorch会直接忽略该样本的损失,输出0。

2. 你的自定义实现与PyTorch官方的差异

看你第二个双样本的例子,结果差异的根源同样是reduction的计算方式:

你的实现逻辑:

return torch.mean((x_terms+log_terms)*weights)

这里是将每个样本的加权损失相加后,除以batch size(样本数量),属于简单平均。

PyTorch官方的默认逻辑:

当reduction='mean'且传入weight时,损失计算为:
$$\text{loss} = \frac{\sum_{i=0}^{N-1} \text{weight}[y_i] \times \text{cross_entropy}(x_i, y_i)}{\sum_{i=0}^{N-1} \text{weight}[y_i]}$$
也就是加权损失总和除以batch中样本对应类别的权重之和,属于加权平均。

我们用你的例子验证:

  • 样本0的原始损失≈1.172,权重0.1 → 加权损失=1.172*0.1=0.1172
  • 样本1的原始损失≈0.644,权重1 → 加权损失=0.644*1=0.644
  • 加权损失总和=0.1172+0.644=0.7612
  • 权重总和=0.1+1=1.1
  • PyTorch最终损失=0.7612/1.1≈0.692,和你得到的结果一致
  • 你的实现结果=0.7612/2≈0.3806,也和你的输出一致

3. 如何让PyTorch的行为和你的自定义实现一致?

如果你希望PyTorch采用简单平均(除以batch size),可以将reduction设为'sum'后手动除以batch size:

criterion = nn.CrossEntropyLoss(weight=weight, reduction='sum')
loss = criterion(logits, label) / len(label)
print(loss.item()) # 结果会和你的自定义实现一致

4. 额外补充:权重的归一化误区

很多人误以为PyTorch会对权重做全局归一化(除以所有权重的总和),但实际上并没有——权重的作用是直接乘以对应类别的损失,最终的平均是基于batch中实际用到的权重之和,而非全局权重总和。

内容的提问来源于stack exchange,提问作者AveryLiu

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.11 09:07:06