Unet模型训练时Test Loss与Dice系数出现NaN值的解决求助
问题分析与解决步骤
1. 模型输出未经过Sigmoid就直接取阈值
你的模型输出是Logits(未经过Sigmoid激活的原始输出),直接用torch.round(output)完全错误:
- Logits范围是(-∞, +∞),比如输出10.0时round后为10,不符合分割任务0/1的标签要求;
- 当Logits趋近于±∞时,后续计算Dice或Loss极易出现NaN。
修改方案:
先对输出做Sigmoid激活,转换为[0,1]之间的概率值后再取阈值:
pred = torch.round(torch.sigmoid(output))
2. Dice系数计算的数值稳定性问题
当某个batch中无病灶(target全0)且预测也全0时,Dice系数公式(2*TP/(2*TP+FP+FN))会出现分子分母均为0的情况,直接导致NaN。
修改方案:
在compute_meandice函数中添加极小常数(如1e-6)避免除以0:
def compute_meandice(pred, target, include_background=False): smooth = 1e-6 intersection = torch.sum(pred * target) union = torch.sum(pred) + torch.sum(target) dice = (2. * intersection + smooth) / (union + smooth) # 保留原函数中多类/背景处理逻辑 return dice
如果使用第三方库(如MONAI)的compute_meandice,可直接设置其smooth参数(若支持)。
3. BCEWithLogitsLoss的数值稳定性与输入合法性
- 目标类型修正:BCEWithLogitsLoss要求目标值为浮点型,即使是0/1标签也需转换,修改数据加载后的目标处理:
target = target.unsqueeze(1).float() # 替换原有的target.unsqueeze(1) - 学习率调整:0.001的学习率对Unet偏高,易引发梯度爆炸,导致模型参数极端化、Logits趋近于±∞,最终Loss出现NaN。建议降低学习率:
optimizer = optim.Adam(model.parameters(), lr=1e-4) - Loss累加方式优化:直接累加张量
loss可能保留计算图中的NaN,改用数值累加更安全:totalTrainLoss += loss.item() # 替换原totalTrainLoss += loss totalTestLoss += loss.item()
4. 数据合法性检查
检查训练/测试数据(train_x, train_y, test_x, test_y)是否存在NaN或异常值:
在ProcessDataset的__getitem__方法中添加检查逻辑:
def __getitem__(self, idx): x = self.x[idx] y = self.y[idx] assert not torch.isnan(x).any(), f"数据索引{idx}存在NaN" assert not torch.isnan(y).any(), f"标签索引{idx}存在NaN" return x, y
同时确保标签target的取值范围为[0,1],整数0/1需转成浮点型。
5. 训练过程中的NaN监控
在训练循环中添加检查,快速定位NaN出现的环节:
# 训练阶段 loss = criterion(output, target) if torch.isnan(loss): print(f"Epoch {epoch+1}, Batch {i+1}出现NaN Loss") print("输入数据极值:", data.min().item(), data.max().item()) print("标签极值:", target.min().item(), target.max().item()) print("模型输出极值:", output.min().item(), output.max().item()) break
修改后的测试阶段代码示例
with torch.no_grad(): model.eval() total_dice = 0 totalTestLoss = 0 for data, target in test_loader: data, target = data.to(device), target.to(device) target = target.unsqueeze(1).float() output = model(data) loss = criterion(output, target) if torch.isnan(loss): print(f"Epoch {epoch+1}出现NaN测试Loss") continue totalTestLoss += loss.item() pred = torch.round(torch.sigmoid(output)) dice = compute_meandice(pred, target, include_background=False) if torch.isnan(dice): print(f"Epoch {epoch+1}出现NaN Dice") print("当前batch标签总和:", target.sum().item()) print("当前batch预测总和:", pred.sum().item()) continue total_dice += dice.item() avg_dice = total_dice / len(test_loader) avgTestLoss = totalTestLoss / testSteps
内容的提问来源于stack exchange,提问作者anastasia
相关产品推荐
相关产品推荐

