为何CrossEntropyLoss等价于LogSoftmax+NLLLoss的测试断言失败
问题原因分析
- CrossEntropyLoss输入不符合要求
CrossEntropyLoss的内部逻辑已经包含了LogSoftmax计算,它要求输入是模型最后一层输出的原始logits,不需要额外做归一化。你给model_ce的最后一层加了nn.LogSoftmax(dim=1),相当于把已经做过LogSoftmax的结果传给CrossEntropyLoss,它内部会再做一次归一化,计算出来的损失自然和NLLLoss的结果不一致。 - 两个模型参数未对齐
PyTorch的nn.Linear层默认是随机初始化权重的,你分别实例化了两个结构相同的模型,它们的权重参数初始值完全不同,哪怕输入完全一致,输出结果也会有差异,损失值自然对不上。要验证等价性需要把其中一个模型的权重直接赋值给另一个,保证参数完全一致。 - 浮点数等值判断方式错误
torch.eq是严格的数值相等判断,而浮点数计算过程中会有微小的精度误差,哪怕逻辑上完全相等的两个值,也可能因为计算顺序或者精度截断出现末尾几位的差异,用torch.eq会直接判定为不相等,应该用允许一定误差范围的torch.allclose做对比。
修正后的测试代码
import torch import torch.nn as nn # NLLLoss搭配的模型:末尾需要加LogSoftmax model_nll = nn.Sequential(nn.Linear(3072, 1024), nn.Tanh(), nn.Linear(1024, 512), nn.Tanh(), nn.Linear(512, 128), nn.Tanh(), nn.Linear(128, 2), nn.LogSoftmax(dim=1)) # CrossEntropyLoss搭配的模型:直接输出原始logits,不需要加LogSoftmax model_ce = nn.Sequential(nn.Linear(3072, 1024), nn.Tanh(), nn.Linear(1024, 512), nn.Tanh(), nn.Linear(512, 128), nn.Tanh(), nn.Linear(128, 2)) # 对齐两个模型的权重参数,保证输出完全一致 model_ce.load_state_dict(model_nll.state_dict(), strict=False) loss_fn_ce = nn.CrossEntropyLoss() loss_fn_nll = nn.NLLLoss() t = torch.rand(1, 3072) target = torch.tensor([1]) with torch.no_grad(): loss_nll = loss_fn_nll(model_nll(t), target) loss_ce = loss_fn_ce(model_ce(t), target) # 浮点数对比使用允许误差的allclose assert torch.allclose(loss_nll, loss_ce) print("等价性验证通过,两个损失计算结果一致")
内容的提问来源于stack exchange,提问作者Marcel Coetzee
相关产品推荐
相关产品推荐

