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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.28 18:02:52