PyTorch ResNet图像分割训练中损失恒为零的问题解决
解决PyTorch二分类图像分割中损失为0但准确率极高的问题
嘿,我仔细扒了你的代码和问题描述,这情况根本不是模型真的完美,而是代码里几个关键逻辑错得离谱——假高准确率和零损失都是计算错误导致的。咱们一步步来修复:
核心问题排查与修复
1. 标签维度错误(直接导致损失异常)
你在训练循环里写了y = y[:, 0,0],这相当于把每个样本的整幅掩码只截取了左上角一个像素的标签!而图像分割是像素级任务,模型输出应该是每个像素的类别预测(形状[Nbatch, 2, H, W]),标签也应该保持[Nbatch, H, W]的像素级维度。这种维度不匹配加上类别权重的设置,直接让CrossEntropyLoss计算出零损失。
- 修复:删掉
y = y[:, 0,0]这行,保留标签的原始像素级形状。
2. 准确率计算完全错误(假高准确率的来源)
你写的_, preds = torch.max(X, 1),居然是对输入图像X取通道最大值,而不是模型的预测结果prediction!这就相当于拿输入RGB通道的最大值来和标签做对比,刚好巧合出现了高匹配度,所以才会有0.99的假准确率。
- 修复:改成
_, preds = torch.max(prediction, 1),这样才是取模型预测的类别维度(第二维)的最大值作为每个像素的预测类别;同时删掉preds = preds[:,0,0],让preds保持[Nbatch, H, W]的像素级形状,和标签维度匹配。
3. 损失累加与权重应用错误
你设置了reduce=False让损失返回每个像素的损失值,但累加时只取了loss.data[0](第一个样本的损失),而且定义的y_weight完全没用到。
- 修复:
- 如果要应用像素权重(比如边缘权重),把损失改成
loss = criterion(prediction, y) * y_weight; - 累加损失时用
running_loss += loss.mean().item(),确保累加的是整个批次的平均损失; - 统计准确率时,要计算所有像素的正确数,而不是样本数:总正确数除以总像素数(
样本数 × patch_size × patch_size)。
- 如果要应用像素权重(比如边缘权重),把损失改成
4. 数据变换冗余问题
你同时用了RandomCrop和RandomResizedCrop,相当于先裁剪再随机缩放,会让样本的尺寸变换不可控;而且颜色增强的参数(brightness=0等)等于没开,无法增加数据多样性。
- 修复:去掉重复的
RandomCrop,只保留RandomResizedCrop;适当调整ColorJitter的参数(比如brightness=0.2)增强数据随机性。
修正后的关键代码片段
训练循环核心部分
for ii , (X, y, y_weight) in enumerate(dataLoader[phase]): optim.zero_grad() X = X.to(device) # [Nbatch, 3, H, W] y_weight = y_weight.type('torch.FloatTensor').to(device) y = y.type('torch.LongTensor').to(device) # [Nbatch, H, W],类别索引为(0, 1) with torch.set_grad_enabled(phase == 'train'): prediction = model_ft(X) # 确保输出形状为[Nbatch, 2, H, W] # 去掉错误的标签维度压缩 loss = criterion(prediction, y) # 应用像素权重(如果需要) loss = loss * y_weight # 基于模型预测获取类别,而非输入图像 _, preds = torch.max(prediction, 1) # preds形状[Nbatch, H, W] if phase=="train": loss.mean().backward() optim.step() # 累加整个批次的平均损失 running_loss += loss.mean().item() # 统计所有像素的正确预测数 running_corrects += torch.sum(preds == y).item() # 计算epoch级别的损失和准确率 epoch_loss = running_loss / len(dataLoader[phase]) total_pixels = len(dataLoader[phase].dataset) * patch_size * patch_size epoch_acc = running_corrects / total_pixels print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))
损失函数与类别权重优化
# 更直观的二分类权重计算:反向类别频率 class_count = dataset["train"].numpixels[1,0:2] class_weight = torch.from_numpy((class_count.sum() - class_count) / class_count.sum()).type('torch.FloatTensor').to(device) print(class_weight) # 不需要每个像素损失的话,不用设置reduce=False,默认返回批次平均损失 criterion = nn.CrossEntropyLoss(weight=class_weight, ignore_index=ignore_index)
总结
这些错误本质是把像素级的图像分割任务当成了样本级的分类任务来处理,导致所有指标计算都偏离了实际情况。修正后你应该能看到损失正常波动,准确率也会反映模型的真实性能。
内容的提问来源于stack exchange,提问作者Sarah
相关产品推荐
相关产品推荐

