PyTorch中MaxPool2d设return_indices=True时报错及解码器问题求助
PyTorch自编码器报错解决:MaxPool2d/MaxUnpool2d索引问题
问题分析
- 编码器初始报错原因:当
nn.MaxPool2d设置return_indices=True时,该层会返回**(输出张量, 池化索引)**两个结果,但nn.Sequential容器无法自动处理多返回值,直接在forward里调用self.encoder_cnn(x)会得到一个元组,后续的Flatten层无法处理元组,导致报错。 - 解码器报错原因:
nn.MaxUnpool2d必须接收对应MaxPool2d生成的索引才能完成反池化操作,如果你去掉了编码器里的return_indices=True,就没有生成索引,解码器调用MaxUnpool2d时自然会提示缺少indices参数。
解决方案
方案1:保留MaxPool2d+MaxUnpool2d(需要传递索引)
修改编码器,手动拆分CNN层,保存池化索引并和编码一起返回;解码器接收索引,在对应位置传入MaxUnpool2d。
修改后的编码器:
import torch import torch.nn as nn # B = Batch size # encoder (B, 3, 224, 224) => (B, 8), 同时返回两次池化的索引 class Encoder(nn.Module): def __init__(self): super().__init__() # 拆分CNN层,把带return_indices的MaxPool单独处理 self.conv1 = nn.Conv2d(3, 8, kernel_size=3, stride=1, padding=0) self.relu1 = nn.ReLU(True) self.conv2 = nn.Conv2d(8, 16, kernel_size=3, stride=2, padding=0) self.relu2 = nn.ReLU(True) self.bn1 = nn.BatchNorm2d(16) self.pool1 = nn.MaxPool2d(2, return_indices=True) self.conv3 = nn.Conv2d(16, 32, kernel_size=3, stride=2, padding=1) self.relu3 = nn.ReLU(True) self.conv4 = nn.Conv2d(32, 64, kernel_size=3, stride=2, padding=1) self.relu4 = nn.ReLU(True) self.bn2 = nn.BatchNorm2d(64) self.pool2 = nn.MaxPool2d(2, return_indices=True) self.flat = nn.Flatten(start_dim=1) self.encoder_fc = nn.Sequential( nn.Linear(64*7*7, 1024), nn.ReLU(True), nn.Linear(1024, 8), nn.Sigmoid() ) def forward(self, x): x = self.conv1(x) x = self.relu1(x) x = self.conv2(x) x = self.relu2(x) x = self.bn1(x) x, idx1 = self.pool1(x) # 保存第一次池化索引 x = self.conv3(x) x = self.relu3(x) x = self.conv4(x) x = self.relu4(x) x = self.bn2(x) x, idx2 = self.pool2(x) # 保存第二次池化索引 x_flat = self.flat(x) codes = self.encoder_fc(x_flat) return codes, idx1, idx2 # 返回编码和两个索引
修改后的解码器:
# decoder (B, 8) + 索引 => (B, 3, 224, 224) class Decoder(nn.Module): def __init__(self): super().__init__() self.decoder_fc = nn.Sequential( nn.Linear(8, 1024), nn.ReLU(True), nn.Linear(1024, 64*7*7), nn.ReLU(True) ) self.unflat = nn.Unflatten(dim=1, unflattened_size=(64, 7, 7)) # 拆分CNN层,手动传入索引到MaxUnpool2d self.unpool2 = nn.MaxUnpool2d(2) self.bn2 = nn.BatchNorm2d(64) self.deconv1 = nn.ConvTranspose2d(64, 32, kernel_size=3, stride=2, padding=1) self.relu1 = nn.ReLU(True) self.deconv2 = nn.ConvTranspose2d(32, 16, kernel_size=3, stride=2, padding=1) self.unpool1 = nn.MaxUnpool2d(2) self.bn1 = nn.BatchNorm2d(16) self.relu2 = nn.ReLU(True) self.deconv3 = nn.ConvTranspose2d(16, 8, kernel_size=3, stride=2, padding=0) self.relu3 = nn.ReLU(True) self.deconv4 = nn.ConvTranspose2d(8, 3, kernel_size=3, stride=1, padding=0) def forward(self, codes, idx1, idx2): x = self.decoder_fc(codes) x = self.unflat(x) x = self.unpool2(x, idx2) # 传入第二次池化的索引 x = self.bn2(x) x = self.deconv1(x) x = self.relu1(x) x = self.deconv2(x) x = self.unpool1(x, idx1) # 传入第一次池化的索引 x = self.bn1(x) x = self.relu2(x) x = self.deconv3(x) x = self.relu3(x) x = self.deconv4(x) return x
测试代码:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu") encoder = Encoder().to(device) decoder = Decoder().to(device) # 假设train_data[0]是(3,224,224)的张量 test_img = torch.unsqueeze(train_data[0], dim=0).to(device) codes, idx1, idx2 = encoder(test_img) output = decoder(codes, idx1, idx2) print(output.shape) # 应该输出torch.Size([1, 3, 224, 224])
方案2:替换为无需索引的上采样(更简单)
如果不想处理索引,可以把编码器里的MaxPool2d去掉return_indices=True(直接用nn.MaxPool2d(2)),同时把解码器里的MaxUnpool2d换成nn.Upsample或者nn.ConvTranspose2d来实现上采样,避免依赖索引。
修改后的解码器(对应无索引的编码器):
class Decoder(nn.Module): def __init__(self): super().__init__() self.decoder_fc = nn.Sequential( nn.Linear(8, 1024), nn.ReLU(True), nn.Linear(1024, 64*7*7), nn.ReLU(True) ) self.unflat = nn.Unflatten(dim=1, unflattened_size=(64, 7, 7)) self.decoder_cnn = nn.Sequential( # 把MaxUnpool2d换成Upsample,scale_factor对应原来的池化 stride nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True), nn.BatchNorm2d(64), nn.ConvTranspose2d(64, 32, kernel_size=3, stride=2, padding=1), nn.ReLU(True), nn.ConvTranspose2d(32, 16, kernel_size=3, stride=2, padding=1), nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True), nn.BatchNorm2d(16), nn.ReLU(True), nn.ConvTranspose2d(16, 8, kernel_size=3, stride=2, padding=0), nn.ReLU(True), nn.ConvTranspose2d(8, 3, kernel_size=3, stride=1, padding=0) ) def forward(self, x): x = self.decoder_fc(x) x = self.unflat(x) x = self.decoder_cnn(x) return x
对应的编码器(去掉return_indices):
class Encoder(nn.Module): def __init__(self): super().__init__() self.encoder_cnn = nn.Sequential( nn.Conv2d(3, 8, kernel_size=3, stride=1, padding=0), nn.ReLU(True), nn.Conv2d(8, 16, kernel_size=3, stride=2, padding=0), nn.ReLU(True), nn.BatchNorm2d(16), nn.MaxPool2d(2), # 去掉return_indices nn.Conv2d(16, 32, kernel_size=3, stride=2, padding=1), nn.ReLU(True), nn.Conv2d(32, 64, kernel_size=3, stride=2, padding=1), nn.ReLU(True), nn.BatchNorm2d(64), nn.MaxPool2d(2), # 去掉return_indices ) self.flat = nn.Flatten(start_dim=1) self.encoder_fc = nn.Sequential( nn.Linear(64*7*7, 1024), nn.ReLU(True), nn.Linear(1024, 8), nn.Sigmoid() ) def forward(self, x): x = self.encoder_cnn(x) x = self.flat(x) x = self.encoder_fc(x) return x
测试代码:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu") encoder = Encoder().to(device) decoder = Decoder().to(device) test_img = torch.unsqueeze(train_data[0], dim=0).to(device) codes = encoder(test_img) output = decoder(codes) print(output.shape) # 应该输出torch.Size([1, 3, 224, 224])
内容的提问来源于stack exchange,提问作者monstar
相关产品推荐
相关产品推荐

