用PyTorch复现InfoGAN:手把手教你控制生成图像的‘数字’和‘笔画粗细’

在生成对抗网络(GAN)的世界里,InfoGAN一直是个独特的存在。它不仅能够生成逼真的图像,还能让我们像调音台一样精确控制生成结果的特定属性。想象一下,你正在训练一个生成手写数字的模型,突然发现可以通过几个简单的参数控制生成数字的类别和笔画粗细——这不是魔法,而是InfoGAN带给我们的现实能力。

对于已经掌握GAN基础但渴望更深入实践的开发者来说,复现InfoGAN就像获得了一把打开生成模型黑箱的钥匙。本文将带你从零开始,用PyTorch构建一个完整的InfoGAN模型,重点解决两个实际问题:如何让模型理解"数字类别"和"笔画粗细"这两个语义概念,以及如何在训练过程中稳定这些控制信号。

1. 环境准备与数据加载

在开始构建模型前,我们需要确保环境配置正确。推荐使用Python 3.8+和PyTorch 1.10+版本,这些版本在自动微分和GPU加速方面都有良好支持。如果你使用CUDA加速,别忘了检查torch.cuda.is_available()的输出。

import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader

# 检查GPU可用性
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")

MNIST数据集是我们的实验对象,这个包含手写数字的经典数据集非常适合验证InfoGAN的控制能力。但要注意,标准的MNIST加载方式需要做些调整,以适应InfoGAN的特殊需求:

# 数据预处理
transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5,), (0.5,))  # 将像素值从[0,1]归一化到[-1,1]
])

# 加载数据集
train_dataset = datasets.MNIST(
    root='./data', 
    train=True,
    download=True,
    transform=transform
)

train_loader = DataLoader(
    dataset=train_dataset,
    batch_size=64,
    shuffle=True,
    num_workers=4
)

这里有几个关键点值得注意:

  • 归一化到[-1,1]范围是为了匹配生成器tanh激活函数的输出范围
  • batch_size设置为64是个不错的起点,太大可能导致训练不稳定
  • num_workers可以加速数据加载,但要根据你的CPU核心数合理设置

2. 模型架构设计

InfoGAN的核心创新在于它的三网络结构:生成器(G)、判别器(D)和辅助网络(Q)。与普通GAN不同,我们需要特别设计Q网络来预测隐变量c,这是实现控制的关键。

2.1 生成器网络

生成器需要接收两种输入:随机噪声z和可解释隐变量c。在MNIST案例中,我们可以这样设计c:

  • 一个10维的one-hot向量表示数字类别(0-9)
  • 一个1维的连续变量控制笔画粗细
class Generator(nn.Module):
    def __init__(self, latent_dim=64, num_classes=10):
        super(Generator, self).__init__()
        # 隐变量c包含:10维类别 + 1维笔画粗细 = 11维
        self.total_latent_dim = latent_dim + num_classes + 1
        
        self.main = nn.Sequential(
            nn.Linear(self.total_latent_dim, 256),
            nn.BatchNorm1d(256),
            nn.ReLU(),
            
            nn.Linear(256, 512),
            nn.BatchNorm1d(512),
            nn.ReLU(),
            
            nn.Linear(512, 1024),
            nn.BatchNorm1d(1024),
            nn.ReLU(),
            
            nn.Linear(1024, 784),
            nn.Tanh()  # 输出范围[-1,1]
        )
    
    def forward(self, z, c_discrete, c_continuous):
        # 拼接噪声z和隐变量c
        input = torch.cat([z, c_discrete, c_continuous], dim=1)
        img = self.main(input)
        return img.view(-1, 1, 28, 28)  # 重塑为图像尺寸

提示:BatchNorm层对生成器至关重要,它能稳定训练过程。但要注意在测试时使用generator.eval()来固定统计量。

2.2 判别器与Q网络

判别器不仅要判断图像真伪,还要与Q网络共享特征提取层。这种设计既能节省计算资源,又能确保两个任务共享视觉特征。

class Discriminator(nn.Module):
    def __init__(self):
        super(Discriminator, self).__init__()
        
        # 共享的特征提取层
        self.feature_extractor = nn.Sequential(
            nn.Conv2d(1, 64, 4, 2, 1),
            nn.LeakyReLU(0.2),
            
            nn.Conv2d(64, 128, 4, 2, 1),
            nn.BatchNorm2d(128),
            nn.LeakyReLU(0.2),
            
            nn.Conv2d(128, 256, 4, 2, 1),
            nn.BatchNorm2d(256),
            nn.LeakyReLU(0.2)
        )
        
        # 判别真伪的头部
        self.discriminator_head = nn.Sequential(
            nn.Linear(256*7*7, 1),
            nn.Sigmoid()
        )
        
        # Q网络头部(预测隐变量c)
        self.Q_head = nn.Sequential(
            nn.Linear(256*7*7, 128),
            nn.BatchNorm1d(128),
            nn.LeakyReLU(0.2),
            
            nn.Linear(128, 10 + 1)  # 10类分类 + 1维回归
        )
    
    def forward(self, x):
        features = self.feature_extractor(x)
        features = features.view(features.size(0), -1)
        
        validity = self.discriminator_head(features)
        
        # Q网络输出:类别logits和连续值
        q_logits = self.Q_head(features)
        q_class = q_logits[:, :10]  # 前10维是类别
        q_cont = q_logits[:, 10:11]  # 最后一维是笔画粗细
        
        return validity, q_class, q_cont

这种共享设计有几个优势:

  1. 特征提取只需计算一次,效率更高
  2. 判别任务和Q任务共享底层视觉特征,相互促进
  3. 减少了模型参数量,降低过拟合风险

3. 训练策略与损失函数

InfoGAN的训练比标准GAN更复杂,因为它需要同时优化三个目标:生成质量、判别准确性和隐变量预测准确性。我们需要精心设计损失函数和训练流程。

3.1 损失函数组成

InfoGAN的损失由三部分组成:

  1. 对抗损失:与标准GAN相同,让生成样本尽可能真实
  2. 分类损失:预测离散隐变量(数字类别)的交叉熵
  3. 回归损失:预测连续隐变量(笔画粗细)的均方误差
# 初始化模型
generator = Generator().to(device)
discriminator = Discriminator().to(device)

# 定义优化器
g_optimizer = optim.Adam(generator.parameters(), lr=0.0002, betas=(0.5, 0.999))
d_optimizer = optim.Adam(discriminator.parameters(), lr=0.0001, betas=(0.5, 0.999))

# 损失函数
adversarial_loss = nn.BCELoss()
classification_loss = nn.CrossEntropyLoss()
continuous_loss = nn.MSELoss()

# 互信息权重
lambda_discrete = 1.0  # 类别控制的权重
lambda_continuous = 0.1  # 笔画粗细控制的权重

注意:连续变量的损失权重通常设置得比离散变量小,因为回归任务通常比分类更难优化。

3.2 训练循环实现

训练过程需要交替更新生成器和判别器/Q网络。下面是训练循环的关键代码:

for epoch in range(epochs):
    for i, (real_imgs, _) in enumerate(train_loader):
        real_imgs = real_imgs.to(device)
        batch_size = real_imgs.size(0)
        
        # 准备真实和假标签
        real_labels = torch.ones(batch_size, 1).to(device)
        fake_labels = torch.zeros(batch_size, 1).to(device)
        
        # 准备隐变量c
        # 离散变量:数字类别
        c_discrete = torch.zeros(batch_size, 10).to(device)
        random_classes = torch.randint(0, 10, (batch_size,))
        c_discrete.scatter_(1, random_classes.unsqueeze(1), 1)
        
        # 连续变量:笔画粗细 (-1到1之间)
        c_continuous = torch.rand(batch_size, 1).to(device) * 2 - 1
        
        # 随机噪声
        z = torch.randn(batch_size, 64).to(device)
        
        # ===== 训练生成器 =====
        g_optimizer.zero_grad()
        
        # 生成假图像
        fake_imgs = generator(z, c_discrete, c_continuous)
        
        # 判别器对假图像的判断
        validity, q_class, q_cont = discriminator(fake_imgs)
        
        # 计算损失
        g_loss_adv = adversarial_loss(validity, real_labels)
        g_loss_discrete = classification_loss(q_class, random_classes)
        g_loss_continuous = continuous_loss(q_cont, c_continuous)
        
        total_g_loss = g_loss_adv + \
                     lambda_discrete * g_loss_discrete + \
                     lambda_continuous * g_loss_continuous
        
        total_g_loss.backward()
        g_optimizer.step()
        
        # ===== 训练判别器 =====
        d_optimizer.zero_grad()
        
        # 真实图像损失
        real_validity, _, _ = discriminator(real_imgs)
        d_real_loss = adversarial_loss(real_validity, real_labels)
        
        # 假图像损失
        fake_validity, _, _ = discriminator(fake_imgs.detach())
        d_fake_loss = adversarial_loss(fake_validity, fake_labels)
        
        d_loss = (d_real_loss + d_fake_loss) / 2
        
        d_loss.backward()
        d_optimizer.step()

训练过程中有几个容易出错的点:

  • 忘记对fake_imgs使用detach()会导致判别器更新影响生成器
  • 连续变量的范围要与生成器的输出范围匹配
  • 损失权重需要根据训练动态调整

4. 控制生成与结果分析

训练完成后,我们可以通过调整隐变量c来控制系统生成特定属性的数字。以下是控制生成的关键代码:

# 固定噪声z,只改变隐变量c
z = torch.randn(1, 64).to(device).repeat(10, 1)  # 对10个数字使用相同噪声

# 控制数字类别
c_discrete = torch.zeros(10, 10).to(device)
for i in range(10):
    c_discrete[i, i] = 1  # 每个数字类别one-hot编码

# 控制笔画粗细 (-1:细, 1:粗)
c_continuous = torch.linspace(-1, 1, 10).view(-1, 1).to(device)

# 生成图像
with torch.no_grad():
    control_imgs = generator(z, c_discrete, c_continuous)

通过这种方法,我们可以观察到:

  1. 改变c_discrete会改变生成数字的类别
  2. 改变c_continuous会改变笔画的粗细程度
  3. 相同的噪声z保证了其他视觉特征的一致性

实际训练中可能会遇到的一些问题及解决方案:

问题现象 可能原因 解决方案
生成图像模糊 判别器太强 降低判别器学习率
控制不准确 互信息权重不足 增大lambda_discrete/lambda_continuous
模式崩溃 生成器太强 增加判别器更新频率
训练不稳定 学习率太高 使用更小的学习率

在成功训练后,你可以尝试以下进阶实验:

  • 添加更多的控制变量(如倾斜角度、数字大小)
  • 尝试不同的网络架构(如ResNet块)
  • 应用到其他数据集(如FashionMNIST)
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐