ARTICLE DETAIL

资讯详情

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

PyTorch实现CycleGAN:无配对图像风格迁移与循环一致性实战

PyTorch实现CycleGAN:无配对图像风格迁移与循环一致性实战 简介这是一份基于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(inplaceTrue), 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_channels3, out_channels3, n_blocks9): super().__init__() model [ nn.ReflectionPad2d(3), nn.Conv2d(in_channels, 64, 7), nn.InstanceNorm2d(64), nn.ReLU(inplaceTrue), ] # 下采样把 256x256 降到 64x64通道升到 256 model [ nn.Conv2d(64, 128, 3, stride2, padding1), nn.InstanceNorm2d(128), nn.ReLU(inplaceTrue), nn.Conv2d(128, 256, 3, stride2, padding1), nn.InstanceNorm2d(256), nn.ReLU(inplaceTrue), ] # 转换阶段9 个残差块保持尺寸和通道不变 for _ in range(n_blocks): model.append(ResidualBlock(256)) # 上采样从 64x64 恢复到 256x256 model [ nn.ConvTranspose2d(256, 128, 3, stride2, padding1, output_padding1), nn.InstanceNorm2d(128), nn.ReLU(inplaceTrue), nn.ConvTranspose2d(128, 64, 3, stride2, padding1, output_padding1), nn.InstanceNorm2d(64), nn.ReLU(inplaceTrue), 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×256中cycleGAN 最稳定的默认配置Unet Generator256×256高保留空间结构但容易把源域纹理带过去2.2 PatchGAN 判别器实现与感受野判别器如果用普通二分类把整张图压成一个概率生成器很容易钻空子全局颜色对了局部纹理全是糊的。cycleGAN 通常搭配 PatchGAN 判别器它不输出单个概率而是输出一个二维特征图每个位置对应原图一个局部区域的真假判断最后取均值作为最终分数。class PatchDiscriminator(nn.Module): def __init__(self, in_channels3): super().__init__() def conv_block(in_f, out_f, stride, use_normTrue): layers [nn.Conv2d(in_f, out_f, 4, stridestride, padding1)] if use_norm: layers.append(nn.BatchNorm2d(out_f)) layers.append(nn.LeakyReLU(0.2, inplaceTrue)) return nn.Sequential(*layers) self.model nn.Sequential( # 第一层不用 BatchNorm避免小 batch 下统计量抖动 conv_block(in_channels, 64, 2, use_normFalse), conv_block(64, 128, 2), conv_block(128, 256, 2), conv_block(256, 512, 1), nn.Conv2d(512, 1, 4, padding1), ) 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 内所有样本的均值和方差混在一起相当于抹掉了每张图独立的颜色风格信息。马的照片和斑马的照片如果出现在同一个 batchBN 会强行把两者的亮度分布拉齐。InstanceNorm 则对每一张图的每一个通道单独做归一化保留图片自身的色温和纹理对比度。把生成器里的 InstanceNorm 换成 BatchNorm 后最典型的症状是训练曲线看起来正常但生成图整体发灰、边缘出现波纹。这类问题不是调学习率能救回来的属于网络结构层面的退化。另外一个容易踩的坑是残差块中间的卷积不要加偏置因为后面紧跟归一化层偏置会被归一化抵消白白浪费参数。cycleGAN 官方实现里的升采样层用转置卷积如果你发现生成图有棋盘格伪影可以考虑换成 PixelShuffle但要注意它会改变输出通道布局需要在代码里做一次重排。3. cycleGAN 损失函数设计对抗损失、循环一致性 L1 与身份损失参数怎么定cycleGAN 能收敛损失函数的组合方式占了七成功劳。只靠对抗损失生成器很容易找到“骗过判别器”的捷径把所有马的照片都生成同一张带斑纹的图。判别器看不出破绽但生成的内容已经和输入的马毫无关系。这种情况在生成对抗网络里叫模式坍缩循环一致性损失就是为了压制它。3.1 对抗损失选 MSE 还是 BCELSGAN 与普通二分类判别器的输出有两种主流封装方式。用 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 而不是 L2L2 对像素偏差做平方惩罚对少量大偏差过度敏感梯度在小扰动区域会变得很小生成图容易偏模糊。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_idtidentity loss 不是所有任务都适合。当两个域差异极大比如素描变照片identity loss 过大会压制转换强度输出的图片还是原图样子。一般来说色彩相关任务先开 0.5观察生成图颜色是否正确出现颜色漂移就提到 1.0转换不到位就降到 0.1。各损失项的配置最终可以归纳成下面这张表损失项实现方式默认权重调参方向对抗损失MSELSGAN1.0训练初期波动大时降到 0.5循环一致性L110.0模糊降到 5保持不了结构提到 15身份损失L10.5颜色漂移提到 1.0转换不足降到 0.14. 跑通 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, interpolationtransforms.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, transformtransform) dataset_B datasets.ImageFolder(data/trainB, transformtransform) loader_A iter(DataLoader(dataset_A, batch_size1, shuffleTrue)) loader_B iter(DataLoader(dataset_B, batch_size1, shuffleTrue))先放大到 286 再随机裁 256等于给训练集引入了随机尺度和随机位置增强能显著降低判别器对边缘纹理的过拟合。归一化到 [-1, 1] 是因为生成器最后一层是 Tanh输出值域与输入必须一致。插值方式我显式写成 BICUBICPyTorch 不同版本默认插值不同不指定的话容易出现细微的预处理差异导致换版本后结果对不上。batch size 设成 1 是 cycleGAN 最稳定的选择增大 batch size 虽然能加速但会改变 BatchNorm 在判别器里的统计行为不建议一上来就改。4.2 Adam 超参、学习率衰减与完整训练循环优化器参数几乎被固定成一个惯例Adamlr0.0002beta10.5beta20.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_lambdalambda_rule) sched_D torch.optim.lr_scheduler.LambdaLR(opt_D, lr_lambdalambda_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, size50): 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 迅速趋近 0G 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(dim1) 得到单通道误差图误差集中在边缘是正常现象说明模型在做纹理迁移误差集中在整张图的全局区域说明模型在做像素拷贝风格根本没有迁移出去。最后一个实用技巧是给每个 epoch 固定住随机种子把同一个输入反复喂给生成器这样生成的输出序列可用于前后对比。如果训练到中段出现明显跳变优先怀疑学习率衰减曲线在那个阶段变化过快而不是生成器结构坏了。本文还有配套的精品资源点击获取
返回列表