资讯动态

WGAN-GP动漫头像生成实战:256×256稳定训练的Python源码与踩坑指南

发布时间:2026/10/10 1:27:01 来源:尧图企业网站定制
简介这是一套基于WGAN-GP算法的动漫头像生成系统源码面向对生成对抗网络和图像生成感兴趣的开发者可用于学习并实际生成256×256像素的高清动漫人物头像。针对传统GAN训练不稳定、易出现模式崩塌的问题代码采用Wasserstein距离作为分布度量并引入梯度惩罚项来增强训练稳定性从而提升生成头像的清晰度和多样性。资源包共26个文件压缩包大小仅1.32MB包含2个Python核心算法文件、11张生成结果示例PNG、6个XML配置文件、项目说明文本以及Git忽略文件等XML配置可用于保存训练参数与路径设置Idea工程文件便于在集成开发环境中直接打开。目前已有306人学习下载。读者通过研读主程序、模型结构与生成样例可以理解WGAN-GP的损失函数、梯度惩罚实现和网络构建思路并可根据需要调整输入条件生成不同性别、发型、表情和配饰的动漫头像适用于游戏素材、虚拟形象、表情包制作等场景也为后续算法改进提供了可扩展的基线代码。1. 一张 256×256 的动漫脸为什么偏要选 WGAN-GP能稳定跑出 256×256 动漫头像的 WGAN-GP 源码远比一个只能出 32×32 缩略图的 Demo 值钱。很多人找免费 python 源码其实是想拿到一个能直接跑通的基线而不是重新推导一遍定理模型已经写好、数据管线已经备齐、训练参数给了一组能用的默认值开机就能看着损失曲线往下走。这套方案适合两类人一是训练过 GAN 但一直困在小分辨率、一放大就训练崩掉的人二是会写模型但没调过大图训练、想省掉几个月试错成本的人。下面按我自己的拆解习惯把 WGAN-GP 在 256×256 动漫头像生成这件事上的原理、数据、训练参数和踩坑点一次讲透。2. 先立住原理再上代码Wasserstein 距离与梯度惩罚2.1 从判别器到评论家输出层的 sigmoid 必须拆掉经典 GAN 的判别器末端是 sigmoid输出一个 0 到 1 的概率。这个设计在 64×64 以下的世界没什么大问题一旦分辨率抬到 256×256训练就变成玄学。根因是 JS 散度在真实分布和生成分布几乎没有重叠时梯度会趋近于 0——判别器越强生成器越学不到东西表现为损失卡死、图像永远是噪声。WGAN 把判别器改名叫评论家末端 sigmoid 拆掉输出一个无界标量。这个标量的含义不再是概率而是“这个样本有多接近真实分布”的相对分数。真实样本分数高、生成样本分数低两者的均值差近似 Wasserstein 距离这个距离即使两个分布完全不重叠也能给出有意义的梯度方向。下面是我在源码里写评论家输出层的惯用写法class Critic(nn.Module): def __init__(self, in_channels3): super().__init__() # 主干用六层卷积把 256x256 一路压到 4x4通道数逐层翻倍 self.body nn.Sequential( nn.Conv2d(in_channels, 32, 4, 2, 1), # 256 - 128 nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(32, 64, 4, 2, 1), # 128 - 64 nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(64, 128, 4, 2, 1), # 64 - 32 nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(128, 256, 4, 2, 1), # 32 - 16 nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(256, 512, 4, 2, 1), # 16 - 8 nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(512, 1024, 4, 2, 1), # 8 - 4 nn.LeakyReLU(0.2, inplaceTrue), ) # 末尾不接 sigmoid直接输出标量 self.fc nn.Linear(1024 * 4 * 4, 1) def forward(self, x): return self.fc(self.body(x).view(x.size(0), -1))这里最容易翻车的地方是照搬 DCGAN 的判别器代码忘了把最后一层 sigmoid 删掉。sigmoid 会把输出压到 0 到 1 之间W 距离的估计就被截断了生成器拿到的梯度信号会弱一大截。LeakyReLU 的斜率用 0.2 是 DCGAN 留下来的惯例判别的下采样路径上不要加 BatchNorm原因后面避坑章节细说。2.2 梯度惩罚把 Lipschitz 约束写进损失而不是写进裁剪原始 WGAN 用权重裁剪把评论家限制在 1-Lipschitz工程上很粗暴权重被裁到 [-0.01, 0.01] 之后参数几乎全部挤在边界上评论家的表达能力废掉大半训练照样不稳定。WGAN-GP 的思路是在损失里加一项软约束让评论家对“真实样本和生成样本之间的插值点”的梯度范数尽量接近 1。这个点在测地线上约束它就能让全局梯度变化不会太陡。梯度惩罚的完整实现不长但细节特别容易写错源码里我一般长这样def gradient_penalty(critic, real, fake, lambda_gp10): real: 真实图像 batch fake: 生成器输出且 detach 过的图像 batch batch real.size(0) # 每个样本独立采样插值系数避免整批共享同一个比例 alpha torch.rand(batch, 1, 1, 1, devicereal.device) interpolated (alpha * real (1 - alpha) * fake).requires_grad_(True) d_out critic(interpolated) grad torch.autograd.grad( outputsd_out, inputsinterpolated, grad_outputstorch.ones_like(d_out), create_graphTrue, # 必须保留计算图loss 里还要求二阶导 retain_graphTrue, # 训练循环里还要复用同一张图做反向 )[0] grad grad.view(batch, -1) grad_norm grad.norm(2, dim1) penalty ((grad_norm - 1) ** 2).mean() return lambda_gp * penalty参数说明lambda_gp 是惩罚系数原论文给的 10 在动漫头像这种分布相对集中的任务上同样适用interpolated 的 requires_grad_(True) 必须在调用 critic 之前设置否则 autograd 不会对输入求梯度create_graphTrue 是硬性要求因为惩罚项本身要参与反向传播计算图必须完整保留。提示很多人把 gradient_penalty 里的 alpha 写成一个标量整批共享随机值。实测下来每个样本独立采样能让约束更平滑训练更稳。2.3 两个 loss 的完整形状n_critic 为什么是 5 又为什么常被改小评论家和生成器的损失在 WGAN-GP 里非常干净。评论家要让真实样本的分数尽量高、生成样本的分数尽量低同时满足梯度约束生成器要让评论家对生成样本的输出尽量高。写成代码就是下面几行d_loss -(d_real.mean() - d_fake.mean()) gp g_loss -d_fake_new.mean()注意 d_loss 里的 gp 是加到负号外面的很多人手滑写成了 -(d_real.mean() - d_fake.mean() gp)等于把梯度惩罚也反转了评论家会朝着不满足 Lipschitz 约束的方向跑训练必崩。n_critic 指每更新一次生成器之前先更新几次评论家。原论文默认 5原理是让评论家先逼近真实的 Wasserstein 距离生成器再跟着走。但 256×256 的图前向反向都很贵n_critic5 意味着评论家要做 5 次完整的前向反向训练时间直接翻倍。常见的工程做法是改成 1 到 3同时把学习率降到 1e-4 到 2e-4 来补偿评论家“没练够”带来的波动。我的经验是第一次跑通用 3观察 W 距离估计曲线平稳后再尝试降到 1。3. 数据与分辨率256×256 动漫头像的预处理和显存预算3.1 动漫头像数据集的整理统一尺寸前先做一次头图筛选动漫头像数据集和自然图像数据集有个明显的差异来源五花八门可能混着整页漫画截图、带 UI 的截图、表情包、甚至带水印的同人图。如果直接把所有图片 resize 到 256×256生成器会学到一堆噪点和边缘伪影。我做数据集的第一步不是写代码而是把图片按边长过滤小于 256×256 的直接丢因为放大后边缘是虚的评论家很容易靠“边缘发虚”这个假特征区分真假图生成器学到的是假纹理。数据管线一般长这样from glob import glob import os from PIL import Image from torch.utils.data import Dataset from torchvision import transforms class AnimeFaceDataset(Dataset): def __init__(self, root_dir, img_size256, min_size256): self.paths [] for ext in (*.png, *.jpg, *.jpeg): self.paths.extend(glob(os.path.join(root_dir, ext))) # 排序让多卡训练时每个进程的数据顺序一致 self.paths.sort() self.transform transforms.Compose([ transforms.Resize((img_size, img_size)), transforms.RandomHorizontalFlip(p0.5), transforms.ToTensor(), # 把 [0,1] 归一化到 [-1,1]匹配生成器输出层 Tanh transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]), ]) def __len__(self): return len(self.paths) def __getitem__(self, idx): try: img Image.open(self.paths[idx]).convert(RGB) if img.width min_size or img.height min_size: return self.__getitem__((idx 1) % len(self.paths)) return self.transform(img) except Exception: # 损坏图片在清洗不彻底的数据集里很常见返回相邻样本 return self.__getitem__((idx 1) % len(self.paths))代码里的 min_size 过滤我建议保留而不是完全依赖 resize小图放大后的虚边和压缩伪影是 256×256 动漫头像生成里最容易忽视的暗坑。还有一点数据读取用 PIL 的 convert(RGB) 统一三通道避免出现灰度图或者带 alpha 通道的 PNG 混进训练集导致通道数报错。3.2 归一化到 [-1,1]Tanh 输出与预训练特征提取的差异生成器的输出层用的是 Tanh值域是 [-1,1]。如果真实图片还在 [0,1] 区间评论家从第一眼就能靠“数值分布不对”区分真假生成器会陷入一个很奇怪的处境明明图像内容是对的但整体色调对不上梯度一直在位相上纠正。把输入归一化到 [-1,1] 是跟 Tanh 对齐的必要条件。Normalize 用均值 0.5、标准差 0.5 做的事是 (x - 0.5) / 0.5也就是把 ToTensor 得到的 [0,1] 张量映射到 [-1,1]。这里跟 ImageNet 预训练模型那套 mean[0.485, 0.462, 0.406], std[0.229, 0.224, 0.225] 的归一化不是一回事GAN 训练用的是对称归一化让数据以 0 为中心分布评论家和生成器的初始化才匹配。如果你后面要用 FID 评估注意 InceptionV3 有它自己的一套预处理不要在生成器上套两遍归一化。另外 transforms.Resize 的插值方式我习惯默认的 BILINEAR。最近邻会让线条出现明显锯齿而这些锯齿会被评论家当成“真实感”的信号生成器就会去模仿锯齿而不是画干净的线条。3.3 256×256 的显存预算先按 batch8 跑通再谈放大生成器网络我用 DCGAN 骨架做六层转置卷积从 4×4 一路升到 256×256通道数从 1024 递减到 3。每层 kernel4、stride2、padding1输出尺寸按 (in-1)*2 翻倍4→8→16→32→64→128→256 正好六层。这个结构在显存占用上比 ResNet 骨架小训练速度更快256×256 下是性价比最高的选择。显存预算要按峰值算不能按静态算。评论家每步要同时对真实图、生成图、插值图三个分支做前向梯度惩罚还要 create_graph 保留二阶计算图实际峰值占用是单分支前向的两倍以上。我常用的起步配置如下配置方向网络规模batch size显存预估结论起步配置32→1024 通道 DCGAN 骨架88~10 GB先跑通再优化提质量同上1614~16 GB想提 batch 先确认卡够省显存评论家通道数减半 32→512167~9 GB优先减通道不减 batch显存不够时优先减评论家通道数而不是减 batch size。batch 太小会让插值采样变的稀疏且小 batch 下的统计波动会让 W 距离估计的方差变大训练更容易抖动。4. 训练主循环n_critic、梯度惩罚系数与 EMA 参数4.1 train.py 主循环评论家先走几步生成器再走一步训练循环是整套源码的骨架也是大多数人改了参数就跑不动的地方。下面是我常用的最小可跑版本import torch from torch.optim import Adam # models.py 里定义的 Generator 和 Critic G Generator(latent_dim128) C Critic(in_channels3) g_opt Adam(G.parameters(), lr2e-4, betas(0.5, 0.9)) c_opt Adam(C.parameters(), lr2e-4, betas(0.5, 0.9)) n_critic 3 lambda_gp 10 for step, real_imgs in enumerate(train_loader): real_imgs real_imgs.to(device) batch real_imgs.size(0) z torch.randn(batch, 128, devicedevice) # 1) 更新评论家 n_critic 次 for _ in range(n_critic): fake_imgs G(z).detach() # detach 阻止梯度流入生成器 d_real C(real_imgs).mean() d_fake C(fake_imgs).mean() gp gradient_penalty(C, real_imgs, fake_imgs, lambda_gp) d_loss -(d_real - d_fake) gp # 注意 gp 的符号 c_opt.zero_grad() d_loss.backward() c_opt.step() # 2) 更新生成器一次 fake_imgs G(z) g_loss -C(fake_imgs).mean() g_opt.zero_grad() g_loss.backward() g_opt.step() if step % 500 0: w_dist d_real.item() - d_fake.item() print(fstep {step} d_loss {d_loss.item():.4f} g_loss {g_loss.item():.4f} W {w_dist:.4f})逻辑说明评论家更新时用的 fake_imgs 必须 detach否则评论家反向传播会把梯度也带进生成器生成器更新时用的是新一次前向的 fake_imgs可以直接用评论家的输出算损失不需要重新 detach因为生成器本来就要从这条路拿到梯度。打印里的 w_dist 是 Wasserstein 距离估计这个值比 d_loss 本身更有观察价值。注意如果 w_dist 单调飙升且没有回落迹象先检查 lambda_gp 是否生效再降学习率不要先去调网络结构。4.2 梯度惩罚系数 λ 与学习率默认值之外还要盯什么WGAN-GP 对学习率的容忍度比经典 GAN 高但并不意味着可以随便用 1e-3。经验值是评论家和生成器都用 2e-4Adam 的 beta1 必须取 0.5 而不是默认的 0.9——GAN 训练里动量的惯性太大会让损失绕大圈beta10.5 是 DCGAN 时期就验证过的惯例。lambda_gp 在原论文里对多个数据集都用的 10这个值在动漫头像上够用不要一开始就动它。真正需要调的是 n_critic 与学习率的配合n_critic 降到 1 时评论家的约束更新频率变低学习率也相应降到 1e-4反过来如果 n_critic 保持 5学习率可以试着提到 3e-4但收益通常不明显。参数速查表参数起步值调整方向观察信号lr2e-4降一半W 距离抖动幅度变大betas(0.5, 0.9)不要动 beta0保持 0.5lambda_gp105~20生成图出现明显噪点n_critic31~5W 距离不收敛latent_dim12864~256模式坍塌还有一个新手容易忽略的点WGAN-GP 的 d_loss 数值本身没有绝对意义它不像分类任务的准确率有参考区间。只看 d_loss 从 10 掉到 0.5 就以为训练成功很容易被骗。结合 w_dist 和每 500 步保存的生成图一起看才不至于把均值模糊的脸当成收敛。4.3 EMA 生成器训练后期最便宜的一次质量提升EMA指数移动平均是很多生成模型源码里默认开启的模块做法是对生成器权重做滑动平均推理时用平均后的权重而不是最新权重。GAN 训练后期生成器参数会在一个最优区域附近来回震荡EMA 把这个震荡抹平生成图更干净颜色也更稳定。它不增加任何训练时间只多占一份模型内存。实现很简单import copy class EMA: def __init__(self, model, decay0.999): # 深拷贝一份权重作为滑动平均的载体 self.model copy.deepcopy(model) self.decay decay def update(self, model): with torch.no_grad(): for ema_p, p in zip(self.model.parameters(), model.parameters()): # 滑动平均新的权重只贡献 (1-decay) 的比例 ema_p.mul_(self.decay).add_(p.detach(), alpha1 - self.decay) ema_G EMA(G, decay0.999) # 训练循环里每次 g_opt.step() 后调一次 # ema_G.update(G)decay 参数说明0.999 表示每步只吸收新权重的 0.1%适合训练步数在几万以上的场景如果数据集只有几千张图、总步数很短decay 降到 0.99 让 EMA 跟得快一些。推理和保存 checkpoint 时都用 ema_G.model测试时不要忘了 .eval()否则 BatchNorm 的手滑问题会再次出现。我一般把 EMA 版本的生成器单独保存一份文件名加 ema 后缀防止后面想对比时找不到原始版本。EMA 不是万能药它解决的是“训练后期震荡”和“单帧生成结果不稳定”的问题如果你的模型本身欠拟合EMA 只会把模糊变得更平滑。5. 256×256 训练排障五个必踩的坑与排查顺序5.1 训练中的三个翻车现场从损失爆炸到模式坍塌第一个坑是损失爆炸。现象d_loss 几百步就冲到负几千生成图是纯色或噪点。原因基本是两种学习率过高或者 lambda_gp 因为括号问题被减掉变成负号梯度惩罚形同虚设。排查时先把 d_loss 打印出来人工核对公式确认是 -(d_real - d_fake) gp而不是 -(d_real - d_fake gp)。然后看 w_dist 的绝对值如果上百且持续发散直接把 lr 降到 1e-4 重新跑。这是最容易踩的血泪经验十次崩溃有八次出在这里。第二个坑是模式坍塌。现象固定噪声向量生成的几十张脸轮廓都差不多换 z 只是换个发色或换个眼睛角度。原因一般是 latent_dim 太低64 或 32导致生成器的映射空间不够或者评论家训练得太快、每步都把生成器压死。解决方式按顺序试latent_dim 提到 128n_critic 降到 1打开 EMA。如果 EMA 版本的图像比实时权重版本更有多样性说明是训练后期震荡导致的坍塌EMA 已经能缓解大半。第三个坑是生成图出现纯色块或大面积色斑。现象图片大部分区域颜色均匀发丝和轮廓糊成一团。多数情况下是评论家用了 BatchNorm。BN 在评论家里会把 batch 的统计信息混入特征评论家学会“利用 batch 统计量”而不是“利用图像内容”来打分生成器就被带偏。我实测评论家里干脆不用任何归一化层纯卷积加 LeakyReLU 最稳如果网络太深不收敛再考虑 LayerNorm而不是 BatchNorm。5.2 两个容易被误诊的 eval 与显存问题第四个坑是 eval 偏灰。现象训练过程中每 500 步保存的 sample 图颜色正常但训练完用保存的 pth 权重单独推理出来的图整体偏灰、饱和度偏低。原因不是模型没训练好而是加载生成器后没有切换 eval 模式BatchNorm 在训练模式下用 batch 统计量推理模式下才用 running 统计量如果模型里还残留 BN输出颜色就会整体漂移。解决就是推理代码里加两行ema_G.model.eval() with torch.no_grad(): fake ema_G.model(z)第五个坑是显存峰值 OOM 出现在训练中途。现象程序跑过前几百步突然在评论家更新时报 CUDA out of memory而不是一开始就爆。原因是梯度惩罚的插值样本和二阶计算图在峰值时的显存占用远超静态估算某些 batch 的边缘情况会触碰峰值。解决路径先把 batch 降到 8 跑完一轮确认峰值稳定再考虑加大如果仍然 OOM把评论家通道数减半32→512 封顶不要用混合精度来救——AMP 和二阶梯度叠加会让排查难度翻倍。排查顺序我建议固定先看生成图、再看 w_dist 曲线、最后才看 d_loss。生成图能直接告诉你问题在哪一类w_dist 能区分是训练不收敛还是评估阶段出错d_loss 更像是辅助信号单独看很容易误判。6. 让结果拿得出手固定噪声演化记录与 FID 验证训练调参过程中最值得做的一件事是在训练开始前固定一组噪声向量比如 8×8 共 64 个 z每训练 500 步用它们生成一张拼图并保存。这样训练结束后把几百张拼图按顺序合成视频你能亲眼看到生成器从噪声到清晰脸的完整演化过程比任何 loss 曲线都直观。from torchvision.utils import save_image # 训练前固定一组 z整个训练周期都用它 vis_z torch.randn(64, 128, devicedevice) # 训练循环里每隔 500 步执行一次 if step % 500 0: ema_G.model.eval() with torch.no_grad(): samples ema_G.model(vis_z) save_image(samples, fvis/step_{step:06d}.png, nrow8, normalizeTrue) ema_G.model.train()这套固定噪声记录还有一个作用就是判断“什么时候该停”。如果你发现生成图在 step 30000 之后只是局部抖动已经不再出现结构性的改进说明模型进入平台期继续烧显卡不值得。我早期的教训是只看 FID 数值结果它降到一个区间后就不再变化我以为还能继续训其实视觉上早就停在了同一水平。FID 由于依赖 InceptionV3 特征空间它对动漫风格图的绝对分数并不与 ImageNet 对标应该在训练过程中多次评估、看相对趋势。评估时从验证集随机抽 2000 张生成器用 EMA 版本生成 2000 张计算两者特征分布的均值和协方差距离实现可以直接用 torchmetrics 的 FID 接口注意你安装版本的参数签名以官方文档为准。这套流程走下来你手里的东西就从“一个别人写的源码包”变成了“一个自己能控制训练节奏、知道每个参数在做什么、也能解释异常现象的系统”。我自己的习惯是固定噪声的演化视频留到训练彻底结束再删它是我判断这套配置值不值得复用的第一依据。希望这些参数和踩坑记录能帮你少走几步弯路希望帮到你。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑