PyTorch中nn.CrossEntropyLoss标签全为ignore_index时损失为nan的解决方法
解决PyTorch CrossEntropyLoss全ignore标签导致nan的问题
以下是几种可行的解决办法:
提前判断全ignore场景
在调用损失函数前,先检查标签是否全部为ignore_index(即-100),如果是直接返回0.0的张量,避免触发内置计算的nan问题:# 假设label为你的标签张量,criterion是已初始化的CrossEntropyLoss if (label == -100).all(): loss_ent = torch.tensor(0.0, device=label.device) else: loss_ent = criterion(output, label)手动筛选有效样本计算损失
跳过所有ignore的标签,只对有效标签对应的输出计算损失,彻底避开全ignore的异常情况:valid_mask = (label != -100) if valid_mask.any(): valid_output = output[valid_mask] valid_label = label[valid_mask] loss_ent = criterion(valid_output, valid_label) else: loss_ent = torch.tensor(0.0, device=label.device)升级PyTorch版本
你当前使用的PyTorch 1.12.0存在全ignore标签下返回nan的已知问题,升级到1.13.0及以上版本后,内置的CrossEntropyLoss会处理这种场景,直接返回0而不是nan,升级前注意确认环境依赖的兼容性。
内容的提问来源于stack exchange,提问作者ScoTT
相关产品推荐
相关产品推荐

