PyTorch图像补全模型输入报错:通道不匹配与数据类型不兼容问题求助
解决PyTorch图像补全(Outpaint)模型的两类运行时错误
让我一步步帮你排查并解决这两个运行时问题:
错误1:通道数不匹配 (RuntimeError: expected input to have 3 channels, but got 96 channels instead)
问题根源
你用skimage.io.imread()读取的图像默认是HWC格式(高度→宽度→通道),比如你的96x96图像形状是(96, 96, 3),组成批次后就变成了(5, 96, 96, 3)。但PyTorch的卷积层nn.Conv2d要求输入必须是NCHW格式(批次大小→通道数→高度→宽度),此时模型会把第二个维度的96误认为是通道数,和你定义的第一个卷积层输入通道数3完全不匹配,所以抛出了这个错误。
修复方案
在数据预处理阶段把图像维度从HWC转成CHW,同时转成PyTorch张量。修改你的OutpaintDataset类的__getitem__方法:
def __getitem__(self, index): image = io.imread(self.image_names[index]) # 处理灰度图:确保图像是3通道 if image.ndim == 2: image = np.expand_dims(image, axis=-1) image = np.repeat(image, 3, axis=-1) # 调整图像尺寸 input_image = self.custom_resize(image, self.input_size) ground_image = self.custom_resize(image, self.output_size) masked_image = self.outpaint(image=ground_image.copy()) # 关键:将HWC转成CHW,并转成torch.Float32张量 input_image = torch.from_numpy(input_image).permute(2, 0, 1).float() masked_image = torch.from_numpy(masked_image).permute(2, 0, 1).float() ground_image = torch.from_numpy(ground_image).permute(2, 0, 1).float() # 可选:归一化到[-1,1],匹配模型最后一层Tanh的输出范围 input_image = (input_image * 2) - 1.0 masked_image = (masked_image * 2) - 1.0 ground_image = (ground_image * 2) - 1.0 return input_image, masked_image, ground_image
这样处理后,每个样本的形状会变成(3, 96, 96)或(3, 192, 192),组成批次后就是(5, 3, 96, 96),完全符合模型的输入要求。
错误2:数据类型不匹配 (RuntimeError: expected scalar type Double but found Float)
问题根源
PyTorch的模型参数默认是torch.float32(单精度浮点)类型,但你在数据加载代码里把输入强制转成了torch.double(双精度浮点,即torch.float64),导致模型参数和输入数据类型不兼容,从而触发错误。
修复方案
有两种可行的解决方式,推荐第一种:
方式1:保持输入为Float32(推荐)
去掉代码里转成double的语句,让输入和模型参数类型一致。修改你的数据加载循环:
for i, data in enumerate(train_loader, 0): input_image, masked_image, ground_image = data # 因为Dataset里已经处理好维度和类型,直接传入即可 output = generator(masked_image) break
如果暂时不想修改Dataset类,也可以在加载时转成float32而非double:
reshaped = masked_image.permute(0, 3, 1, 2).float() # 用.float()替代.type(torch.double) output = generator(reshaped)
方式2:将模型转为Double类型(不推荐,除非有特殊需求)
如果你确实需要使用双精度浮点训练,可以把整个模型的参数转成double类型:
generator = Generator().double()
额外注意点
skimage.transform.resize会自动把图像像素值归一化到[0, 1],而你的模型最后用了nn.Tanh()(输出范围[-1, 1]),所以我在Dataset的修复代码里加上了归一化到[-1, 1]的步骤,这能让模型的输入输出范围匹配,提升训练效果。- 确保你的
train_loader使用了合适的collate_fn(如果需要),不过默认的collate_fn已经能正确处理我们修改后的张量格式。
内容的提问来源于stack exchange,提问作者Mert Arda Asar
相关产品推荐
相关产品推荐

