自定义交叉熵损失myCEE与nn.CrossEntropyLoss差异及nan问题咨询
自定义交叉熵损失与官方实现差异及NaN问题原因
核心差异与NaN产生的根本原因
- 数值稳定性处理缺失
你的实现中直接调用torch.exp(outputs)计算指数,当模型输出的logits存在较大正值时,torch.exp()会直接上溢为inf,后续torch.log(torch.sum(exp, dim=1))也会得到inf,最终A-B的计算结果就会变为inf或NaN。
官方nn.CrossEntropyLoss内部实现了LogSumExp数值优化技巧:计算时会先减去每个样本logits的最大值,再做指数、求和、对数计算,既完全不影响损失的数学正确性,又从根源上避免了指数上溢的问题,同时也兼容了logits全为极小值导致的指数下溢边界场景。 - 前向逻辑的边界兼容差异
你测试时两者结果完全一致,是因为测试所用的样本logits数值范围较小,没有触发溢出的边界条件,所以前向计算结果匹配。但训练过程中模型参数不断更新,logits的数值范围会动态波动,一旦出现异常的极大/极小值,你的实现就会出现计算错误,而官方实现可以稳定处理。
增加卷积层后未报错的原因
增加更多卷积层后,通常会伴随更多非线性激活、批量归一化(BN)层或者参数正则约束,会压缩模型输出logits的数值范围,让logits不会出现足以触发指数溢出的极大/极小值,暂时避开了问题触发条件,但并未解决你实现的损失函数本身的数值稳定性缺陷,后续训练如果出现异常logits值仍然可能出现NaN。
内容的提问来源于stack exchange,提问作者Jake
相关产品推荐
相关产品推荐

