VAE输出形状与输入不匹配问题求助
问题原因分析
你的VAE解码器是基于固定尺寸的特征图((6,12,12))构建的:经过4次scale=2的上采样后,输出尺寸固定为(6×16, 12×16, 12×16) = (96,192,192)。只有当输入尺寸经过编码器4次下采样后恰好得到(6,12,12)时(即输入为(96,192,192)),解码器才能输出和输入一致的尺寸;其他输入尺寸都会因为解码器的固定输出逻辑,导致输出与输入尺寸不匹配。
你提到输入(80,96,80)能正常工作,大概率是预处理时将其resize到了(96,192,192),或者后续对输出做了裁剪,掩盖了尺寸不匹配的问题。
解决方案
方案1:输出阶段自适应调整(最简单)
在VAE的forward方法末尾,将解码器输出直接调整为输入的原始尺寸,无需修改模型核心结构:
def forward(self, x): original_size = x.size()[2:] # 获取输入的空间维度 (h, w, d) x = self.encoder(x) x = torch.flatten(x, start_dim=1) z_mean = self.z_mean(x) z_log_sigma = self.z_log_sigma(x) z = z_mean.to("cpu") + z_log_sigma.exp().to("cpu") * self.epsilon.to("cpu") y = self.decoder(z) # 将输出resize为输入原始尺寸 y = F.interpolate(y, size=original_size, mode='trilinear', align_corners=False) return y, z_mean, z_log_sigma
方案2:动态适配输入尺寸(更灵活)
修改模型初始化逻辑,让编码器输出尺寸根据输入动态计算,解码器基于该尺寸构建,确保上采样后还原输入尺寸:
- 修改VAE类的初始化方法,传入输入尺寸并计算编码器输出尺寸:
class VAE(nn.Module): def __init__(self, input_size, latent_dim=128): super(VAE, self).__init__() self.device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu') self.latent_dim = latent_dim # 计算编码器4次池化后的输出尺寸 h, w, d = input_size for _ in range(4): # 池化尺寸公式:floor((h + 2*padding - kernel)/stride) +1 = (h+1)//2 h = (h + 1) // 2 w = (w + 1) // 2 d = (d + 1) // 2 self.encoder_output_size = (h, w, d) # 根据编码器输出尺寸定义线性层 self.z_mean = nn.Linear(256 * h * w * d, latent_dim) self.z_log_sigma = nn.Linear(256 * h * w * d, latent_dim) self.epsilon = torch.normal(size=(1, latent_dim), mean=0, std=1.0, device=self.device) self.encoder = Encoder() self.decoder = Decoder(latent_dim, self.encoder_output_size) self.reset_parameters()
- 修改Decoder类,接收编码器输出尺寸作为参数:
class Decoder(nn.Module): """ Decoder Module """ def __init__(self, latent_dim, encoder_output_size): super(Decoder, self).__init__() self.latent_dim = latent_dim self.h, self.w, self.d = encoder_output_size self.linear_up = nn.Linear(latent_dim, 256 * self.h * self.w * self.d) self.relu = nn.ReLU() self.upsize4 = up_conv(ch_in=256, ch_out=128, k_size=1, scale=2) self.res_block4 = ResNet_block(ch=128, k_size=3, num_groups=16) self.upsize3 = up_conv(ch_in=128, ch_out=64, k_size=1, scale=2) self.res_block3 = ResNet_block(ch=64, k_size=3, num_groups=16) self.upsize2 = up_conv(ch_in=64, ch_out=32, k_size=1, scale=2) self.res_block2 = ResNet_block(ch=32, k_size=3, num_groups=16) self.upsize1 = up_conv(ch_in=32, ch_out=1, k_size=1, scale=2) self.res_block1 = ResNet_block(ch=1, k_size=3, num_groups=1) self.reset_parameters() def forward(self, x): x4_ = self.linear_up(x) x4_ = self.relu(x4_) x4_ = x4_.view(-1, 256, self.h, self.w, self.d) x4_ = self.upsize4(x4_) x4_ = self.res_block4(x4_) x3_ = self.upsize3(x4_) x3_ = self.res_block3(x3_) x2_ = self.upsize2(x3_) x2_ = self.res_block2(x2_) x1_ = self.upsize1(x2_) x1_ = self.res_block1(x1_) return x1_
- 初始化模型时传入目标输入尺寸:
# 示例:输入尺寸为(94,180,180) vae = VAE(input_size=(94,180,180), latent_dim=128)
方案3:固定输入尺寸(适合特定场景)
如果业务允许固定输入尺寸,可将所有输入预处理为(96,192,192)(即解码器固定输出尺寸),这样编码器下采样4次后恰好得到(6,12,12),解码器上采样后完全还原输入尺寸。
内容的提问来源于stack exchange,提问作者Julio Lopez
相关产品推荐
相关产品推荐

