PyTorch手动实现带ignore_index的CrossEntropyLoss结果不一致问题
你的自定义交叉熵损失和PyTorch内置实现的差异来自平均损失时的分母取值不同。
PyTorch的nn.CrossEntropyLoss默认使用reduction='mean'配置,这种情况下计算平均损失时,只会统计标签不等于ignore_index的有效样本数作为分母,而不是除以总batch大小。你当前的实现直接除以了完整的batch样本数,当存在被忽略的填充标签时,计算出来的损失就会比内置实现更小,刚好等于内置结果乘以(有效样本数/总样本数)。
你给出的测试用例中总样本是100,x[40:]都是-100,有效样本只有40个,内置实现除以40,你除以100,25.55 * (40/100) ≈ 10.22,和你得到的结果完全吻合。当没有-100标签时,有效样本数等于总样本数,所以二者结果一致。
修复后的实现代码
import torch import torch.nn as nn class compute_crossentropyloss_manual: """ y0是模型输出的logits,形状为 (batch_size,C) x是标签,形状为 (batch_size),元素为0到C-1的整数,或ignore_index """ def __init__(self, ignore_index=-100) -> None: self.ignore_index=ignore_index def __call__(self, y0, x): loss = 0. n_batch, n_class = y0.shape valid_count = 0 # 统计有效样本数 for y1, x1 in zip(y0, x): class_index = int(x1.item()) if class_index == self.ignore_index: continue # 也可以替换为loss += nn.functional.log_softmax(y1, dim=0)[class_index] 避免数值溢出 loss = loss + torch.log(torch.exp(y1[class_index])/(torch.exp(y1).sum())) valid_count +=1 # 避免有效样本为0时除以0 loss = - loss/valid_count if valid_count >0 else 0. return loss
验证效果
用你提供的测试用例验证,修复后的代码输出会和nn.CrossEntropyLoss完全一致:
criterion = nn.CrossEntropyLoss() criterion2 = compute_crossentropyloss_manual() y = torch.randint(1, 50, (100, 50), dtype=torch.float) x = torch.randint(0, 50, (100,)) x[40:] = -100 print(criterion(y, x).item()) print(criterion2(y, x).item()) # 二者输出数值完全相同
优化建议
直接使用torch.exp计算容易出现数值溢出问题,实际使用时建议用torch.nn.functional.log_softmax来计算对数概率,数值稳定性更好。
内容的提问来源于stack exchange,提问作者Minions
相关产品推荐
相关产品推荐

