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

PyTorch自编码器图像重建边缘伪影问题求助

自编码器重建边缘伪影问题求助

输入为100×100像素的灰度图像,重建后的图像右侧和底部出现伪影。

自编码器架构

class Encoder(nn.Module):

    def __init__(self,
                 encoded_space_dim,
                 lin_dims):
        super().__init__()

        self.encoder_cnn = nn.Sequential(
            nn.Conv2d(1, 8, 3, stride=2, padding=1),    #1 image divided to 8 feature maps,                                                                 kernel 3x3
            nn.ReLU(True),
            nn.Conv2d(8, 16, 3, stride=2, padding=1),   #8 feature maps divided to 16, kernel 3x3
            nn.ReLU(True),
            nn.Conv2d(16, 32, 2, stride=2, padding=0),  #16 feature maps divided to 32, kernel 2x2 and padding 0 in order to achieve image size 12x12 pixels
            nn.ReLU(True),
            nn.Conv2d(32, 64, 3, stride=2, padding=1),  #32 feature maps to 64, kernel 3x3
            nn.ReLU(True),
            nn.Conv2d(64, 128, 3, stride=2, padding=1), #64 feature maps to 128, kernel 3x3
            nn.ReLU(True)
        )

        self.flatten = nn.Flatten(start_dim=1)

        self.encoder_lin = nn.Sequential(
            nn.Linear(128*3*3, lin_dims[0]),            #3x3(kernel size) x  number of feature maps in last layer so in this case 128
            nn.ReLU(True),
            nn.Linear(lin_dims[0], lin_dims[1]),
            nn.ReLU(True),
            nn.Linear(lin_dims[1], encoded_space_dim)
        )

    def forward(self, x):
        x = self.encoder_cnn(x)
        x = self.flatten(x)
        x = self.encoder_lin(x)
        return x

class Decoder(nn.Module):

    def __init__(self,
                 encoded_space_dim,
                 lin_dims):
        super().__init__()

        self.decoder_lin = nn.Sequential(
            nn.Linear(encoded_space_dim, lin_dims[0]),
            nn.ReLU(True),
            nn.Linear(lin_dims[0], lin_dims[1]),
            nn.ReLU(True),
            nn.Linear(lin_dims[1], 128*3*3),
            nn.ReLU(True)
        )

        self.decoder_cnn = nn.Sequential(
            nn.ConvTranspose2d(128, 64, 3, stride=2, padding=1, output_padding=1),   #128 maps to 64, kernel 3x3
            nn.ReLU(True),
            nn.ConvTranspose2d(64, 32, 3, stride=2, padding=1, output_padding=1),    #64 maps to 32, kernel 3x3
            nn.ReLU(True),
            nn.ConvTranspose2d(32, 16, 2, stride=2, padding=0, output_padding=1),    #32 maps to 16, kernel 2x2
            nn.ReLU(True),
            nn.ConvTranspose2d(16, 8, 3, stride=2, padding=1, output_padding=1),     #16 maps to 8, kernel 3x3
            nn.ReLU(True),
            nn.ConvTranspose2d(8, 1, 3, stride=2, padding=1, output_padding=1)       #8 maps to 1, kernel 3x3
        )

    def forward(self, x):
        x = self.decoder_lin(x)
        x = x.view(-1, 128, 3, 3)       
        x = self.decoder_cnn(x)
        x = torch.sigmoid(x)
        return x

编码器输出摘要

# print encoder summary
summary(encoder, (1, 100, 100))
----------------------------------------------------------------
        Layer (type)               Output Shape         Param #
================================================================
            Conv2d-1            [-1, 8, 50, 50]              80
              ReLU-2            [-1, 8, 50, 50]               0
            Conv2d-3           [-1, 16, 25, 25]           1,168
              ReLU-4           [-1, 16, 25, 25]               0
            Conv2d-5           [-1, 32, 12, 12]           2,080
              ReLU-6           [-1, 32, 12, 12]               0
            Conv2d-7             [-1, 64, 6, 6]          18,496
              ReLU-8             [-1, 64, 6, 6]               0
            Conv2d-9            [-1, 128, 3, 3]          73,856
             ReLU-10            [-1, 128, 3, 3]               0
          Flatten-11                 [-1, 1152]               0
           Linear-12                 [-1, 1000]       1,153,000
             ReLU-13                 [-1, 1000]               0
           Linear-14                  [-1, 750]         750,750
             ReLU-15                  [-1, 750]               0
           Linear-16                  [-1, 500]         375,500
================================================================
Total params: 2,374,930
Trainable params: 2,374,930
Non-trainable params: 0
----------------------------------------------------------------
Input size (MB): 0.04
...
Forward/backward pass size (MB): 0.62
Params size (MB): 9.06
Estimated Total Size (MB): 9.72
----------------------------------------------------------------

解码器输出摘要

# print decoder summary
summary(decoder, (1, encoded_space_dim))
----------------------------------------------------------------
        Layer (type)               Output Shape         Param #
================================================================
            Linear-1               [-1, 1, 750]         375,750
              ReLU-2               [-1, 1, 750]               0
            Linear-3              [-1, 1, 1000]         751,000
              ReLU-4              [-1, 1, 1000]               0
            Linear-5              [-1, 1, 1152]       1,153,152
              ReLU-6              [-1, 1, 1152]               0
   ConvTranspose2d-7             [-1, 64, 6, 6]          73,792
              ReLU-8             [-1, 64, 6, 6]               0
   ConvTranspose2d-9           [-1, 32, 12, 12]          18,464
             ReLU-10           [-1, 32, 12, 12]               0
  ConvTranspose2d-11           [-1, 16, 25, 25]           2,064
             ReLU-12           [-1, 16, 25, 25]               0
  ConvTranspose2d-13            [-1, 8, 50, 50]           1,160
             ReLU-14            [-1, 8, 50, 50]               0
  ConvTranspose2d-15          [-1, 1, 100, 100]              73
================================================================
Total params: 2,375,455
Trainable params: 2,375,455
Non-trainable params: 0
----------------------------------------------------------------
Input size (MB): 0.00
Forward/backward pass size (MB): 0.68
Params size (MB): 9.06
Estimated Total Size (MB): 9.75
----------------------------------------------------------------

训练配置

loss_fn = nn.L1Loss()
batch_size = 128

#latent space 
encoded_space_dim = 500
enc_dims = [1000, 750]
dec_dims = [750, 1000]

# Encoder & Decoder
encoder = Encoder(
    encoded_space_dim=encoded_space_dim,
    lin_dims=enc_dims
).to(device)

decoder = Decoder(
    encoded_space_dim=encoded_space_dim,
    lin_dims=dec_dims
).to(device)

# Optimizer
params_to_optimize = [
    {'params': encoder.parameters()},
    {'params': decoder.parameters()}
]
optim = torch.optim.Adam(params_to_optimize)

已用公式((n+2p-f)/s)+1计算网络维度(n为图像尺寸,p为padding,s为stride,f为kernel size),调整padding会引发维度不匹配错误。已排除过采样问题(若为过采样,整幅图像都会失真而非仅边缘),目前无思路,恳请提供解决建议。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 12:05:16