You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.27 21:38:11