PyTorch实现CycleGAN:无配对图像风格迁移与循环一致性实战
2026/9/16 2:59:47 网站建设 项目流程

简介:这是一份基于PyTorch的CycleGAN完整项目实现,主要面向深度学习中图像风格迁移、生成对抗网络方向的学习者与研究者,可帮助理解无监督图像到图像翻译的核心机制。压缩包共30个文件,以Python源码为主,覆盖生成器、判别器、模型搭建、训练与测试脚本,同时包含示例图片、数据下载脚本和说明文档,整体仅900KB,轻量易用,便于快速阅读和调试。项目清晰拆分了双向生成器与两个判别器,并实现循环一致性损失,可完成马与斑马互转等经典风格迁移任务;配套样例图片和数据脚本能辅助读者直接运行与验证效果。目前已有1213人学习下载,适合希望从代码层面掌握CycleGAN原理,并进一步开展二次应用的PyTorch使用者。

1. cycleGAN 是什么:没有配对照片,也能让马变成斑马

训练一个风格迁移模型,最头疼的不是网络写不出来,而是找不到成对的训练数据。让卫星图变成地图、让马的普通照片变成斑马照片,这类任务几乎拿不到像素级对齐的图片对。cycleGAN 只用两个图片集合就能完成训练:一个目录放马的照片,另一个目录放斑马的照片,格式、数量都不需要一一对应。它不是靠“长得像”硬拼,而是引入循环一致性,让两张图片先后经过两个方向的生成器之后还能还原回原图。这个约束把无配对问题变成了可监督问题,也让风格迁移、图像翻译、域适配这类任务有了统一的落地路线。这套用 PyTorch 训练生成对抗网络的流程,适合入门深度学习的图像生成方向,也适合处理真实项目里数据永远凑不齐配对图的情况。

2. 用 PyTorch 搭 cycleGAN 生成器与判别器:ResNet-9Block、PatchGAN 与实例归一化

cycleGAN 的完整结构包括两个生成器和两个判别器,四个网络在训练时互相制约。生成器 G 负责把 A 域图片翻译成 B 域,F 负责把 B 域图片翻译回 A 域;判别器 D_B 判断输入是不是真实 B 域图片,D_A 对应判断 A 域。生成器的结构决定了最终画面能保留多少细节,也直接决定显存占用和训练速度。

2.1 生成器的残差块实现与反射填充

256×256 分辨率下,cycleGAN 生成器最常见的做法是 ResNet-9Block 结构:卷积把图片降到 64×64,经过 9 个残差块做域转换,再上采样回 256×256。残差块本身不改变尺寸和通道数,作用是在内容特征上叠加目标域的风格信息。先看残差块实现:

import torch.nn as nn class ResidualBlock(nn.Module): def __init__(self, in_channels): super().__init__() self.block = nn.Sequential( nn.ReflectionPad2d(1), # 镜像填充,避免边缘伪影 nn.Conv2d(in_channels, in_channels, 3), nn.InstanceNorm2d(in_channels), nn.ReLU(inplace=True), nn.ReflectionPad2d(1), nn.Conv2d(in_channels, in_channels, 3), nn.InstanceNorm2d(in_channels), ) def forward(self, x): # 残差连接:输入和卷积结果直接相加 return x + self.block(x)

ReflectionPad2d 是镜像填充,比常见的零填充少很多边缘黑边和振铃效应。卷积后没有在残差连接之后再加 ReLU,因为常规做法是激活只放在残差分支内部,输出的恒等映射部分由下一个模块继续处理。InstanceNorm2d 在风格迁移任务里几乎属于标配,具体原因在 2.3 展开。

把残差块组装成完整生成器时,需要把下采样和上采样段的通道变化对应好。下面是一份可直接跑的 ResNet-9Block 生成器:

class ResNetGenerator(nn.Module): def __init__(self, in_channels=3, out_channels=3, n_blocks=9): super().__init__() model = [ nn.ReflectionPad2d(3), nn.Conv2d(in_channels, 64, 7), nn.InstanceNorm2d(64), nn.ReLU(inplace=True), ] # 下采样:把 256x256 降到 64x64,通道升到 256 model += [ nn.Conv2d(64, 128, 3, stride=2, padding=1), nn.InstanceNorm2d(128), nn.ReLU(inplace=True), nn.Conv2d(128, 256, 3, stride=2, padding=1), nn.InstanceNorm2d(256), nn.ReLU(inplace=True), ] # 转换阶段:9 个残差块保持尺寸和通道不变 for _ in range(n_blocks): model.append(ResidualBlock(256)) # 上采样:从 64x64 恢复到 256x256 model += [ nn.ConvTranspose2d(256, 128, 3, stride=2, padding=1, output_padding=1), nn.InstanceNorm2d(128), nn.ReLU(inplace=True), nn.ConvTranspose2d(128, 64, 3, stride=2, padding=1, output_padding=1), nn.InstanceNorm2d(64), nn.ReLU(inplace=True), nn.ReflectionPad2d(3), nn.Conv2d(64, out_channels, 7), nn.Tanh(), ] self.model = nn.Sequential(*model) def forward(self, x): return self.model(x)

代码中 n_blocks 可以改成 6、9、18 等值。128×128 的输入用 6 个残差块就够,256×256 用 9 个,再往上增加残差块会让训练时间接近线性增长,但对生成质量的提升往往有限。最后一层用 Tanh 是因为输入图片已经归一化到 [-1, 1],输出范围要严格对齐。

生成器方案适用分辨率显存占用特点
ResNet-6Block128×128快速验证网络是否走通
ResNet-9Block256×256cycleGAN 最稳定的默认配置
Unet Generator256×256保留空间结构,但容易把源域纹理带过去

2.2 PatchGAN 判别器实现与感受野

判别器如果用普通二分类,把整张图压成一个概率,生成器很容易钻空子:全局颜色对了,局部纹理全是糊的。cycleGAN 通常搭配 PatchGAN 判别器,它不输出单个概率,而是输出一个二维特征图,每个位置对应原图一个局部区域的真假判断,最后取均值作为最终分数。

class PatchDiscriminator(nn.Module): def __init__(self, in_channels=3): super().__init__() def conv_block(in_f, out_f, stride, use_norm=True): layers = [nn.Conv2d(in_f, out_f, 4, stride=stride, padding=1)] if use_norm: layers.append(nn.BatchNorm2d(out_f)) layers.append(nn.LeakyReLU(0.2, inplace=True)) return nn.Sequential(*layers) self.model = nn.Sequential( # 第一层不用 BatchNorm,避免小 batch 下统计量抖动 conv_block(in_channels, 64, 2, use_norm=False), conv_block(64, 128, 2), conv_block(128, 256, 2), conv_block(256, 512, 1), nn.Conv2d(512, 1, 4, padding=1), ) def forward(self, x): return self.model(x)

输入 256×256 时,输出大约是 30×30 的评分图。每个评分格子的感受野约等于 70×70 像素,这意味着判别器判断的不是全局风格对不对,而是每个局部窗口是否真实。PatchGAN 的另一个好处是和输入分辨率解耦,换到 512×512 输入只需要调整下采样层数,不需要重写网络结构。LeakyReLU 的负斜率 0.2 是 GAN 里的常见取值,保留负区间的梯度,避免判别器某些神经元训死。判别器内部使用 BatchNorm,这里不像生成器那样换成 InstanceNorm,因为判别器的任务是分类真假,需要一个稳定的 batch 统计量来统一尺度。

2.3 实例归一化在风格迁移中的作用

BatchNorm 在风格迁移里有一个很微妙的问题:它会把一个 batch 内所有样本的均值和方差混在一起,相当于抹掉了每张图独立的颜色风格信息。马的照片和斑马的照片如果出现在同一个 batch,BN 会强行把两者的亮度分布拉齐。InstanceNorm 则对每一张图的每一个通道单独做归一化,保留图片自身的色温和纹理对比度。

把生成器里的 InstanceNorm 换成 BatchNorm 后,最典型的症状是训练曲线看起来正常,但生成图整体发灰、边缘出现波纹。这类问题不是调学习率能救回来的,属于网络结构层面的退化。另外一个容易踩的坑是:残差块中间的卷积不要加偏置,因为后面紧跟归一化层,偏置会被归一化抵消,白白浪费参数。cycleGAN 官方实现里的升采样层用转置卷积,如果你发现生成图有棋盘格伪影,可以考虑换成 PixelShuffle,但要注意它会改变输出通道布局,需要在代码里做一次重排。

3. cycleGAN 损失函数设计:对抗损失、循环一致性 L1 与身份损失参数怎么定

cycleGAN 能收敛,损失函数的组合方式占了七成功劳。只靠对抗损失,生成器很容易找到“骗过判别器”的捷径,把所有马的照片都生成同一张带斑纹的图。判别器看不出破绽,但生成的内容已经和输入的马毫无关系。这种情况在生成对抗网络里叫模式坍缩,循环一致性损失就是为了压制它。

3.1 对抗损失选 MSE 还是 BCE:LSGAN 与普通二分类

判别器的输出有两种主流封装方式。用 BCEWithLogitsLoss 表示真假的二分类概率,用 torch.nn.MSELoss 表示对真假的评分。MSE 版本等价于 LSGAN——最小二乘生成对抗网络。LSGAN 会对远离决策边界的样本同样提供梯度,训练初期比 BCE 稳定,模式坍缩的概率更低。cycleGAN 原版实现用的就是 MSE,如果你从网上找到的代码里损失函数写的是 NLLLoss 或 BCELoss,那大概率是某个旧版本改动后的结果。

# 判别器对真实斑马图的评分 real_pred = D_B(real_B) # 判别器对生成斑马图的评分 fake_pred = D_B(fake_B.detach()) # detach 阻断反向传播到生成器 loss_D = 0.5 * (F.mse_loss(real_pred, torch.ones_like(real_pred)) + F.mse_loss(fake_pred, torch.zeros_like(fake_pred)))

生成器那边只计算 fake_pred 与 1 的 MSE,不反向传播判别器。用 detach 把假样本从计算图中摘出来,是为了让生成器梯度不会串到判别器参数上。这是训练顺序里最容易漏的一个细节,漏掉之后两个网络会同时更新,参数互相拉扯,loss 曲线看起来就像噪声。判别器损失前的 0.5 是缩放系数,让判别器步长和生成器保持在同一个量级,这个系数可加可不加,但只要加了,学习率对应也要做细微调整。

3.2 循环一致性损失代码与 L1 选择理由

循环一致性的核心逻辑是:马的图片生成一张假斑马,再用反向生成器把假斑马还原成马,还原结果应当与原图接近。另一个方向同理。周期损失直接用 L1 距离,代码很短:

recon_A = F_G2A(fake_B) # 假斑马还原回假马 recon_B = F_A2B(fake_A) # 假马还原回假斑马 cyc_loss = (F.l1_loss(recon_A, real_A) + F.l1_loss(recon_B, real_B)) * lambda_cyc

为什么用 L1 而不是 L2?L2 对像素偏差做平方惩罚,对少量大偏差过度敏感,梯度在小扰动区域会变得很小,生成图容易偏模糊。L1 的梯度恒定为 1,对边缘细节更友好。lambda_cyc 默认取 10,这是原版实验里比较稳的值。如果你发现生成图学会了风格但丢了轮廓,把 lambda_cyc 调到 15;如果画面纹理丰富但整体发虚,调到 5 试试。

3.3 身份损失要不要开:用于保住源域颜色

identity loss 的含义是:把一张目标域图片输入源域生成器,期望输出尽量保持不变。比如把真实的斑马照片输入马生成器,理想结果应该还是斑马,而不是被强行加一匹马的样子。这个约束强制生成器不要乱改颜色和光照。

idt_B = G_A2B(real_B) # 真实斑马输入马生成器,理想输出仍是斑马 idt_A = G_B2A(real_A) idt_loss = (F.l1_loss(idt_B, real_B) + F.l1_loss(idt_A, real_A)) * lambda_idt

identity loss 不是所有任务都适合。当两个域差异极大,比如素描变照片,identity loss 过大会压制转换强度,输出的图片还是原图样子。一般来说,色彩相关任务先开 0.5,观察生成图颜色是否正确;出现颜色漂移就提到 1.0,转换不到位就降到 0.1。各损失项的配置最终可以归纳成下面这张表:

损失项实现方式默认权重调参方向
对抗损失MSE(LSGAN)1.0训练初期波动大时降到 0.5
循环一致性L110.0模糊降到 5,保持不了结构提到 15
身份损失L10.5颜色漂移提到 1.0,转换不足降到 0.1

4. 跑通 cycleGAN 的最小可复现配置:数据加载、Adam 超参与 epoch 规划

环境侧建议直接用 Anaconda 单独建一个 pytorch 环境,装好 torch、torchvision 和 tensorboard 就能开始。CPU 也可以跑通,只是 256×256 的 epoch 时间会很长;有 GPU 时单卡就能训练,cycleGAN 对显存的要求不算极端。下面这套配置是我在本地复现时固定下来的,改动尽可能少,适合先跑通再看效果。

4.1 torchvision 数据加载与 286 到 256 随机裁剪

数据目录按两个域分开:trainA 放源域图片,trainB 放目标域图片。用 torchvision.datasets.ImageFolder 读取最省事,配合随机左右翻转增强。

from torchvision import transforms, datasets from torch.utils.data import DataLoader transform = transforms.Compose([ transforms.Resize(286, interpolation=transforms.InterpolationMode.BICUBIC), transforms.RandomCrop(256), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]) ]) dataset_A = datasets.ImageFolder("data/trainA", transform=transform) dataset_B = datasets.ImageFolder("data/trainB", transform=transform) loader_A = iter(DataLoader(dataset_A, batch_size=1, shuffle=True)) loader_B = iter(DataLoader(dataset_B, batch_size=1, shuffle=True))

先放大到 286 再随机裁 256,等于给训练集引入了随机尺度和随机位置增强,能显著降低判别器对边缘纹理的过拟合。归一化到 [-1, 1] 是因为生成器最后一层是 Tanh,输出值域与输入必须一致。插值方式我显式写成 BICUBIC,PyTorch 不同版本默认插值不同,不指定的话容易出现细微的预处理差异,导致换版本后结果对不上。batch size 设成 1 是 cycleGAN 最稳定的选择,增大 batch size 虽然能加速,但会改变 BatchNorm 在判别器里的统计行为,不建议一上来就改。

4.2 Adam 超参、学习率衰减与完整训练循环

优化器参数几乎被固定成一个惯例:Adam,lr=0.0002,beta1=0.5,beta2=0.999。beta1 从 PyTorch 默认的 0.9 改成 0.5,是为了减少历史梯度对当前更新的影响,让对抗过程更稳。训练计划常见的是前 100 个 epoch 保持学习率不变,后 100 个 epoch 线性衰减到 0。用 LambdaLR 可以少写很多手写逻辑:

def lambda_rule(epoch): total_epochs = 200 if epoch < 100: return 1.0 return max(0.0, (total_epochs - epoch) / 100) sched_G = torch.optim.lr_scheduler.LambdaLR(opt_G, lr_lambda=lambda_rule) sched_D = torch.optim.lr_scheduler.LambdaLR(opt_D, lr_lambda=lambda_rule)

训练循环按照“一次判别器更新、一次生成器更新”的顺序交替执行:

# 生成器前向 fake_B = G_A2B(real_A) # 马 -> 假斑马 fake_A = G_B2A(real_B) # 斑马 -> 假马 # 先更新判别器 D_B d_real_B = D_B(real_B) d_fake_B = D_B(fake_B.detach()) loss_D_B = 0.5 * (F.mse_loss(d_real_B, torch.ones_like(d_real_B)) + F.mse_loss(d_fake_B, torch.zeros_like(d_fake_B))) opt_D.zero_grad() loss_D_B.backward() opt_D.step()

判别器 D_A 的更新方式完全相同。生成器更新时把对抗损失、循环一致性损失和身份损失全部加在一起:

g_adv = F.mse_loss(D_B(fake_B), torch.ones_like(d_fake_B)) + \ F.mse_loss(D_A(fake_A), torch.ones_like(d_fake_A)) recon_A = G_B2A(fake_B) recon_B = G_A2B(fake_A) g_cyc = F.l1_loss(recon_A, real_A) + F.l1_loss(recon_B, real_B) idt_A = G_B2A(real_A) idt_B = G_A2B(real_B) g_idt = F.l1_loss(idt_A, real_A) + F.l1_loss(idt_B, real_B) loss_G = g_adv + 10.0 * g_cyc + 0.5 * g_idt opt_G.zero_grad() loss_G.backward() opt_G.step()

这段顺序里的关键点是:判别器先看当前 batch 的假图,生成器再根据这张假图的判别结果做反向传播。如果调换顺序,生成器更新时用的还是上一轮判别器对旧图的判断,损失曲线会剧烈震荡。每个 epoch 结束后调用一次 sched_G.step() 和 sched_D.step(),不要在 batch 内部反复衰减学习率。整个 200 epoch 的规划适合大多数风格迁移任务;如果你的数据量很小,比如一个域只有几十张图,建议把前 100 epoch 改成前 50,总 epoch 减到 120,否则后期学习率衰减过程占掉太多时间。

4.3 判别器先崩了:历史池 buffer 的写法

训练早期的典型问题是判别器聪明过头,立刻把生成图全部判为假,生成器梯度爆炸,输出变成噪声。经典解法是历史池:给判别器喂一批上一轮生成的旧图,让它不能只依赖当前 batch 的特征来做判断。实现可以用一个简单队列:

import random class ImagePool: def __init__(self, size=50): self.size = size self.images = [] def query(self, img): if self.size == 0: return img if len(self.images) < self.size: self.images.append(img) return img # 一半概率保留旧图,一半概率返回当前图 if random.random() > 0.5: idx = random.randint(0, self.size - 1) old = self.images[idx].clone() self.images[idx] = img return old return img

每次计算判别器损失之前,把 fake_B 和 fake_A 先过一遍池,返回结果再喂给判别器。池子大小 50 在 256×256 任务上是比较稳妥的选择,太大则生成器更新速度被拖慢,太小起不到缓冲作用。从池里取出的图片在送入判别器前要再做一次 detach,防止梯度从这个分支回流到生成器。这个技巧不能完全替代学习率调节,但它能把训练早期的崩溃概率降低一大截。

5. cycleGAN 训练结果验证技巧:判别器先崩、生成图偏模糊怎么修

模型能不能用,跑到第 30 个 epoch 基本能看出来。验证时一定要固定住同一张测试图片,不要每轮随机采样,不然你看到的变化分不清是模型进步了还是输入变了。我一般每 5 个 epoch 保存一张三行拼图:源图、生成图、重建图,三个并排看变化。重建图如果一直模糊,说明循环一致性在起作用,但像素细节没有完全兜住。

判别器先崩的现象很典型:D loss 迅速趋近 0,G loss 一路上涨,生成图全是噪声。先检查是不是忘了加历史池,其次把判别器的学习率从 2e-4 降到 1e-4,或者改成每两个 batch 才更新一次判别器。另一个更隐蔽的原因是判别器网络太强,比如把 PatchGAN 的最后一层卷积换成了全连接,判别能力远超生成器,怎么调都救不回来。遇到这种情况直接换回标准 PatchGAN,不要继续加网络深度。

生成图偏模糊时,问题大多在损失权重而不是网络结构。循环一致性的 L1 权重偏大,模型为了把重建损失压下去,不敢在纹理上做大修改。把 lambda_cyc 从 10 降到 5,同时观察重建图的边缘是否变清晰。颜色整体发灰时先别动损失,检查生成器的 Tanh 输出在 TensorBoard 里是不是没做反归一化。TorchVision 的 make_grid 会把 [-1, 1] 的值原样转到 8bit 显示,看到发灰的图是显示问题,不是模型问题。

想进一步定位是哪一类像素还原不好,可以把循环一致性误差画成热力图:直接用 torch.abs(recon_A - real_A).mean(dim=1) 得到单通道误差图,误差集中在边缘是正常现象,说明模型在做纹理迁移;误差集中在整张图的全局区域,说明模型在做像素拷贝,风格根本没有迁移出去。最后一个实用技巧是给每个 epoch 固定住随机种子,把同一个输入反复喂给生成器,这样生成的输出序列可用于前后对比。如果训练到中段出现明显跳变,优先怀疑学习率衰减曲线在那个阶段变化过快,而不是生成器结构坏了。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询