简介这份源码资源面向深度学习入门者与动漫图像生成爱好者提供了一套基于WGAN-GP算法生成256×256像素动漫头像的完整实现可用于理解生成对抗网络的训练流程与梯度惩罚机制并作为二次开发或课程实验的实操起点。压缩包共26个文件、约1.32MB其中2个Python源文件承担生成器、判别器与训练逻辑的核心实现11个PNG图片展示不同训练阶段生成的头像效果6个XML与1个iml文件用于IDE项目及参数配置另有readme说明与Git忽略文件辅助项目管理。目前已有306人学习。读者可借此掌握WGAN-GP相较传统GAN在缓解模式崩塌、提升训练稳定性方面的具体做法观察生成样本随训练推进的清晰度变化并基于现有目录结构快速替换数据集或调整网络层用于在线游戏、虚拟形象、表情包等场景的头像定制实验。1. 拆开这份 WGAN-GP 动漫头像源码256×256 像素到底能不能直接跑拿到一个只有两个 Python 文件的 GAN 项目第一反应通常是「这么点代码能生成 256×256 的动漫头像」我一开始也这么想。这份基于 WGAN-GP 的动漫头像生成源码核心逻辑就压在main.py和check.py两个文件里配套 11 张已经生成好的 PNG 结果图、若干 XML 配置和一份 readme。它解决的不是「从零训练一个工业级动漫头像模型」的问题而是让你在一台普通带显卡的机器上把 WGAN-GP 的完整训练链路跑通并且肉眼看到 256×256 分辨率的头像从噪声里长出来。适合谁想搞懂 WGAN-GP 梯度惩罚项到底怎么写进 loss 的人、课程设计需要一份能演示的 GAN 源码的人、以及手里有动漫头像数据集想快速验证训练稳定性的人。不适合谁指望开箱就产出商用级高清头像的人——256×256 是它的上限不是起点。下面按「资源结构 → 环境与数据 → 训练与调参 → 排错 → 进阶验证」的顺序拆能抄的地方我直接给命令和参数。2. 资源结构与 WGAN-GP 关键实现两个 py 文件里藏了什么2.1 目录清单与各文件职责先把压缩包解开目录结构大致是这样我按实际用途归类而不是照抄文件树文件/目录类型作用main.pyPython训练主入口含生成器、判别器、训练循环check.pyPython结果检查/推理脚本加载权重出图data/目录训练数据存放位置内含.keep占位save/目录训练过程与最终生成图输出位置result/目录已生成的 11 张 PNG 结果图save6/save10/save15…readme.txt文本项目说明、运行方式.idea/配置IDE 工程配置含gan256.iml、misc.xml等.gitignore配置版本控制忽略规则这里有个容易被忽略的点data/和save/里都只有.keep说明作者是把空目录结构提交上来、把真实数据和权重排除在外的。也就是说这份源码给的是「代码骨架 结果样例」不是「数据 权重」。你得自己准备数据集这一点在动手前必须清楚否则跑起来第一件事就是报找不到图片。2.2 生成器与判别器的结构选型256×256 这个分辨率不算大但对纯卷积 GAN 来说已经需要足够深的网络。常见做法是生成器用「全连接升维 多层转置卷积上采样」判别器用「多层步长卷积下采样 全连接输出标量」。为什么不用 ResNet 残差块因为这个体量的项目残差带来的收益有限反而增加调试成本。我一般会按下面的思路组织生成器import torch import torch.nn as nn class Generator(nn.Module): def __init__(self, z_dim128, base64): super().__init__() # 从噪声 z 映射到 4x4 的特征图再逐级上采样到 256x256 self.fc nn.Linear(z_dim, base * 8 * 4 * 4) self.net nn.Sequential( # 4x4 - 8x8 nn.ConvTranspose2d(base * 8, base * 4, 4, 2, 1, biasFalse), nn.BatchNorm2d(base * 4), nn.ReLU(True), # 8x8 - 16x16 nn.ConvTranspose2d(base * 4, base * 2, 4, 2, 1, biasFalse), nn.BatchNorm2d(base * 2), nn.ReLU(True), # 16x16 - 32x32 nn.ConvTranspose2d(base * 2, base, 4, 2, 1, biasFalse), nn.BatchNorm2d(base), nn.ReLU(True), # 32x32 - 64x64 nn.ConvTranspose2d(base, base // 2, 4, 2, 1, biasFalse), nn.BatchNorm2d(base // 2), nn.ReLU(True), # 64x64 - 128x128 nn.ConvTranspose2d(base // 2, base // 4, 4, 2, 1, biasFalse), nn.BatchNorm2d(base // 4), nn.ReLU(True), # 128x128 - 256x256最后一层用 tanh 把像素压到 [-1,1] nn.ConvTranspose2d(base // 4, 3, 4, 2, 1, biasFalse), nn.Tanh(), ) def forward(self, z): x self.fc(z).view(z.size(0), -1, 4, 4) return self.net(x)逻辑说明z_dim128是噪声维度太小会导致多样性不足太大收敛慢128 是这个分辨率的常用值。base64控制通道基数显存不够就降到 32。每一层ConvTranspose2d的kernel4, stride2, padding1是标准的两倍上采样配置能把特征图边长翻倍。最后一层输出 3 通道对应 RGBTanh把值域对齐到[-1,1]这一点必须和后面数据归一化方式一致否则生成图会整体偏灰或过曝。判别器则相反逐级下采样到 4×4 后接全连接输出一个标量。注意 WGAN-GP 的判别器严格叫 Critic最后一层不要加 Sigmoid因为它输出的是 Wasserstein 距离里的评分不是概率。这是新手最容易翻车的地方之一。2.3 梯度惩罚项的实现位置WGAN-GP 相比原始 WGAN 的核心改动是用梯度惩罚替代权重裁剪。惩罚项要在「真实样本和生成样本的连线上随机插值」处计算梯度范数逼它接近 1。常见写法def gradient_penalty(critic, real, fake, device, lambda_gp10): batch_size real.size(0) # 在 0~1 之间采样插值系数形状对齐到图像 alpha torch.rand(batch_size, 1, 1, 1, devicedevice) # 真实样本与生成样本的线性插值 interpolated (alpha * real (1 - alpha) * fake).requires_grad_(True) score critic(interpolated) # 对插值样本求梯度 grad torch.autograd.grad( outputsscore, inputsinterpolated, grad_outputstorch.ones_like(score), create_graphTrue, retain_graphTrue, only_inputsTrue )[0] grad grad.view(batch_size, -1) # 梯度范数逼近 1偏离部分作为惩罚 gp ((grad.norm(2, dim1) - 1) ** 2).mean() return lambda_gp * gp参数说明lambda_gp10是原论文给出的经验值调大惩罚更强、训练更稳但收敛慢调小则约束不足。create_graphTrue必须开因为惩罚项本身要参与反向传播。interpolated要requires_grad_(True)否则autograd.grad拿不到梯度。这段是整个项目最值钱的部分也是main.py里最该逐行读的地方。3. 环境搭建与数据准备让 256×256 训练真正跑起来3.1 依赖安装与版本对齐这份源码是 PyTorch 体系先确认环境。我一般用 conda 隔离避免和系统里的其他框架打架conda create -n gan256 python3.8 -y conda activate gan256 # 按自己 CUDA 版本装对应 torch这里以 cu118 为例 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install numpy pillow tqdm逻辑说明Python 3.8 是这类老项目的稳妥选择太新的版本有时会在某些算子或依赖上出兼容问题。torchvision用来做图像读取和 transform。pillow处理 PNGtqdm看进度。装完先跑一句python -c import torch; print(torch.cuda.is_available())返回True才说明 GPU 可用返回False就先去查驱动和 CUDA 版本别急着往下走。3.2 数据集组织与预处理data/目录默认是空的你得自己塞图。动漫头像数据集常见来源是各类头像爬取集合但要注意版权和合规自己用就本地放着。组织方式建议data/ faces/ 0001.png 0002.png ...所有图片统一到 256×256。如果原图尺寸不一用脚本批量处理别指望训练时 transform 帮你兜底——GAN 对输入尺寸很敏感from PIL import Image import os src_dir data/raw dst_dir data/faces os.makedirs(dst_dir, exist_okTrue) for name in os.listdir(src_dir): if not name.lower().endswith((.png, .jpg, .jpeg)): continue img Image.open(os.path.join(src_dir, name)).convert(RGB) # 先短边缩放再中心裁剪避免拉伸变形 w, h img.size scale 256 / min(w, h) img img.resize((int(w * scale), int(h * scale)), Image.BICUBIC) w, h img.size left, top (w - 256) // 2, (h - 256) // 2 img img.crop((left, top, left 256, top 256)) img.save(os.path.join(dst_dir, name))逻辑说明直接resize((256,256))会把非正方形头像压扁人脸比例失真生成结果也会跟着歪。先按短边等比缩放、再中心裁剪是头像类数据的标准做法。convert(RGB)统一通道避免混入灰度图或带 alpha 的 PNG 导致ToTensor后通道数不一致。3.3 数据加载与归一化对齐数据进网络前要归一化到[-1,1]和生成器最后一层Tanh对齐from torchvision import transforms from torch.utils.data import DataLoader from torchvision.datasets import ImageFolder tf transforms.Compose([ transforms.Resize((256, 256)), transforms.ToTensor(), # 转到 [0,1] transforms.Normalize([0.5]*3, [0.5]*3) # 映射到 [-1,1] ]) dataset ImageFolder(data, transformtf) loader DataLoader(dataset, batch_size16, shuffleTrue, num_workers4, drop_lastTrue)参数说明batch_size16是 256×256 分辨率下 8G 显存的稳妥值显存大可以上 32。drop_lastTrue很重要WGAN-GP 的梯度惩罚按 batch 计算最后一批数量不足会让插值系数形状对不上。num_workers4按 CPU 核数调Windows 下如果报多进程错误就设成 0。4. 训练循环与超参调优WGAN-GP 的 n_critic 怎么设4.1 训练循环骨架WGAN-GP 的典型训练节奏是「判别器更新 n_critic 次生成器更新 1 次」。这个比例直接决定训练稳不稳import torch from torch import optim device cuda if torch.cuda.is_available() else cpu G, D Generator().to(device), Critic().to(device) opt_G optim.Adam(G.parameters(), lr1e-4, betas(0.5, 0.9)) opt_D optim.Adam(D.parameters(), lr1e-4, betas(0.5, 0.9)) z_dim, n_critic, lambda_gp 128, 5, 10 for epoch in range(200): for i, (real, _) in enumerate(loader): real real.to(device) bs real.size(0) # ---- 训练判别器 n_critic 次 ---- for _ in range(n_critic): z torch.randn(bs, z_dim, devicedevice) fake G(z).detach() loss_D D(fake).mean() - D(real).mean() \ gradient_penalty(D, real, fake, device, lambda_gp) opt_D.zero_grad(); loss_D.backward(); opt_D.step() # ---- 训练生成器 1 次 ---- z torch.randn(bs, z_dim, devicedevice) fake G(z) loss_G -D(fake).mean() opt_G.zero_grad(); loss_G.backward(); opt_G.step()逻辑说明判别器 loss 是D(fake).mean() - D(real).mean()WGAN 里判别器要最大化真假评分差所以最小化它的相反数再加上梯度惩罚。生成器只关心把D(fake)推高所以 loss 取负。betas(0.5, 0.9)是 GAN 训练的经典设置动量别用默认的 0.9否则容易震荡。4.2 关键超参对照参数推荐值作用与调整方向z_dim128噪声维度小则多样性差大则收敛慢n_critic5判别器每轮次数训练不稳可加到 5~10lambda_gp10梯度惩罚权重原论文经验值lr1e-4学习率太大直接发散batch_size16受显存限制drop_last 必须开betas(0.5, 0.9)Adam 动量别用默认4.3 训练过程怎么判断好坏WGAN-GP 的好处是 loss 有物理意义-loss_D大致反映 Wasserstein 距离理论上应该缓慢下降并趋于稳定。如果它一路狂掉或者剧烈震荡多半是学习率太大或梯度惩罚没生效。每训练若干轮存一次图用check.py或自己写个采样脚本把固定噪声生成的图存到save/肉眼对比。result/里那 11 张 PNG 就是作者留下的效果参照可以拿来对比自己的输出质量。5. 避坑与排查跑 WGAN-GP 最常见的五个翻车点5.1 生成图全是噪点或纯色块现象训练几十轮后生成图仍是雪花噪点或单一色块看不出头像轮廓。原因判别器太强生成器梯度消失或者数据归一化没对齐[-1,1]导致输入分布和生成分布差太远。解决先确认 transform 里Normalize([0.5]*3,[0.5]*3)没漏再把n_critic从 5 降到 1~2给生成器更多更新机会检查判别器最后一层有没有误加 Sigmoid。5.2 loss 变成 NaN现象训练几百步后 loss 突然变nan之后全废。原因梯度惩罚里grad.norm出现极端值或学习率过大导致参数爆炸。解决把lr降到 5e-5 试在gradient_penalty里对grad加一句grad grad 1e-8防除零开启torch.nn.utils.clip_grad_norm_做梯度裁剪兜底。5.3 显存爆掉 OOM现象跑到一半报CUDA out of memory。原因batch_size太大或 256×256 下特征图通道基数base设太高。解决batch_size降到 8base从 64 降到 32用torch.cuda.empty_cache()清理缓存确认没有在循环里累积计算图生成器更新时fake别带detach之外的残留图。5.4 模式崩塌生成的头像长得都一样现象不管输入什么噪声出来的头像几乎一模一样。原因生成器找到了「骗过判别器」的单一解多样性丢失。解决适当增大z_dim检查是不是n_critic太小导致判别器没学好在生成器 loss 里可以尝试加入小的多样性正则但 WGAN-GP 本身对模式崩塌已有缓解优先排查数据是否过于单一。5.5 加载权重或续训报错现象check.py加载保存的权重时报size mismatch或unexpected key。原因保存和加载时的网络结构定义不一致比如改了base或z_dim却没同步。解决权重和结构定义必须成对保存建议存torch.save({G: G.state_dict(), config: {...}}, path)加载时先按 config 重建网络再load_state_dict。6. 进阶验证用固定噪声和插值看模型到底学到了什么训练跑通只是第一步真正判断一个 WGAN-GP 模型学没学到东西我习惯做两件事固定噪声看一致性、噪声插值看连续性。固定噪声选一组固定的z每隔若干轮用同一组z出图存成序列。如果模型在进步同一组噪声生成的图应该从噪点逐渐收敛成稳定头像而不是每轮都大变样。这能帮你判断训练是否在收敛而不是在原地抖动。import torch from PIL import Image import torchvision.utils as vutils G.eval() fixed_z torch.randn(64, 128, devicedevice) # 固定下来别每次重采样 with torch.no_grad(): fake G(fixed_z) # 反归一化回 [0,1] 再存图 fake (fake * 0.5 0.5).clamp(0, 1) vutils.save_image(fake, save/fixed_grid.png, nrow8)逻辑说明fixed_z一定要在循环外生成一次并固定否则每轮噪声不同你根本分不清变化来自模型还是来自输入。(fake*0.50.5)是把[-1,1]还原到[0,1]和前面的归一化严格对应。nrow8排成 8 列网格方便一眼看多样性。噪声插值取两个噪声z1、z2在它们之间线性插值生成一串z出图后如果画面是平滑过渡的比如发型、表情渐变说明潜空间是连续的模型学到了有意义的表征如果中间突然跳变或糊成一团说明潜空间还有断裂。这是检验 GAN 质量很直观的一招比只看 loss 曲线靠谱得多。z1 torch.randn(1, 128, devicedevice) z2 torch.randn(1, 128, devicedevice) alphas torch.linspace(0, 1, 10, devicedevice).view(-1, 1) z_interp z1 * (1 - alphas) z2 * alphas # 10 个插值点 with torch.no_grad(): imgs G(z_interp) imgs (imgs * 0.5 0.5).clamp(0, 1) vutils.save_image(imgs, save/interp.png, nrow10)参数说明torch.linspace(0,1,10)生成 10 个插值系数view(-1,1)是为了和z的维度广播对齐。nrow10把 10 张图排成一行从左到右就是z1到z2的渐变过程。血泪经验是别只盯着 loss 数字GAN 的 loss 和图像质量之间没有线性关系固定噪声网格图才是我每次必看的「后悔药」。从那以后我每次训 GAN都强制先跑一遍固定噪声出图确认模型真的在收敛再谈调参。希望帮到你。本文还有配套的精品资源点击获取
