ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

从零实现DCGAN:用PyTorch生成动漫头像的实战指南

从零实现DCGAN:用PyTorch生成动漫头像的实战指南 第一次接触GAN的人大多都是从MNIST入手的但说实话28x28的灰度数字生成器很难让你体会到生成模型的真正乐趣。这次项目我把目标换成了二次元动漫头像生成用PyTorch从零实现一个DCGAN训练一个能输出64x64动漫脸的生成模型。为什么选这个方向因为数据好找、训练量级对普通显卡足够友好、生成效果一眼就能看出好坏不需要像FID那样上指标也能判断模型到底有没有学到东西。这一篇会完整串起数据预处理、网络结构、训练循环、踩坑排错和优化方向。代码基于PyTorch做DCGAN实现如果你之前跑通过MNIST版但一换数据集就翻车或者想找一个真正能玩起来的GAN实战项目这篇应该能帮你省下不少时间。1. 数据源选择与预处理管线Anime Face Dataset的取舍1.1 为什么用现成数据集而不是自己爬图动漫头像生成最关键的前置工作不是搭模型而是搞定数据。我见过不少朋友第一步就想着写爬虫去图站抓图再用OpenCV的人脸检测器裁剪这套流程听起来很工程化实际跑起来全是坑目标站点会封IP、图片授权和合规问题让人头疼、OpenCV的人脸检测模型对动漫风格误检率极高检测框经常框住脖子或者半个脸清洗成本比训练成本还高。所以个人项目我非常推荐直接用公开数据集。最常用的是Kaggle上的Anime Face Dataset大约有6.3万张已经裁剪好的动漫脸部图尺寸统一是96x96数据来自社区图片的筛选裁剪具体来源细节以数据集页面的说明为准。选它的核心理由是省掉了整个数据采集链路拿到手只需要做尺寸统一和归一化就能进模型这对跑通一个完整项目来说体验完全不同。如果你的目标是更高分辨率的生成效果也可以考虑从中筛选子集或者用更大规模的Danbooru系列数据集。但对DCGAN这个级别的模型来说6.3万张已经能支撑出不错的结果数据再多模型容量和训练时长也会成为瓶颈。1.2 从96x96到64x64Resize、水平翻转和归一化数据集的原始尺寸是96x96但DCGAN最经典、最稳的训练尺寸是64x64。这不是说96x96不行而是盲目放大分辨率会带来两个问题第一模型通道数和训练时长成倍上涨第二GAN在低分辨率下更容易收敛作为基线项目先把64x64跑通、跑稳比一上来追求高清更有价值。预处理管线我建议按下面这套来from torch.utils.data import DataLoader from torchvision import datasets from torchvision.transforms import v2 as T transform T.Compose([ T.Resize((64, 64)), T.RandomHorizontalFlip(p0.5), T.ToTensor(), T.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]), ]) dataset datasets.ImageFolder(rootdata/anime_faces, transformtransform) dataloader DataLoader( dataset, batch_size128, shuffleTrue, num_workers4, pin_memoryTrue, )这里有两个细节很容易被忽略。第一Resize直接用双线性插值把一个96x96的图缩到64x64会损失部分细节但换来的是训练速度和稳定性对基线项目是划算的。第二归一化必须用[0.5, 0.5, 0.5]这个均值方差这会把像素范围映射到[-1, 1]因为生成器最后一层是Tanh输出范围就是[-1, 1]。如果你图省事用了ImageNet的均值和方差训练初期判别器会非常困惑损失震荡得厉害。数据增强方面我只加了水平翻转。动漫头像左右翻转不会破坏语义而且能等效地把训练数据翻倍。像随机裁剪、色彩抖动这类增强在GAN的训练里要谨慎过度增强会让判别器学习到错误的真实感信号反而拖慢收敛。1.3 DataLoader参数和训练集划分DataLoader参数里建议num_workers4到8batch_size128。在Windows上num_workers设成0可以避免多进程的序列化报错Linux上可以放心开4以上。pin_memoryTrue配合CUDA训练能减少数据从CPU到GPU的拷贝时间这个参数在数据加载占比较大的场景里有实际收益。这个数据集没有划分验证集和测试集GAN也不需要传统意义上的验证集因为训练目标不是预测精度而是学到训练数据的分布。如果要看生成质量和多样性只需要固定一个随机噪声每个epoch结束后用这个噪声生成一批图对比不同epoch之间的变化这比任何metric都直观。2. 生成器和判别器的PyTorch实现DCGAN设计原则逐层拆解2.1 生成器从100维噪声一步步反卷积出64x64图像DCGAN的核心设计理念是拿卷积网络替代原始GAN里的全连接网络。生成器的输入是100维高斯噪声向量经过一个线性映射reshape成1024x4x4的特征图之后逐步用转置卷积把空间尺寸从4x4放大到8x8、16x16、32x32最后到64x64通道数则从1024逐层降到3。我自己实现的生成器结构如下import torch import torch.nn as nn class Generator(nn.Module): def __init__(self, latent_dim100): super().__init__() self.main nn.Sequential( nn.ConvTranspose2d(latent_dim, 1024, 4, 1, 0, biasFalse), nn.BatchNorm2d(1024), nn.ReLU(True), nn.ConvTranspose2d(1024, 512, 4, 2, 1, biasFalse), nn.BatchNorm2d(512), nn.ReLU(True), nn.ConvTranspose2d(512, 256, 4, 2, 1, biasFalse), nn.BatchNorm2d(256), nn.ReLU(True), nn.ConvTranspose2d(256, 128, 4, 2, 1, biasFalse), nn.BatchNorm2d(128), nn.ReLU(True), nn.ConvTranspose2d(128, 3, 4, 2, 1, biasFalse), nn.Tanh(), ) def forward(self, z): return self.main(z.unsqueeze(2).unsqueeze(3))有几个设计细节直接来自DCGAN论文我当时的理解是生成器里所有层都用ReLU激活除了输出层用Tanh所有卷积层都不加bias因为在后面接BatchNorm时BN层自带可学习的偏置参数多余的bias只会浪费参数转置卷积的stride2、kernel_size4、padding1这个组合能让特征图尺寸精确翻倍这个结论自己推一遍输出尺寸公式就能记住比死记硬背可靠。如果遇到z.unsqueeze(2).unsqueeze(3)不理解的读者这里是把(batch, 100)的噪声变成(batch, 100, 1, 1)的四维张量让转置卷积能直接处理它。第一次写DCGAN的人很容易卡在这一行。2.2 判别器用stride卷积代替池化判别器的结构和生成器几乎是对称的但方向反过来。输入是3通道的64x64图像经过4个步长为2的普通卷积逐步把尺寸压缩到4x4通道数从3升到512最后用一个不padding的4x4卷积把所有信息压成一个标量再过Sigmoid输出这是真实图片的概率。这里最关键的取舍是DCGAN原始论文明确说了判别器里不要用池化层而是用带stride的卷积做降采样。原因是池化是固定规则的操作参数不可学习卷积降采样让判别器自己决定要保留什么信息判别能力更强梯度回传也更顺畅。class Discriminator(nn.Module): def __init__(self): super().__init__() self.main nn.Sequential( nn.Conv2d(3, 64, 4, 2, 1, biasFalse), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(64, 128, 4, 2, 1, biasFalse), nn.BatchNorm2d(128), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(128, 256, 4, 2, 1, biasFalse), nn.BatchNorm2d(256), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(256, 512, 4, 2, 1, biasFalse), nn.BatchNorm2d(512), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(512, 1, 4, 1, 0, biasFalse), nn.Sigmoid(), ) def forward(self, x): return self.main(x)判别器里激活函数换成了LeakyReLU负斜率0.2。这个选择的原因是ReLU在反向传播时负半轴梯度恒为0判别器一旦输出落入负区域对应位置的梯度就完全消失这会加剧训练不稳定LeakyReLU保留了负半轴的梯度虽然只是一个小斜率但足以维持梯度流动。判别器的BN位置也很有讲究。第一层输入层不加BatchNorm最后一层输出层不加BatchNorm只在中国三层加。原因是输入层做BN会把原始图像分布打乱输出层加BN会影响Sigmoid输出对真实概率的表达论文里这个设计是为了稳定训练。2.3 参数初始化为什么是均值为0、标准差为0.02的正态分布模型定义好之后必须做参数初始化这是DCGAN训练能收敛的隐藏条件。DCGAN论文里对所有权重和偏置的初始化策略是N(0, 0.02)。如果用PyTorch的默认初始化直接开始训练训练前期会非常痛苦梯度爆炸和生成器崩溃的概率大大增加。def weights_init(m): if isinstance(m, (nn.Conv2d, nn.ConvTranspose2d)): nn.init.normal_(m.weight.data, 0.0, 0.02) elif isinstance(m, nn.BatchNorm2d): nn.init.normal_(m.weight.data, 1.0, 0.02) nn.init.constant_(m.bias.data, 0) G Generator().to(device) D Discriminator().to(device) G.apply(weights_init) D.apply(weights_init)初始化影响这么大原因在于GAN的训练本质上是两个网络在对抗中共同进化任何一方的初始状态离谱都会让另一方快速学会躺赢。BN层里weight用均值1、标准差0.02bias设为0这样初始阶段BN层的输出基本是标准化的不会一下子把激活值推到饱和区。3. 训练循环与关键超参数为什么Adam的beta1要设成0.53.1 BCE损失与交替训练节奏DCGAN的损失函数用二分类交叉熵PyTorch里直接调nn.BCELoss()即可。训练节奏是交替进行的先更新判别器再更新生成器每个batch执行一次完整更新。判别器的目标很直接真实图片输出接近1生成图片输出接近0。生成器的目标则是骗过判别器让判别器对生成图片输出尽量接近1。注意生成器训练时用的标签是1.0而不是0.0很多第一次上手的人会在这里写错。criterion nn.BCELoss() real_label 1.0 fake_label 0.0 opt_G torch.optim.Adam(G.parameters(), lr2e-4, betas(0.5, 0.999)) opt_D torch.optim.Adam(D.parameters(), lr2e-4, betas(0.5, 0.999)) for epoch in range(epochs): for i, (images, _) in enumerate(dataloader): # ---------- 训练判别器 ---------- D.zero_grad() real images.to(device) batch_size real.size(0) output D(real).view(-1) errD_real criterion(output, torch.full_like(output, real_label)) z torch.randn(batch_size, latent_dim, 1, 1, devicedevice) fake G(z) output D(fake.detach()).view(-1) errD_fake criterion(output, torch.full_like(output, fake_label)) errD errD_real errD_fake errD.backward() opt_D.step() # ---------- 训练生成器 ---------- G.zero_grad() output D(fake).view(-1) errG criterion(output, torch.full_like(output, real_label)) errG.backward() opt_G.step()fake.detach()这一行是很多人会忽略的关键。更新判别器时生成器输出的fake图只是为了给判别器喂数据梯度不应该回传到生成器detach之后反向传播只更新判别器参数不会白白把梯度算进生成器。而到了生成器更新阶段用的又是新的z或者复用之前的fake此时不detach让梯度穿过判别器回传到生成器。3.2 超参数配置学习率2e-4和beta10.5DCGAN的超参数配置非常反直觉Adam默认的beta10.9在这里不能用必须设成0.5学习率两个网络都固定为0.0002。beta1是Adam的一阶动量衰减系数控制着梯度历史平均多大程度上影响当前更新。默认的0.9会让梯度更新带了很大的惯性这在普通分类任务里是好事但GAN的损失面高度非凸且动态变化惯性太大会导致两个玩家在对抗中剧烈震荡甚至发散。设成0.5相当于大幅降低动量让每次更新更看重当前梯度这是GAN训练里公认的稳定手段。学习率我这里G和D都取2e-4。如果你觉得训练不稳第一个尝试的不是各自乱调而是把两个网络的学习率同步调低到1e-4。GAN最忌讳的就是G和D的学习率不平衡一个学太快一个学太慢最后一定有一方躺平。3.3 PyTorch安装与运行环境代码本身不需要特别新奇的库PyTorch基础框架就够用。如果你的显卡驱动是近两年的直接在PyTorch官网找到对应CUDA版本的安装命令比如pip3 install torch torchvision --index-url https://download.pytorch.org/whl/cu121这种形式装完用torch.cuda.is_available()验证一下就能开工。CPU训练这套模型也能跑但单epoch速度大概是GPU的20倍以上强烈不建议宁可在云端租一张T4。我的实际运行配置是PyTorch 2.x加CUDA 12.x6万张图的规模在T4上单epoch大约1到2分钟RTX 3060级别会更快训练50个epoch大约一个多小时。这个量级对调试来说非常友好等一次训练出结果不会等到怀疑人生。3.4 Checkpoint与固定噪声可视化训练中途必须定期保存模型而且要固定一个随机噪声向量z_fixed每个epoch结束后用同一个噪声生成图像。这两件事缺一不可。固定噪声的作用是提供一个控制变量同一个输入在不同epoch的生成结果能直接看出模型学习的方向是否正确。如果你每个epoch都用新噪声看到的结果每次都不一样根本分辨不出是模型进步了还是噪声本身就变了。z_fixed torch.randn(64, latent_dim, 1, 1, devicedevice) # 每个epoch结束时 G.eval() with torch.no_grad(): fake_fixed G(z_fixed).cpu() G.train() # 用torchvision.utils.save_image保存成8x8网格这里必须强调G.eval()的用意。PyTorch中BatchNorm在train()和eval()两种模式下的行为完全不同训练时用当前batch的均值方差推理时用累计的running统计量。如果保存图像时忘记切到eval()BN层用的还是训练模式生成结果会有额外的随机性你看到的进步曲线会掺杂噪声影响判断。保存模型建议用字典形式一次存全量状态torch.save({ epoch: epoch, G_state: G.state_dict(), D_state: D.state_dict(), optG_state: opt_G.state_dict(), optD_state: opt_D.state_dict(), }, fcheckpoints/dcgan_epoch_{epoch:03d}.pt)只保存state_dict而丢掉优化器状态会导致中断后无法严格恢复训练Adam的动量信息全丢了相当于重新开始。4. 训练过程中的两大翻车现场模式崩塌与判别器过强4.1 翻车现场一生成器把所有噪声都映射成同一张脸我第一次在这个数据集上训练跑到第10个epoch左右发现固定噪声生成的64张图从第6个epoch开始变得越来越像到第10个epoch几乎就是同一个人的不同表情。这就是典型的模式崩塌Mode Collapse生成器找到了一个能骗过判别器的安全点于是不管输入什么噪声都往那个点输出多样性完全丢失。判断模式崩塌不能只看损失值。我当时的G loss和D loss都维持在一个看似健康的水平G在1.5左右震荡D在0.4左右光看曲线完全看不出问题。真正暴露问题的是可视化8x8网格里所有图像的高光、发色、朝向几乎一致。这种视觉上的重复感是模式崩塌最直接的信号。排查思路和修复手段我按顺序试了三个。第一降低学习率到1e-4给生成器更小的更新步长避免它大步流星冲进安全点。第二把z的维度从100提到128增加输入噪声的表达自由度。第三也是最有效的一招给判别器加标签平滑真实图片标签从1.0改成0.9让判别器不要对真实样本过于自信从而留出梯度空间让生成器继续探索。real_label_smooth 0.9 errD_real criterion(output, torch.full_like(output, real_label_smooth))把标签平滑加进去之后再训练20个epoch固定噪声的生成结果恢复出了发型、发色、脸型的多样性。模式崩塌在训练早期出现很正常关键是要有手段把它拉回来。4.2 翻车现场二判别器loss直接归零生成器开始输出噪点另一次翻车发生在换用更大的batch size之后。判别器在30个epoch后loss掉到0.01以下而生成器的loss飙到8以上生成的图像变成了一片片色彩斑驳的噪点。这是典型的判别器过强它已经把真实和生成两个分布完全分开梯度消失生成器什么都学不到。这种状态下如果你继续训练判别器的loss会越来越低但生成器永远无法翻身因为判别器的梯度在接近饱和区Sigmoid两端时几乎为0生成器根本没有有效的学习信号。我的修复方案分成两步。第一步把判别器的通道数从64/128/256/512降为32/64/128/256降低判别器的模型容量让它别那么聪明。第二步在判别器的输入上加上标准差为0.05的高斯噪声。这个思路的直觉是给判别器的考题增加一点难度让它不能躺赢必须在噪声干扰下学会区分真实和生成样本生成器才有了追赶的空间。def add_noise(x, std0.05): if std 0: return x torch.randn_like(x) * std return x output_real D(add_noise(real)).view(-1)这里有个细节加噪声只加在判别器输入上生成器输出的fake图不用加因为生成器本来就自带一定的噪声早期输出质量差对判别器就是天然干扰。噪声标准差按epoch衰减训练后期可以降为0让判别器能精确判断真实样本。4.3 稳定训练的排查清单经历了这两次翻车之后我整理了一个排查清单每次训练异常就按这个顺序过一遍固定噪声图像是否出现大面积重复是优先查模式崩塌先加标签平滑再考虑动结构。G和D的loss是否长期不变且图像无改善大概率是学习率过低尝试翻倍到4e-4观察几个epoch。D的loss长期接近0判别器过强降低D容量或给D输入加噪。G的loss长期接近0生成器过强少见给G的输出加噪声或削减G容量。所有图像都是纯色或全黑检查归一化是否用0.5/0.5以及生成器输出层是否在Tanh之后。这套清单的核心思路是不要盲调参数先用可视化确认是哪一方占了上风再有针对性地干预。训练GAN就像调解两个人吵架你得先搞清楚谁在欺负谁再决定拉偏架的方向。5. 效果观察与后续优化方向从DCGAN到更香的架构5.1 如何判断训练状态而不是迷信loss曲线GAN的loss曲线是所有深度学习项目里最不直观的因为G和D的loss是零和博弈的两个绝对值它们的相对关系随时会变。我见过太多人在训练群里晒出漂亮的loss下降曲线结果生成的图完全不能看原因就是只盯着数值忽略了分布学习的目标。我的经验是三分看曲线七分看图。看曲线主要确认没有极端值D的loss没有长期贴着0G的loss没有爆炸式增长。看图则要注意三件事五官轮廓是否清晰欠拟合还是轮廓模糊发丝和眼睛的高光是否有细节模糊往往意味着模型容量不够固定噪声的不同样本之间是否有表情和角度的差异多样性是判断模式崩塌的窗口。如果跑完50个epoch输出在64x64的尺寸下五官准确、肤色自然、头发轮廓没有明显撕裂这个DCGAN项目就算成功了。清晰度上不要期待它能和Pixiv上的高清原图相比64x64本身就是一种风格很多独立游戏人物头像就是这个分辨率。5.2 低成本可落地的几个优化方向如果你不满足于当前效果想继续往上提升我建议按性价比从高到低依次尝试这几个方向。第一分辨率从64x64提升到128x128。做法是在生成器和判别器里各加一层卷积转置或卷积把通道设计改为1024/512/256/128/64和64/128/256/512/1024。显存占用大约会增加一倍T4的16G完全扛得住但要注意更大分辨率下BN的batch size最好保持在64以上否则统计量不稳定。第二把BCELoss换成hinge loss或者WGAN-GP。这些损失函数在设计上专门针对训练不稳定做了改良尤其是WGAN-GP配合梯度惩罚机制能显著减少模式崩塌的发生频率。代价是实现复杂度高一些要处理梯度惩罚项但对于已经跑通DCGAN的人来说代码量增加不多收益明显。第三给判别器加谱归一化Spectral Normalization。这个方法比WGAN-GP更容易集成只需要把判别器里的卷积层包一层nn.utils.spectral_norm能有效限制判别器的Lipschitz常数训练稳定性提升明显而且几乎不增加训练时间。第四如果对生成质量有更高追求直接考虑迁移到StyleGAN或StyleGAN2。DCGAN是理解GAN原理的最佳教材但StyleGAN系列在动漫人脸生成上的效果是断层式的领先。迁移成本也没有想象中高有很多开源实现可以直接用我建议至少完整跑通一次DCGAN再上手因为很多训练技巧固定噪声可视化、标签平滑、噪声注入在StyleGAN里依然适用。这次项目完整的复现路径就是这样。从数据预处理到模型结构从训练循环到翻车修复每一步都有可以沉淀的经验。GAN训练确实比普通监督学习更容易让人心态崩但只要你把观察——假设——干预——复现验证的循环建立起来它能带来的正反馈也是其他模型很难比的。动手跑一遍把固定噪声那张8x8网格图慢慢看到五官清晰起来的过程比任何教程都管用。
返回列表