DCGAN 实战:PyTorch 1.13 实现 64x64 人脸生成,FID 分数低于 30
·
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']) # 输出名
更多推荐




所有评论(0)