DCGAN 实战:PyTorch 1.13 实现 64x64 人脸生成,FID 分数低于 30

人脸生成一直是计算机视觉领域最具挑战性的任务之一。传统方法往往难以捕捉人脸细节的复杂分布,而深度卷积生成对抗网络(DCGAN)通过对抗训练机制,能够生成高度逼真的人脸图像。本文将带你从零实现一个完整的DCGAN项目,使用PyTorch 1.13生成64x64分辨率的人脸图像,并确保FID分数低于30。

1. 环境准备与数据加载

1.1 安装依赖

首先确保你的环境已安装PyTorch 1.13及以上版本。推荐使用Anaconda创建虚拟环境:

conda create -n dcgan python=3.8
conda activate dcgan
conda install pytorch==1.13.1 torchvision==0.14.1 -c pytorch
pip install numpy pandas matplotlib scikit-learn

1.2 数据集选择与预处理

我们使用CelebA数据集,包含超过20万张名人面部图像。下载后解压到 data/celeba 目录,然后进行预处理:

import torchvision.transforms as transforms
from torchvision.datasets import ImageFolder

transform = transforms.Compose([
    transforms.Resize(64),
    transforms.CenterCrop(64),
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])

dataset = ImageFolder(root='data/celeba', transform=transform)
dataloader = torch.utils.data.DataLoader(dataset, batch_size=128, 
                                         shuffle=True, num_workers=4)

关键预处理步骤:

  • 统一缩放到64x64分辨率
  • 中心裁剪保证人脸居中
  • 像素值归一化到[-1, 1]范围
  • 使用批量大小为128加速训练

2. DCGAN模型架构设计

DCGAN的核心创新在于将CNN引入GAN框架。以下是生成器(G)和判别器(D)的详细实现:

2.1 生成器网络

生成器接收100维随机噪声,通过转置卷积逐步上采样:

import torch.nn as nn

class Generator(nn.Module):
    def __init__(self, ngpu=1):
        super(Generator, self).__init__()
        self.ngpu = ngpu
        self.main = nn.Sequential(
            # 输入是Z, 进入全连接
            nn.ConvTranspose2d(100, 512, 4, 1, 0, bias=False),
            nn.BatchNorm2d(512),
            nn.ReLU(True),
            
            # 状态大小 (512,4,4)
            nn.ConvTranspose2d(512, 256, 4, 2, 1, bias=False),
            nn.BatchNorm2d(256),
            nn.ReLU(True),
            
            # 状态大小 (256,8,8)
            nn.ConvTranspose2d(256, 128, 4, 2, 1, bias=False),
            nn.BatchNorm2d(128),
            nn.ReLU(True),
            
            # 状态大小 (128,16,16)
            nn.ConvTranspose2d(128, 64, 4, 2, 1, bias=False),
            nn.BatchNorm2d(64),
            nn.ReLU(True),
            
            # 状态大小 (64,32,32)
            nn.ConvTranspose2d(64, 3, 4, 2, 1, bias=False),
            nn.Tanh()
            # 输出状态大小 (3,64,64)
        )

    def forward(self, input):
        return self.main(input)

关键设计要点:

  • 使用转置卷积实现上采样
  • 每层后接BatchNorm稳定训练
  • 输出层使用Tanh将像素值约束到[-1,1]
  • 避免在全连接层使用池化操作

2.2 判别器网络

判别器是标准的CNN分类器,但使用LeakyReLU激活:

class Discriminator(nn.Module):
    def __init__(self, ngpu=1):
        super(Discriminator, self).__init__()
        self.ngpu = ngpu
        self.main = nn.Sequential(
            # 输入 (3,64,64)
            nn.Conv2d(3, 64, 4, 2, 1, bias=False),
            nn.LeakyReLU(0.2, inplace=True),
            
            # 状态大小 (64,32,32)
            nn.Conv2d(64, 128, 4, 2, 1, bias=False),
            nn.BatchNorm2d(128),
            nn.LeakyReLU(0.2, inplace=True),
            
            # 状态大小 (128,16,16)
            nn.Conv2d(128, 256, 4, 2, 1, bias=False),
            nn.BatchNorm2d(256),
            nn.LeakyReLU(0.2, inplace=True),
            
            # 状态大小 (256,8,8)
            nn.Conv2d(256, 512, 4, 2, 1, bias=False),
            nn.BatchNorm2d(512),
            nn.LeakyReLU(0.2, inplace=True),
            
            # 状态大小 (512,4,4)
            nn.Conv2d(512, 1, 4, 1, 0, bias=False),
            nn.Sigmoid()
        )

    def forward(self, input):
        return self.main(input).view(-1)

关键设计要点:

  • 使用带步长的卷积代替池化
  • LeakyReLU斜率设为0.2防止梯度消失
  • 最后一层Sigmoid输出概率值
  • 去除全连接层,直接使用卷积输出

3. 训练策略与技巧

3.1 损失函数与优化器

使用二元交叉熵损失和Adam优化器:

# 初始化网络
netG = Generator().to(device)
netD = Discriminator().to(device)

# 定义优化器
optimizerD = optim.Adam(netD.parameters(), lr=0.0002, betas=(0.5, 0.999))
optimizerG = optim.Adam(netG.parameters(), lr=0.0002, betas=(0.5, 0.999))

# 损失函数
criterion = nn.BCELoss()

# 固定噪声用于可视化
fixed_noise = torch.randn(64, 100, 1, 1, device=device)

3.2 训练循环实现

采用交替训练策略,先更新判别器再更新生成器:

for epoch in range(num_epochs):
    for i, (real_images, _) in enumerate(dataloader):
        # 训练判别器
        netD.zero_grad()
        real_images = real_images.to(device)
        batch_size = real_images.size(0)
        label = torch.full((batch_size,), real_label, device=device)
        
        output = netD(real_images)
        errD_real = criterion(output, label)
        errD_real.backward()
        
        noise = torch.randn(batch_size, 100, 1, 1, device=device)
        fake = netG(noise)
        label.fill_(fake_label)
        output = netD(fake.detach())
        errD_fake = criterion(output, label)
        errD_fake.backward()
        errD = errD_real + errD_fake
        optimizerD.step()

        # 训练生成器
        netG.zero_grad()
        label.fill_(real_label)
        output = netD(fake)
        errG = criterion(output, label)
        errG.backward()
        optimizerG.step()

关键训练技巧:

  • 判别器训练时冻结生成器参数
  • 生成器训练时冻结判别器参数
  • 使用不同的标签值(real_label=1, fake_label=0)
  • 每100批次可视化生成结果

3.3 FID分数监控

FID(Frechet Inception Distance)是评估生成质量的重要指标:

from torchvision.models import inception_v3
from scipy.linalg import sqrtm

def calculate_fid(real_images, fake_images):
    # 加载预训练Inception-v3模型
    inception_model = inception_v3(pretrained=True, transform_input=False).to(device)
    inception_model.eval()
    
    # 提取特征
    with torch.no_grad():
        real_features = inception_model(real_images)[0].view(real_images.size(0), -1)
        fake_features = inception_model(fake_images)[0].view(fake_images.size(0), -1)
    
    # 计算统计量
    mu1, sigma1 = real_features.mean(0), torch_cov(real_features, rowvar=False)
    mu2, sigma2 = fake_features.mean(0), torch_cov(fake_features, rowvar=False)
    
    # 计算FID
    diff = mu1 - mu2
    covmean = sqrtm(sigma1 @ sigma2)
    fid = diff.dot(diff) + torch.trace(sigma1 + sigma2 - 2*covmean)
    return fid

目标是将FID控制在30以下,这需要:

  • 充分训练(通常需要50-100个epoch)
  • 使用足够大的批量(至少64)
  • 定期检查FID避免过拟合

4. 模型调优与问题解决

4.1 常见问题分析

问题现象 可能原因 解决方案
生成图像模糊 判别器过强 降低判别器学习率
模式崩溃 生成器多样性不足 增加噪声维度
训练不稳定 学习率过高 使用更小的学习率(如0.0001)
生成图像有 artifacts 网络容量不足 增加滤波器数量

4.2 高级调优技巧

标签平滑 :防止判别器过度自信

real_labels = torch.FloatTensor(batch_size).uniform_(0.9, 1.0)
fake_labels = torch.FloatTensor(batch_size).uniform_(0.0, 0.1)

梯度惩罚 :改进Wasserstein GAN

def gradient_penalty(D, real_samples, fake_samples):
    alpha = torch.rand(real_samples.size(0), 1, 1, 1).to(device)
    interpolates = (alpha * real_samples + (1 - alpha) * fake_samples).requires_grad_(True)
    d_interpolates = D(interpolates)
    gradients = torch.autograd.grad(
        outputs=d_interpolates,
        inputs=interpolates,
        grad_outputs=torch.ones_like(d_interpolates),
        create_graph=True,
        retain_graph=True,
        only_inputs=True,
    )[0]
    gradients = gradients.view(gradients.size(0), -1)
    penalty = ((gradients.norm(2, dim=1) - 1) ** 2).mean()
    return penalty

学习率调度 :后期微调

schedulerD = torch.optim.lr_scheduler.StepLR(optimizerD, step_size=30, gamma=0.1)
schedulerG = torch.optim.lr_scheduler.StepLR(optimizerG, step_size=30, gamma=0.1)

5. 结果分析与应用

经过充分训练后,我们的DCGAN能够生成高质量的人脸图像。以下是典型生成样本:

生成样本示例

实际应用中可以考虑以下方向:

  • 数据增强:为分类任务生成更多训练样本
  • 图像编辑:通过潜在空间插值实现属性编辑
  • 隐私保护:生成匿名化人脸替代真实照片
  • 艺术创作:生成不存在的人物肖像

完整项目代码已开源在GitHub,包含预训练模型和详细使用说明。在实际部署时,建议使用ONNX格式导出模型以提高推理效率:

torch.onnx.export(netG,               # 模型
                  torch.randn(1,100,1,1),  # 输入样本
                  "dcgan.onnx",       # 保存路径
                  export_params=True,  # 导出参数
                  opset_version=11,    # ONNX版本
                  do_constant_folding=True,  # 优化
                  input_names=['input'],   # 输入名
                  output_names=['output']) # 输出名
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐