ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

PyTorch GAN实现源码解析:从DCGAN到CycleGAN的实践指南

PyTorch GAN实现源码解析:从DCGAN到CycleGAN的实践指南 简介一套基于Pytorch实现多种GAN模型的完整项目源码面向深度学习研究者与算法工程师覆盖GAN、CycleGAN、GRAGAN等主流变体重点解决图像生成、风格迁移及跨域转换等任务的实战落地问题。压缩包共88个文件包含60个py源码、14张效果图、9个动态gif示例及训练脚本、说明文档等整体大小约29.96MB便于快速解析模型结构与复现实验。资源附带完整流程教程讲解从GAN基本原理、Pytorch模型搭建到CycleGAN循环一致性损失设计、GRAGAN全局结构建模再到性能评估与特定任务应用的全链路方法同时涵盖DCGAN、WGAN、pix2pix、StarGAN等代表性架构结合源码与效果图可直观对比不同生成对抗网络的差异。适合具备一定深度学习基础、希望深入源码细节的读者已有105人学习下载可用于图像生成、风格迁移、跨域转换等项目研发与教学参考。1. 从开箱到读源码这堆PyTorch GAN实现为什么值得花时间过一遍如果你在找一份能“一边跑一边改”的 GAN 代码集合这份压缩包比单独对着论文公式造轮子要实在得多。解压后会看到implementations/下列着 dcgan、wgan_gp、cyclegan、discogan、stargan、esrgan 等几十个独立子目录每个都自带数据加载、模型定义和训练入口。对已经跑过 PyTorch 基础框架的开发者最省力的路径是先用它建立“生成器—判别器—目标函数”的对应关系再按任务复制一份改装。这套代码沿用了早期 PyTorch GAN 工程的写法直接跑会遇到兼容性问题后面几章会集中讲踩坑点。如果你能读完并动手把某个子目录改成别的 GAN 变体这份源码的价值就真正落地了。2. 先吃透主干从 DCGAN 到 WGAN-GP 的生成器、判别器和训练循环不论你最终要跑 CycleGAN 还是 ESRGAN源码里所有模型都复用同一套“数据加载—模型定义—训练循环”的三段式结构。先从一个最简单的 DCGAN 入手把网络骨架和目标函数看懂再去翻 CycleGAN 就不会被一堆缩写劝退。这一章同时解决两个问题包里的目录如何映射到具体模型以及训练循环里那几行关键代码抄错后会发生什么。2.1 解压后的目录里到底有什么压缩包解压后没有复杂的工程外壳核心就是implementations/和data/以及一个 README。我第一次打开这个包时建议不要在 IDE 里挨个点文件先在终端里把结构看清楚unzip Pytorch_GAN.zip -d ~/gan-exp cd ~/gan-exp ls implementations/执行后能看到每个模型一个目录assets/里还有训练过程中的 gif 和模型结构图。下面这张表整理了从新手到进阶都适合下手的目录入口目录模型建议先看它的原因gan原始 GAN只含 MLP最容易追踪对抗损失dcgan深度卷积 GAN建立了 Conv 层做生成的默认写法wgan_gp带梯度惩罚的 WGAN展示 Wasserstein 距离和梯度惩罚cyclegan循环一致性 GAN多生成器多判别器的典型代表srgan超分辨率 GAN涉及内容损失和对抗损失的加权README 里会写明各模型的数据集来源和基本用法但我更推荐先打开assets/wgan_gp.gif这类训练动图观察收敛过程。WGAN-GP 的图像往往是从一片噪声到边缘清晰中间很少出现原始 GAN 那种剧烈跳变有了这个印象再看代码里的gradient_penalty更容易理解它在约束什么。2.2 生成器与判别器的代码骨架DCGAN 给出了默认答案在implementations/dcgan/dcgan.py里生成器用一个线性层把随机噪声映射到 8×8 的特征图再用两层上采样卷积放大到 32×32。关键部分集中在conv_blocks# implementations/dcgan/dcgan.py 中生成器结构的关键片段 self.init_size 32 // 4 # 8以此作为卷积输入的初始边长 self.l1 nn.Sequential(nn.Linear(opt.latent_dim, 128 * self.init_size ** 2)) self.conv_blocks nn.Sequential( nn.BatchNorm2d(128), nn.Upsample(scale_factor2), nn.Conv2d(128, 128, 3, stride1, padding1), nn.BatchNorm2d(128, 0.8), nn.LeakyReLU(0.2, inplaceTrue), nn.Upsample(scale_factor2), nn.Conv2d(128, 64, 3, stride1, padding1), nn.BatchNorm2d(64, 0.8), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(64, opt.channels, 3, stride1, padding1), nn.Tanh(), )这里有两个容易忽略的参数。Upsample(scale_factor2)决定输出尺寸每次翻倍从 32×32 开始只能做 32、64、128 这类 2 的整数次幂尺寸如果你要生成 96×96 的图不能只改这一个参数得同时调整init_size或者在上采样流程里加一层。最后的Tanh()把输出压到 [-1,1]对应数据预处理时ToTensor()后还要归一化到 [-1,1]。如果忘了归一化判别器看到真实图像和生成图像的值域不一致loss 会出现持续波动图像却始终模糊。2.3 只替换目标函数就从 GAN 变成 WGAN、LSGAN同样的生成器和判别器只要改变损失计算方式就得到不同模型。这套源码最直观的价值是把这些差异都摊在同一个训练框架下。下面是我在对比时常看的一张表模型生成器目标判别器目标GAN / DCGAN-E[log D(G(z))]E[log D(x)] E[log(1-D(G(z)))]WGAN-E[D(G(z))]E[D(G(z))] - E[D(x)]权重裁剪WGAN-GP-E[D(G(z))]E[D(G(z))] - E[D(x)] λ·GPLSGANE[(D(G(z))-1)^2]E[(D(x)-1)^2] E[D(G(z))^2]DRAGAN-E[D(G(z))]同 WGAN-GP但 GP 在真实样本邻域采样理解这张表后再去看wgan_gp和wgan_div的差别就很容易它们都在解决 WGAN 权重裁剪导致的参数集中在上下界的问题区别只是梯度惩罚的计算方式。在源码里改模型变体时通常不需要动网络结构只需要替换损失函数和判别器最后一层的激活函数。2.4 训练循环中最容易抄错的四行以wgan_gp/wgan_gp.py的训练循环为例我平时改代码时最关心这十几行的顺序for i, (imgs, _) in enumerate(dataloader): imgs imgs.cuda() z torch.randn(imgs.shape[0], opt.latent_dim).cuda() # 更新判别器真实样本分数要高生成样本分数要低 optimizer_D.zero_grad() real_validity discriminator(imgs) fake generator(z).detach() # 判别器更新时隔离生成器梯度 fake_validity discriminator(fake) gradient_penalty compute_gradient_penalty(discriminator, imgs.data, fake.data) d_loss -real_validity.mean() fake_validity.mean() opt.lambda_gp * gradient_penalty d_loss.backward() optimizer_D.step() # 每隔 n_critic 步更新一次生成器 if i % opt.n_critic 0: optimizer_G.zero_grad() fake generator(z) fake_validity discriminator(fake) g_loss -fake_validity.mean() g_loss.backward() optimizer_G.step()这段代码里最值得注意的几个参数lambda_gp是梯度惩罚系数常见取 10n_critic表示判别器每更新 5 次生成器才更新 1 次这是为了让判别器先提供稳定的梯度。上面对fake generator(z)做了detach()避免判别器反向传播把梯度传入生成器如果漏掉.detach()训练时会出现生成器损失也在下降但生成图像始终不清晰的现象因为生成器拿到的梯度不是由当前参数状态计算出来的旧的计算图残留会导致更新方向被污染。3. 解开 CycleGAN循环一致性损失与非配对图像转换的完整流程CycleGAN 在源码包里的位置是implementations/cyclegan/它和前面 DCGAN 最大的不同是同时管理四个网络而且数据不是成对的。理解 CycleGAN 时最忌讳只盯着对抗损失真正决定生成结果能不能保持结构的是循环一致性损失和身份损失。这一章把模型文件、损失计算和训练命令串起来。3.1 模型变量与四个网络的对应关系打开implementations/cyclegan/models.py你会看到一个CycleGAN类里有四个核心属性netG_A2B、netG_B2A、netD_A、netD_B。它们的输入输出关系如下变量名输入 → 输出在训练中的职责netG_A2B真实 A 图 → 伪造 B 图把 A 域风格迁移到 B 域netG_B2A真实 B 图 → 伪造 A 图把 B 域风格迁移回 A 域netD_A真实 A / 伪造 A → 判别分数让伪造的 A 尽量接近真实 AnetD_B真实 B / 伪造 B → 判别分数让伪造的 B 尽量接近真实 B在训练中opt.lambda_cyc和opt.lambda_id两个权重把三部分损失拼在一起。我一般会在第一次运行 CycleGAN 时把models.py里生成器输出层的Tanh()和数据集里Normalize的均值方差对照检查一遍很多生成发灰的问题都出在这里。3.2 循环一致性损失与身份损失的代码形态在train.py里生成器总损失的计算可以抽象成下面这几行# cyclegan/train.py 中生成器损失的核心计算与源码一致的常见写法 loss_GAN_A2B criterion_GAN(self.netD_B(fake_B), valid) loss_GAN_B2A criterion_GAN(self.netD_A(fake_A), valid) loss_cycle_A criterion_cycle(self.netG_B2A(fake_B), real_A) loss_cycle_B criterion_cycle(self.netG_A2B(fake_A), real_B) loss_identity_A criterion_identity(self.netG_B2A(real_A), real_A) loss_identity_B criterion_identity(self.netG_A2B(real_B), real_B) loss_G (loss_GAN_A2B loss_GAN_B2A) \ opt.lambda_cyc * (loss_cycle_A loss_cycle_B) \ opt.lambda_id * (loss_identity_A loss_identity_B)criterion_cycle在代码里通常是 L1 损失也就是把fake_B再经过netG_B2A转回来与real_A做逐像素差距criterion_identity同样用 L1但它直接拿真实图像过生成器要求生成器对已经是目标域的输入尽量不做改动。循环一致性损失不参与对抗过程只约束生成器不能无中生有地改结构身份损失则负责稳定颜色分布一般在风格化任务里能明显改善色偏。参数上CycleGAN 论文常用的组合是lambda_cyc10.0身份损失权重lambda_id从 0.5 开始试即可源码里如果默认没开身份损失先不要急着怀疑模型能力往往是把这里调起来之后整体效果才稳定。3.3 下载非配对数据集与启动训练资源包自带了下载脚本因此在 Linux 或 macOS 下先用命令下载数据集bash data/download_cyclegan_dataset.sh horse2zebra python implementations/cyclegan/train.py --dataset_name horse2zebra \ --n_epochs 200 --decay_epoch 100 --batch_size 1download_cyclegan_dataset.sh会从 CycleGAN 官方数据源把对应数据集下载到data/horse2zebra/。--n_epochs 200 --decay_epoch 100表示前 100 个 epoch 保持初始学习率后 100 个 epoch 线性衰减到 0这是 CycleGAN 最常用的调度策略--batch_size 1是因为 size 256×256 的图像在单卡上占显存较多显存不够时这个值不建议继续调大。下表是这个脚本支持的部分数据集名称数据集名称任务horse2zebra马 ↔ 斑马apple2orange苹果 ↔ 橙子summer2winter_yosemite夏景 ↔ 冬景monet2photo莫奈画作 ↔ 照片训练命令里的dataset_name要和下载时保持一致否则datasets.py找不到对应目录。Windows 下跑这个脚本要注意bash命令需要 Git Bash 或 WSL 环境直接用 cmd 执行会报“无法识别 bash”。3.4 训练过程中看什么CycleGAN 的 loss 曲线不能像分类任务那样只看下降因为对抗损失和循环一致性损失存在博弈关系。我一般把生成样例的图片目录打开每隔 50 个 epoch 手动检查三件事第一fake_B是否出现了 B 域的纹理特征比如马身上开始出现斑纹第二fake_A是否把斑马纹还原成接近真实马的颜色如果还原出来一团糊说明循环一致性没有约束住第三背景结构有没有发生明显形变比如草地被无中生有地重绘。如果loss_cycle一直偏高可能是生成器发现保持输入不变可以降低循环损失此时对抗损失被压制可以适当降低lambda_cyc或者提高判别器学习率。4. 变体对比与训练排错GRAGAN/DRAGAN、梯度惩罚和 FID 验证一份 GAN 合集里最常被问到的不是如何训练而是“为什么我跑出来的结果不像 README 里的 gif”。这一章把源码中容易让人困惑的 GRAGAN 命名理清楚再用梯度惩罚对比、FID 评估和检查清单把训练过程变成可验证的事。4.1 包里的 GRAGAN 到底对应哪个模型先说明一个容易踩的命名坑。摘要里写的 GRAGAN在源码包目录里并没有名为gragan的子目录包含的是dragan和wgan_div两个目录。社区里经常有人把 DRAGAN 打成 GRAGAN如果你是在课程笔记里看到这个词大概率指的就是dragan。DRAGAN 全称是 Deep Regret Analytic GAN它和 WGAN-GP 一样使用梯度惩罚但惩罚的采样区域不同。摘要里提到的 CRF 结构在这个包里没有对应代码不要按这个名字去找 CRF 模块。在使用这套源码时不要把 GRAGAN 当成不相干的第三种模型而是把它放在“梯度惩罚家族”里去理解。这样在你日后阅读其他 GAN 变体时只要看到gradient_penalty字样就知道它和 WGAN-GP、DRAGAN 在数学上处于同一族。4.2 DRAGAN 与 WGAN-GP 的梯度惩罚代码对比在implementations/dragan/dragan.py里梯度惩罚的计算大致如下# dragan.py 中 DRAGAN 梯度惩罚的常见实现 alpha torch.rand(imgs.size(0), 1, 1, 1).cuda() x_hat (alpha * imgs.data (1 - alpha) * fake.data).requires_grad_(True) pred_hat discriminator(x_hat) gradients autograd.grad( outputspred_hat, inputsx_hat, grad_outputstorch.ones_like(pred_hat), create_graphTrue, retain_graphTrue, )[0] gradient_penalty ((gradients.norm(2, dim1) - 1) ** 2).mean()这段代码先对输入做requires_grad_(True)因为我们需要对判别器输出求输入梯度而 PyTorch 默认关闭输入梯度。alpha是 0 到 1 之间的随机数x_hat是真实样本和生成样本之间的插值点DRAGAN 则会在真实样本附近施加扰动而 WGAN-GP 是在真假样本连线上采样。下面这张表让两个方法的适用场景更清楚方法采样位置惩罚目标代码入口WGAN-GP真实与生成的插值点全局 Lipschitz 约束wgan_gp/wgan_gp.pyDRAGAN真实样本附近扰动点局部 Lipschitz 约束dragan/dragan.pyWGAN-DIV多组差值近似Wasserstein 散度wgan_div/wgan_div.py实际训练时WGAN-GP 一般比 DRAGAN 更稳但 DRAGAN 在样本数量少、真实分布覆盖率不够的场景下有时能缓解模式坍缩。如果训练 DCGAN 出现大面积重复生成直接复制wgan_gp目录替换损失函数这个操作成本最低。4.3 用 torchmetrics 计算 FID 判定生成质量源码包里没有现成的 FID 计算脚本但这是判断生成分布是否接近真实分布的关键。我通常会在项目目录下写一个独立的评估脚本用torchmetrics来做# 使用 torchmetrics 计算 FID from torchmetrics.image.fid import FrechetInceptionDistance fid FrechetInceptionDistance(feature2048) fid.update(real_images, realTrue) fid.update(fake_images, realFalse) print(fFID: {fid.compute():.3f})feature2048表示使用 Inception 网络中倒数第二层的 2048 维特征real_images和fake_images需要是 0 到 255 范围内的 RGB 张量形状为(N, C, H, W)。FID 值越小表示两个特征分布越接近。注意在多次评估时每次都要重新实例化FrechetInceptionDistance否则上一次批量里的特征统计会残留。4.4 训练失败的检查单很多人看到 loss 曲线异常就直接调学习率其实更快的办法是把下面这张表打印出来逐项对照现象优先排查项D_loss 快速掉到 0G_loss 不降判别器太强降低判别器学习率或减少n_critic生成图像重复多样性差模式坍缩尝试 DRAGAN、增加噪声或使用 mini-batch discriminationloss 为 NaN检查数据集是否有 NaN、学习率是否过大图像灰度值集中在中部确认生成器输出有没有经过Tanh()数据集是否归一化到 [-1,1]最后再补一点在 PyTorch 2.x 下跑旧代码时如果遇到torch.autograd.Variable相关报错直接把Variable(...)改成普通的torch.Tensor就行。因为 PyTorch 0.4 之后Variable已被废弃但很多早期源码还保留这个写法用pip install -r requirements.txt之前最好先手动装好对应 CUDA 版本的 PyTorch再用 requirements 补缺失依赖。5. 把这套源码改造成自己的项目环境准备、启动顺序与换主干三步掌握了前面的骨架后最后一步是把它变成“能跟着你的数据集走”的工程。这一章给出一套平时从零开始改包的方法先搭环境再选最接近任务的模型最后用三步完成改造。5.1 环境搭建与启动顺序许多 PyTorch 安装教程都会建议用 Anaconda 直接建环境这里只要注意 Python 版本和 CUDA 版本匹配即可。例如conda create -n gan python3.8 -y conda activate gan pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 pip install -r requirements.txt--index-url是 PyTorch 官方的 wheel 源能避免用 conda 默认源装错 CPU 版本。装完之后先用python implementations/gan/gan.py跑一个最小 GAN确认环境没有坑再跑wgan_gp和cyclegan这两个分别验证梯度惩罚和多网络流程。如果某个模型跑起来了但结果没变化先看 README 里对应入口参数的说明再对照前面讲的损失计算不要急着换模型。5.2 换主干的三步改造拿到一个新数据集时我习惯复制一个最接近的模型目录而不是从头写。三步分别是复制目录、改输入输出、替换损失函数。用命令示意cp -r implementations/cgan implementations/my_cgan cd implementations/my_cgan然后把models.py里opt.channels改成新数据集的通道数Generator里的opt.latent_dim改成你的编码维度。最后把判别器输出层的nn.Sigmoid()注释掉因为 LSGAN 等目标函数要求输出为原始 logits# 以 LSGAN 为例去掉 sigmoid 后目标函数只剩均方误差 loss_D 0.5 * ((discriminator(real) - 1) ** 2).mean() 0.5 * (discriminator(fake) ** 2).mean() loss_G 0.5 * ((discriminator(fake) - 1) ** 2).mean()这样改完训练循环里的optimizer.zero_grad()和backward()都不需要动模型就在同一个工程结构上换成了 LSGAN。5.3 推荐的第一组实验起点根据任务类型可以从下面这张表选起点风格迁移或图像翻译直接用cyclegan超分辨率用srgan或esrgan条件生成用cgan或acgan手头没有配对数据但想做域迁移优先看discogan和pixelda。任务起点代码改动幅度风格迁移 / 图像翻译cyclegan中需处理数据集格式图像超分srgan / esrgan中损失权重较敏感条件生成cgan / acgan小只需改条件标签无监督域适应pixelda / discogan大需要准备域标签提示更换数据集后先用两张图跑通前向传播确认输出尺寸是任务期望的(batch, channels, H, W)再开训练GAN 的 bug 往往在第一个 step 就会被放大早验 shape 比早看 loss 更重要。本文还有配套的精品资源点击获取
返回列表