CAE卷积自编码器模型入门指南:从理论到PyTorch实战

1次阅读
没有评论

共计 1875 个字符,预计需要花费 5 分钟才能阅读完成。

image.webp

为什么需要卷积自编码器?

传统全连接自编码器(AE)处理图像时有两个致命伤:

  • 参数爆炸 :假设输入是 28×28 的 MNIST 图像,仅第一层全连接就需要 784×n 个参数(n 为隐藏层维度),当图像尺寸增加到 256×256 时,参数量会呈平方级增长
  • 空间信息丢失 :全连接层将二维图像展平为一维向量,破坏了像素间的空间关联性

而卷积自编码器(CAE)通过局部感受野和权值共享特性,完美解决了这两个问题。

技术选型对比

特性 普通 AE VAE CAE
参数量 极高 较高 极低
空间信息保留 部分 优秀
训练速度 中等
生成质量 一般 优秀 中等
适用场景 低维数据 生成任务 特征提取

PyTorch 实现详解

编码器架构

编码器通过卷积和下采样逐步压缩空间维度:

class Encoder(nn.Module):
    def __init__(self):
        super().__init__()
        # 输入形状: [B, 1, 28, 28]
        self.conv1 = nn.Conv2d(1, 16, 3, padding=1)  # [B,16,28,28]
        self.pool1 = nn.MaxPool2d(2)  # [B,16,14,14]
        self.conv2 = nn.Conv2d(16, 32, 3, padding=1) # [B,32,14,14]
        self.pool2 = nn.MaxPool2d(2)  # [B,32,7,7]

    def forward(self, x):
        x = F.relu(self.conv1(x))
        x = self.pool1(x)
        x = F.relu(self.conv2(x))
        return self.pool2(x)

解码器设计关键

解码器通过转置卷积恢复空间维度,需特别注意:

  1. 最后一层不使用 ReLU,避免像素值被截断到 [0,∞)
  2. 当使用 BCELoss 时,需配合 Sigmoid 激活将输出约束到 [0,1]
class Decoder(nn.Module):
    def __init__(self):
        super().__init__()
        self.upconv1 = nn.ConvTranspose2d(32, 16, 2, stride=2) # [B,16,14,14]
        self.conv1 = nn.Conv2d(16, 16, 3, padding=1)
        self.upconv2 = nn.ConvTranspose2d(16, 1, 2, stride=2)  # [B,1,28,28]

    def forward(self, x):
        x = F.relu(self.upconv1(x))
        x = F.relu(self.conv1(x))
        return torch.sigmoid(self.upconv2(x))  # 重要!

完整 CAE 整合

class CAE(nn.Module):
    def __init__(self):
        super().__init__()
        self.encoder = Encoder()
        self.decoder = Decoder()

    def forward(self, x):
        latent = self.encoder(x)
        return self.decoder(latent)

训练技巧与避坑指南

数据预处理

transform = transforms.Compose([transforms.ToTensor(),  # 自动归一化到 [0,1]
    # 若使用 MNIST 数据集,无需额外归一化
])

损失函数选择

  • MSE 损失

    criterion = nn.MSELoss()

  • BCE 损失 (需配合 Sigmoid):

    criterion = nn.BCELoss()

特征可视化技巧

可视化中间层特征图:

# 获取第一层卷积后的特征图
with torch.no_grad():
    features = model.encoder.conv1(images)
    # 显示前 16 个通道
    plt.figure(figsize=(10,5))
    for i in range(16):
        plt.subplot(4,4,i+1)
        plt.imshow(features[0,i].cpu(), cmap='gray')

MNIST 实战效果

训练 20 个 epoch 后的重建效果对比:

Epoch 1/20 | Loss: 0.2187
Epoch 10/20 | Loss: 0.0342 
Epoch 20/20 | Loss: 0.0215

CAE 卷积自编码器模型入门指南:从理论到 PyTorch 实战

延伸思考

当处理彩色图像(如 CIFAR-10)时,我们需要:

  1. 修改输入通道数:nn.Conv2d(3, 16, 3)
  2. 调整解码器最后一层:nn.ConvTranspose2d(16, 3, 2, stride=2)
  3. 考虑使用感知损失(Perceptual Loss)提升重建质量

思考题 :如果希望 CAE 同时完成分类和重建任务,网络结构应该如何设计?

正文完
 0
评论(没有评论)