PyTorch CNN图像边缘检测模型训练后输出异常问题排查
解决边缘检测模型输出固定位置的问题
嘿,我仔细看了你的代码和问题描述,发现几个核心问题导致模型不管输入什么图像都输出固定位置的边缘——说白了就是你的模型结构和训练逻辑根本没适配边缘检测这个任务,咱们一步步来捋:
1. 最致命的问题:CNN网络结构完全不符合边缘检测需求
你的Net最后输出的是18维张量,这是分类任务的输出形式,但边缘检测是像素级预测任务,需要输出和输入图像同尺寸的特征图(每个像素对应是否是边缘)。另外还有两个硬伤:
- 硬编码的全连接层输入维度
293904:这个值完全依赖于输入图像的固定尺寸,一旦图像尺寸变化就会直接报错,而且全连接层会丢失空间信息,完全不适合像素级任务。 - 最后一行
x = x.view(18):强制把输出固定成18维,不管输入的batch大小和图像尺寸,这直接导致模型输出和输入图像的空间信息完全脱节,自然只会输出固定结果。
2. 训练代码的关键疏漏
- 没有启用训练模式:虽然注释了
net.train(),但实际没执行,PyTorch里像BatchNorm、Dropout这类层在训练和评估模式下行为完全不同,必须显式设置net.train()确保训练正常。 - 数据加载逻辑不合理:直接从列表按顺序取数据,没有打乱,模型很容易过拟合到固定顺序的样本,而且没有用
DataLoader做批量训练,效率低还容易出现维度错误。 - Checkpoint保存时机错误:
if i % 100 == 1的逻辑会跳过前100步的保存,而且触发时机很奇怪,应该改成(i+1) % 100 == 0来每100步保存一次。
3. 测试代码的问题
- 用训练集做测试:这根本看不出模型的泛化能力,必须用独立的测试数据集。
- 没有禁用梯度计算:测试时不需要计算梯度,必须用
with torch.no_grad()包裹前向传播,否则会浪费内存甚至影响结果。 - 输出处理完全不对:你的模型输出是18维张量,根本不是边缘图,自然无法得到正确的边缘位置。
修正后的完整代码示例
适配边缘检测的全卷积网络结构
import torch import torch.nn as nn import torch.nn.functional as F class EdgeDetectionNet(nn.Module): def __init__(self): super(EdgeDetectionNet, self).__init__() # 编码器:下采样提取边缘特征 self.conv1 = nn.Conv2d(3, 16, 3, padding=1) self.conv2 = nn.Conv2d(16, 32, 3, padding=1) self.conv3 = nn.Conv2d(32, 64, 3, padding=1) self.pool = nn.MaxPool2d(2, 2, return_indices=True) # 记录池化索引,用于后续反池化恢复尺寸 # 解码器:上采样恢复到输入图像尺寸 self.unpool = nn.MaxUnpool2d(2, 2) self.deconv1 = nn.Conv2d(64, 32, 3, padding=1) self.deconv2 = nn.Conv2d(32, 16, 3, padding=1) self.deconv3 = nn.Conv2d(16, 1, 3, padding=1) # 输出单通道边缘图(0-1表示边缘概率) def forward(self, x): # 编码阶段 x1 = F.relu(self.conv1(x)) x1_pool, idx1 = self.pool(x1) x2 = F.relu(self.conv2(x1_pool)) x2_pool, idx2 = self.pool(x2) x3 = F.relu(self.conv3(x2_pool)) # 解码阶段 x_unpool2 = self.unpool(x3, idx2, output_size=x2.size()) x_deconv2 = F.relu(self.deconv1(x_unpool2)) x_unpool1 = self.unpool(x_deconv2, idx1, output_size=x1.size()) x_deconv1 = F.relu(self.deconv2(x_unpool1)) output = torch.sigmoid(self.deconv3(x_deconv1)) # 用sigmoid把输出限制在0-1区间 return output
修正后的训练代码
import os import torch from torch.utils.data import DataLoader, Dataset # 自定义数据集类,处理图像和标签的维度转换 class EdgeDataset(Dataset): def __init__(self, inputs, labels): self.inputs = inputs self.labels = labels def __len__(self): return len(self.inputs) def __getitem__(self, idx): # 把图像从H,W,C转成PyTorch需要的C,H,W格式,转成float32 img = torch.as_tensor(self.inputs[idx], dtype=torch.float32).permute(2, 0, 1) # 给标签增加通道维度(变成1,H,W),转成float32 label = torch.as_tensor(self.labels[idx], dtype=torch.float32).unsqueeze(0) return img, label # 初始化数据集和DataLoader,开启打乱和批量训练 train_dataset = EdgeDataset(train_input, train_list) train_loader = DataLoader(train_dataset, batch_size=8, shuffle=True) PATH = './checkpoint.pth' net = EdgeDetectionNet().cuda() # 边缘检测用二分类交叉熵损失(因为每个像素是边缘/非边缘的二分类) criterion = nn.BCELoss() optimizer = torch.optim.Adam(net.parameters(), lr=0.001) start_epoch = 0 start_i = 0 # 加载预训练checkpoint if os.path.isfile(PATH): print(f"加载checkpoint '{PATH}' ...") checkpoint = torch.load(PATH) start_epoch = checkpoint['epoch'] start_i = checkpoint['i'] net.load_state_dict(checkpoint['state_dict']) optimizer.load_state_dict(checkpoint['optimizer']) # 同时保存优化器状态,断点续训更稳定 print(f"=> 成功加载checkpoint,已训练 {checkpoint['epoch']} 轮,{checkpoint['i']} 步") else: print('开始新训练') num_epochs = 50 for epoch in range(start_epoch, num_epochs): net.train() # 显式设置为训练模式 running_loss = 0.0 for i, (inputs, labels) in enumerate(train_loader): inputs = inputs.cuda() labels = labels.cuda() # 前向传播 outputs = net(inputs) loss = criterion(outputs, labels) # 反向传播+优化 optimizer.zero_grad() loss.backward() optimizer.step() running_loss += loss.item() # 每100步打印损失并保存checkpoint if (i + 1) % 100 == 0: print(f'Epoch [{epoch+1}/{num_epochs}], Step [{i+1}/{len(train_loader)}], Loss: {running_loss/100:.4f}') torch.save({ 'epoch': epoch + 1, 'i': start_i + i + 1, 'state_dict': net.state_dict(), 'optimizer': optimizer.state_dict() }, PATH) running_loss = 0.0
修正后的测试代码
PATH = './checkpoint.pth' model = EdgeDetectionNet().cuda() if os.path.isfile(PATH): print('加载checkpoint进行测试...') checkpoint = torch.load(PATH) model.load_state_dict(checkpoint['state_dict']) model.eval() # 切换到评估模式 # 初始化测试数据集(注意要用独立的测试集,不能用训练集) test_dataset = EdgeDataset(test_input, test_list) test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False) with torch.no_grad(): # 禁用梯度计算,节省内存 for inputs, labels in test_loader: inputs = inputs.cuda() outputs = model(inputs) # 转换为numpy数组,去除通道维度,得到H,W的边缘图 edge_result = outputs.cpu().squeeze().numpy() # 这里可以添加可视化或保存代码,比如用matplotlib # import matplotlib.pyplot as plt # plt.imshow(edge_result, cmap='gray') # plt.savefig(f'edge_result_{idx}.png')
关键总结
- 边缘检测是像素级任务,必须用全卷积网络(FCN)结构,不能用分类网络的全连接层输出固定维度。
- 训练时一定要打乱数据、用
DataLoader批量处理,显式设置训练模式。 - 测试时要用独立测试集,禁用梯度,确保输出是和输入同尺寸的特征图。
内容的提问来源于stack exchange,提问作者CHOIKHSS
相关产品推荐
相关产品推荐

