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

PyTorch Hourglass模型张量尺寸不匹配RuntimeError问题求助

解决Deep Image Prior 4通道图像修复的张量尺寸不匹配问题

核心原因

Deep Image Prior的U-Net结构依赖输入尺寸为2的整数次幂(比如64、128、256这类),如果你的4通道输入图像尺寸不符合这个要求,下采样(池化/卷积)再上采样后的特征图尺寸,会和跳跃连接的特征图出现1像素的偏差,导致torch.cat时维度不匹配(比如报错里的17 vs 18)。另外,修改通道数时如果没同步调整所有卷积层的输入输出通道,也可能引发这类问题。

具体解决步骤

1. 把输入图像调整为2的幂次尺寸

预处理阶段直接将图像resize到最近的2的幂次尺寸,比如原尺寸300x300就改成256x256或512x512:

import torchvision.transforms as transforms

# 假设img是你的4通道PIL图像或张量
resize_transform = transforms.Resize((256, 256))
resized_img = resize_transform(img)

要是想保留原图像比例,可以先填充到2的幂次尺寸,修复完成后再裁剪回原大小:

import torch
import math

def pad_to_power_of_two(tensor):
    h, w = tensor.shape[1], tensor.shape[2]
    new_h = 2 ** int(math.ceil(math.log2(h)))
    new_w = 2 ** int(math.ceil(math.log2(w)))
    
    pad_h = new_h - h
    pad_w = new_w - w
    
    # 对H、W维度做对称填充(4通道张量维度为[C, H, W])
    padded = torch.nn.functional.pad(tensor, (pad_w//2, pad_w - pad_w//2, pad_h//2, pad_h - pad_h//2))
    return padded, (h, w)

# 预处理时填充
padded_img, original_size = pad_to_power_of_two(input_tensor)
# 修复后裁剪回原尺寸
restored_img = restored_img[:, :original_size[0], :original_size[1]]

2. 同步调整U-Net所有卷积层的通道数

确保你修改了所有卷积模块的输入输出通道,包括编码器下采样、解码器上采样以及跳跃连接对应的部分:

import torch.nn as nn

class UNet(nn.Module):
    def __init__(self, in_channels=4, out_channels=4):
        super().__init__()
        # 编码器部分
        self.encoder1 = self.conv_block(in_channels, 64)
        self.encoder2 = self.conv_block(64, 128)
        self.encoder3 = self.conv_block(128, 256)
        self.encoder4 = self.conv_block(256, 512)
        self.encoder5 = self.conv_block(512, 1024)
        
        # 解码器部分
        self.decoder1 = self.upconv_block(1024, 512)
        self.decoder2 = self.conv_block(512+512, 256)  # 注意跳跃连接的通道数相加
        self.decoder3 = self.conv_block(256+256, 128)
        self.decoder4 = self.conv_block(128+128, 64)
        self.decoder5 = nn.Conv2d(64+64, out_channels, kernel_size=3, padding=1)
        
    def conv_block(self, in_ch, out_ch):
        return nn.Sequential(
            nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),
            nn.ReLU(inplace=True),
            nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1),
            nn.ReLU(inplace=True)
        )
    
    def upconv_block(self, in_ch, out_ch):
        return nn.Sequential(
            nn.ConvTranspose2d(in_ch, out_ch, kernel_size=2, stride=2),
            nn.ReLU(inplace=True)
        )
    
    def forward(self, x):
        # 编码器前向
        skip1 = self.encoder1(x)
        x = nn.MaxPool2d(2)(skip1)
        skip2 = self.encoder2(x)
        x = nn.MaxPool2d(2)(skip2)
        skip3 = self.encoder3(x)
        x = nn.MaxPool2d(2)(skip3)
        skip4 = self.encoder4(x)
        x = nn.MaxPool2d(2)(skip4)
        skip5 = self.encoder5(x)
        
        # 解码器前向
        x = self.decoder1(skip5)
        x = torch.cat([x, skip4], dim=1)
        x = self.decoder2(x)
        x = torch.cat([x, skip3], dim=1)
        x = self.decoder3(x)
        x = torch.cat([x, skip2], dim=1)
        x = self.decoder4(x)
        x = torch.cat([x, skip1], dim=1)
        x = self.decoder5(x)
        return x

重点确认:每个torch.cat操作前,上采样后的特征图和对应跳跃连接的特征图,H、W维度完全一致。

3. 替换上采样方式(可选)

如果不想修改输入尺寸,可以把转置卷积换成双线性插值+卷积,这种方式更灵活,能减少尺寸偏差:

def upconv_block(self, in_ch, out_ch):
    return nn.Sequential(
        nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
        nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),
        nn.ReLU(inplace=True)
    )

4. 打印张量尺寸定位问题

在模型的forward方法里,每次上采样后打印张量尺寸,和对应跳跃连接的张量对比,明确哪一层出现偏差:

def forward(self, x):
    skip1 = self.encoder1(x)
    print(f"skip1 size: {skip1.shape}")
    x = nn.MaxPool2d(2)(skip1)
    skip2 = self.encoder2(x)
    print(f"skip2 size: {skip2.shape}")
    # ... 其他层同理
    
    x = self.decoder1(skip5)
    print(f"up5 size: {x.shape}, skip4 size: {skip4.shape}")  # 对应报错的up_5和skip_5
    x = torch.cat([x, skip4], dim=1)
    # ... 后续层

验证方法

修改后先单独测试模型前向传播,传入一个4通道、2的幂次尺寸的张量,确认无报错:

model = UNet(in_channels=4, out_channels=4)
test_input = torch.randn(1, 4, 256, 256)  # batch_size=1,4通道,256x256
output = model(test_input)
print(f"Output shape: {output.shape}")  # 应该输出(1,4,256,256)

内容的提问来源于stack exchange,提问作者Khubaib Khawar

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 09:05:33