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

PyTorch中MaxPool2d设return_indices=True时报错及解码器问题求助

PyTorch自编码器报错解决:MaxPool2d/MaxUnpool2d索引问题

问题分析

  1. 编码器初始报错原因:当nn.MaxPool2d设置return_indices=True时,该层会返回**(输出张量, 池化索引)**两个结果,但nn.Sequential容器无法自动处理多返回值,直接在forward里调用self.encoder_cnn(x)会得到一个元组,后续的Flatten层无法处理元组,导致报错。
  2. 解码器报错原因: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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 16:32:24