CycleGAN在MATLAB中的实现:未配对图像转换与风格迁移实战

CycleGAN在MATLAB中的实现:未配对图像转换与风格迁移实战 简介循环一致性对抗网络CycleGAN的MATLAB实现与运行结果资源包主要面向本科、硕士阶段及深度学习、图像处理方向的学习者目标是帮助理解在无需成对训练样本条件下完成图像风格迁移/转换的核心原理。资源包内共包含5个文件其中有2个MATLAB程序文件核心训练脚本、苹果与橙子数据集加载脚本、1个说明文档、1张训练效果图以及1个动态GIF演示压缩包整体约28.5MB结构简洁便于直接对照学习。当前已有238人学习浏览适合作为课堂实验、课程设计或毕业设计的参考项目。代码基于MATLAB 2014a/2019a平台编写并附有运行结果可以直观查看对抗生成网络在具体任务上的训练过程和效果也能通过说明文件快速复现结果并且通过观察训练迭代时的图像变化把抽象的对抗训练过程落到具体代码上。学习者可在现有代码基础上尝试调整网络结构、优化策略或替换数据集进一步扩展图像风格迁移实验逐步掌握生成器、判别器与循环一致性损失的作用为后续研究做铺垫。1. CycleGAN 是什么类型转移与未配对训练的那条主线如果你手里有一批真实照片和一批梵高画作想让模型把照片转成梵高风格常见做法是准备成对数据同一内容先拍照片再临摹成画但现实中几乎找不到这种配对。CycleGAN 解决的就是这个「没有配对也能练」的问题它不需要内容一一对应只需要两个域各自的图像集合就能学会照片到绘画、白昼到黑夜、马到斑马这类风格迁移。这套思路的实际价值在于很多行业场景里配对数据比模型更难搞。比如把 CT 影像转成 MRI 风格做数据增强、把卫星图转成地图样式、把白天照片转成夜间用于自动驾驶测试这些需求遇到的最大瓶颈从来不是网络结构而是「上哪儿找同场景的另一张图」。CycleGAN 在 2017 年被提出后迅速成为这类未配对图像转换的基线方案它的损失函数设计、网络结构安排至今仍是生成对抗网络入门到进阶的重要样本。本文围绕这个标题做全文拆解先讲 CycleGAN 的原理主线再给出 MATLAB 实现的关键路径从数据准备、网络构建、训练循环到结果验收均给出可直接运行的思路与代码。我默认你已经装了 MATLAB R2020b 以上版本并安装了 Deep Learning Toolbox。如果你也想在自己数据集上复现这篇可以作为从零开始的路线图。2. MATLAB 跑 CycleGAN 的数据准备目录结构、加载函数与构造非配对批量2.1 数据组织与加载的常见做法我在处理这类项目时先把两个域的数据分别放进两个文件夹例如 Photo 和 Van Gogh。CycleGAN 训练时是「一批来自 A 域、一批来自 B 域」两个批次的图像内容不需要对应这正是它与 Pix2Pix 的本质差异。实际落地时数据量不必很大每个域 300 到 1000 张就足够看到明显效果因为循环一致性损失起到了较强的约束作用。MATLAB 中加载图像最朴素但最可控的方式是 imageDatastore 配合自定义读图函数。下面的代码负责把两个域的图片路径全部扫出来建立两个数据源并定义读取时要做的基础预处理。% 设置数据根目录A 域为 photoB 域为 van_gogh dataRoot ./data; imdsA imageDatastore(fullfile(dataRoot, photo), IncludeSubfolders, true, ... LabelSource, none, ReadFcn, readAndResize); imdsB imageDatastore(fullfile(dataRoot, van_gogh), IncludeSubfolders, true, ... LabelSource, none, ReadFcn, readAndResize); function img readAndResize(filename) img imread(filename); if size(img, 3) 1 img repmat(img, [1 1 3]); % 灰度图转 3 通道避免维度报错 end img imresize(img, [256 256]); % CycleGAN 经典输入尺寸 img im2single(img); % 转换为 [0,1] 范围内的 single img (img - 0.5) * 2; % 归一化到 [-1,1]与 tanh 输出匹配 end这段代码的关键在于最后的归一化CycleGAN 生成器输出层通常用 tanh值域是 [-1,1]如果输入图像还停在 [0,1] 或 uint8 范围整个训练过程会非常不稳定生成图像会偏灰或出现伪影。imageDatastore 的 ReadFcn 返回的是处理后的图像后续训练循环直接调用 datastore 就能拿数据。2.2 非配对批量采样shuffle 一次就够非配对训练的另一个实现要点在于采样策略。常见做法是一次性读取两个域的全部图像索引每个 epoch 开始前各打乱一次然后按 mini-batch 大小依次取。不需要保证 A 和 B 在语义上有任何关联。下面给出一个简洁的批量采样函数。function [batchA, batchB] sampleBatch(imdsA, imdsB, imgsA, imgsB, batchSize) % 在随机偏移位置取连续 batchSize 张实现类随机采样 idxA randi(numel(imgsA) - batchSize 1); idxB randi(numel(imgsB) - batchSize 1); batchA zeros(256, 256, 3, batchSize, single); batchB zeros(256, 256, 3, batchSize, single); for i 1:batchSize batchA(:,:,:,i) read(imdsA); batchB(:,:,:,i) read(imdsB); end end注意这段代码是一种教学简化如果追求更严格的随机性应该先运行imdsA shuffle(imdsA)再顺序读取而不是用 randi 取连续切片。MATLAB 的 read 函数每次调用会移动 datastore 内部游标直接在上面的函数里反复 read 需要保证 datastore 已被 reset。我实际项目里更常用的是把所有图像预读到内存因为 256x256 的 single 图像每张约 0.8 MB几百张也就几百 MB完全可接受。提示如果你的数据集单张超过 512x512 且数量上千才需要考虑用 augmentedImageDatastore 或自定义 minibatchqueue 做流水线读取否则预读全部数据到内存最简单可靠。3. 用 MATLAB 搭建 CycleGAN 生成器与判别器选哪个网络结构3.1 生成器ResNet 块与 transposedConv 的搭配CycleGAN 的生成器常见有两种U-Net 结构和 ResNet 结构。对于风格迁移这类输入输出同尺寸的任务ResNet 结构更常用因为它通过残差连接保留原始结构避免 U-Net 那种强跳跃连接导致内容过度保留、风格转换不彻底。我一般搭一个 9 个 ResNet 块的生成器每个块包含两个卷积层和残差相加。MATLAB 里可以用 layerGraph 搭配 additionLayer 构建残差连接不用手工写 forward 函数代码结构会比较清晰。核心结构如下layers [ imageInputLayer([256 256 3], Name, input, Normalization, none) convolution2dLayer(7, 64, Padding, 3, Name, conv1) groupNormalizationLayer(channel-wise, 64, Name, norm1) reluLayer(Name, relu1) convolution2dLayer(3, 128, Padding, 1, Stride, 2, Name, conv2) groupNormalizationLayer(channel-wise, 128, Name, norm2) reluLayer(Name, relu2) convolution2dLayer(3, 256, Padding, 1, Stride, 2, Name, conv3) groupNormalizationLayer(channel-wise, 256, Name, norm3) reluLayer(Name, relu3) ]; lgraph layerGraph(layers); % 手动添加一个 ResNet 块示例实际代码用循环加 9 个 lgraph addLayers(lgraph, [ convolution2dLayer(3, 256, Padding, 1, Name, res_conv1) groupNormalizationLayer(channel-wise, 256, Name, res_norm1) reluLayer(Name, res_relu1) convolution2dLayer(3, 256, Padding, 1, Name, res_conv2) groupNormalizationLayer(channel-wise, 256, Name, res_norm2) ]); lgraph addLayers(lgraph, additionLayer(2, Name, res_add)); lgraph connectLayers(lgraph, relu3, res_conv1); lgraph connectLayers(lgraph, res_norm2, res_add/in1); lgraph connectLayers(lgraph, relu3, res_add/in2);这里用 groupNormalizationLayer 而不是 batchNormalizationLayer原因是 CycleGAN 通常 batch size 很小比如 1 到 4BatchNorm 在这种 batch size 下统计量抖动很厉害训练不稳定。GroupNorm 按通道分组做归一化不受 batch size 影响是 CycleGAN 训练更稳的常见选择。MATLAB 从 R2021a 开始支持 groupNormalizationLayer如果版本低就换成 instanceNormalizationLayer 或自己写一个自定义层。3.2 判别器PatchGAN 与 70x70 感受野CycleGAN 的判别器用的是 PatchGAN它的输出不是单个标量而是一个 NxN 的特征图每个元素判断图像中一个局部块的真假。这样设计的好处是参数量小、训练稳定同时能够保留局部纹理的判别能力。70x70 PatchGAN 意味着每个输出像素对应输入图上约 70x70 的感受野。MATLAB 里用卷积层直接构建 PatchGAN最后一层不加 sigmoid训练时用最小二乘损失配合原始分数计算。代码结构如下function lgraph buildDiscriminator() layers [ imageInputLayer([256 256 3], Name, input, Normalization, none) convolution2dLayer(4, 64, Stride, 2, Padding, 1, Name, conv1) leakyReluLayer(0.2, Name, lrelu1) convolution2dLayer(4, 128, Stride, 2, Padding, 1, Name, conv2) groupNormalizationLayer(channel-wise, 128, Name, d_norm2) leakyReluLayer(0.2, Name, lrelu2) convolution2dLayer(4, 256, Stride, 2, Padding, 1, Name, conv3) groupNormalizationLayer(channel-wise, 256, Name, d_norm3) leakyReluLayer(0.2, Name, lrelu3) convolution2dLayer(4, 512, Stride, 1, Padding, 1, Name, conv4) groupNormalizationLayer(channel-wise, 512, Name, d_norm4) leakyReluLayer(0.2, Name, lrelu4) convolution2dLayer(4, 1, Stride, 1, Padding, 1, Name, conv5) ]; lgraph layerGraph(layers); end判别器的 stride 设置很关键。前几层用 stride 2 做下采样最后一层恢复 stride 1 保持空间维度不为 1。PatchGAN 的输出尺寸取决于输入尺寸和卷积步长你把网络放进analyzeNetwork里看一眼就知道输出是 30x30 还是 15x15不影响训练正确性但会改变感受野大小。提示CycleGAN 原文使用了 instance normalizationMATLAB 的 groupNormalizationLayer 当 group 数等于通道数时等价于 Instance Norm 的一种形式。我的经验是 GroupNorm 在 batch size 大于 4 时效果更稳定具体选择可以在小数据集上对比。3.3 完整网络前向传播的封装思路搭建完 layerGraph 后训练时既可以用dlnetwork封装也可以直接在自定义 training loop 里调用forward。我更推荐转成 dlnetwork因为后续计算梯度时接口更统一。如果网络包含残差结构用 layerGraph 的 connectLayers 连接完再转 dlnetwork 会很方便。dlnetG dlnetwork(lgraph); % 生成器 dlnetD dlnetwork(buildDiscriminator()); % 判别器注意生成器里如果有 additionLayerconnectLayers 时要注意输入端口命名additionLayer(2) 的两个输入端口是 in1 和 in2漏掉任何一个都会报连接错误。这类错误在 MATLAB 里很常见排错时在 connectLayers 之后调用analyzeNetwork(lgraph)检查连通性即可。4. MATLAB 训练循环中的关键参数对抗损失、循环一致性与验证策略4.1 损失函数的具体计算CycleGAN 的损失分三块两个域的对抗损失、两个方向的循环一致性损失、以及可选的 identity loss。对抗损失用最小二乘形式LSGAN原因是原始 GAN 的 sigmoid 交叉熵在训练后期容易梯度消失。MATLAB 中直接用mean((d_output - 1).^2)这类计算即可。循环一致性损失是 CycleGAN 的核心把 A 转成 B 再转回 A要和原图尽量一致。这样做的物理意义是强迫生成器学习到内容保持的映射而不是随意改变图像结构。具体损失如下% 前向循环A - fakeB - recA fakeB forward(dlnetG, dlA, Outputs, output); recA forward(dlnetG, fakeB, Outputs, output); % 反向循环B - fakeA - recB fakeA forward(dlnetG, dlB, Outputs, output); recB forward(dlnetG, fakeA, Outputs, output); lambda 10; % 循环一致性权重太多会让风格迁移变弱太少会产生畸形图像 cycLoss mean(abs(recA - dlA), all) mean(abs(recB - dlB), all);forward调用中Outputs指定输出层名称如果生成器输出层名字不是 output这里要改成实际层名。均值绝对误差L1比均方误差L2模糊更少CycleGAN 原文用的就是 L1。lambda 取 10 是论文推荐值但实际使用时如果你的数据风格差异很大可以降到 5 以增强风格化强度如果图像出现大块畸变则升到 15 以加强内容保持。4.2 对抗损失与生成器/判别器交替更新对抗训练在 MATLAB 中需要手动交替更新生成器和判别器。先更新判别器让它对真实图像和生成图像都给出正确判断再更新生成器目标是让判别器对生成图像误判为真。两个网络用各自的 optimize 函数配合dlgradient更新。下面的代码片段展示了单步训练的核心逻辑% 判别器损失对真实图像输出趋近 1对生成图像输出趋近 0 gradientsD dlgradient(dLoss, dlnetD.Learnables); dlnetD dlupdate((W, g) W - lrD * g, dlnetD, gradientsD); % 生成器损失对抗部分 循环一致性 gLossAdv mean((dOutputFake - 1).^2, all); gLoss gLossAdv lambda * cycLoss; gradientsG dlgradient(gLoss, dlnetG.Learnables); dlnetG dlupdate((W, g) W - lrD * g, dlnetG, gradientsG);注意dlupdate是 MATLAB 提供的参数更新函数不需要手写 for 循环遍历层参数。学习率方面生成器和判别器可以使用相同学习率也可以判别器略高。这里有一个实用建议如果把生成器学习率设成 0.0002判别器设成 0.0001训练前期会更稳虽然略偏离原文参数但更容易收敛。4.3 训练进程的可视化与监控指标训练 GAN 最怕的不是 loss 不降而是 loss 降了但图像质量很差。我一般把每个 epoch 的生成图像保存下来同时打印三种 loss 的滑动平均生成器总损失、判别器损失、循环一致性损失。如果循环损失一直降但对抗损失停滞通常是判别器太强生成器学不到有效梯度。if mod(iter, 50) 0 figure(1); subplot(1,2,1); imshow(extractdata(fakeA)); title(B - A); subplot(1,2,2); imshow(extractdata(fakeB)); title(A - B); drawnow; end图中如果出现大面积噪点先检查生成器输出层激活函数是不是 tanh以及输入数据是否也做了 [-1,1] 归一化。常见误用是生成器输出层用了 relu 或直接线性激活导致值域不匹配。5. 运行结果怎么看质量验收、常见失败模式与三个调试技巧5.1 风格迁移的三个验收层次拿到训练结果后不要只看 loss 下降曲线要建立三个验收层次。第一层是查看单张图像是否引入了目标风格比如照片转梵高风格时是否有笔触、色彩是否偏暖第二层是内容是否保留主体物体边缘不应严重扭曲第三层是多样性同一张输入图在多次推理时结果应基本稳定。这个项目里的运行结果文件夹就应该包含训练完成后的生成器权重文件以及测试集上的批量转换效果图。数值指标方面常见做法是计算 FIDFréchet Inception Distance但 FID 对风格迁移并不完全适用它更偏向评价生成多样性。对于 CycleGAN 这类任务我更建议做人工盲评准备 10 张测试图让 5 个人选出风格更像目标域的结果用选择比例作为最终验收依据。5.2 三个高频失败模式训练不收敛或者生成图像有棋盘格伪影是 CycleGAN 最容易遇到的问题。棋盘格伪影多半来自 transposedConv解决办法是换成 upsample 加 convolution2dLayer 的组合或者在转置卷积后面加一个高斯模糊层。模式崩溃则表现为所有输入都生成相似图像这时需要调大判别器学习率或者检查是不是循环一致性损失权重过高压死了多样性。最后一个常见问题是生成器根本没有改变图像内容输出几乎等于输入。这通常是 identity loss 权重过高导致的CycleGAN 某些实现里 identity loss 的默认权重是 0.5如果你的数据域差异已经很小identity 项会抑制风格迁移。我实际用的时候经常直接把它关掉。5.3 用一个具体技巧验证训练是否到位最直接的验证方式是用一个固定的预训练生成器对 20 张测试图像做前向推理然后计算这 20 张结果的像素级方差。如果方差很小说明模型陷入模式 collapse。这个检测方法简单且不需要额外安装任何工具箱代码如下imgs imageDatastore(./testA); % 测试图片 results zeros(256,256,3,20,single); for i 1:20 img read(imgs); img (single(img) / 255 - 0.5) * 2; dlImg dlarray(img, SSCB); fake predict(dlnetG, dlImg); results(:,:,:,i) extractdata(fake); end varianceScore var(results, 0, 4); fprintf(平均像素方差: %.4f\n, mean(varianceScore, all));如果方差低于 0.01说明生成结果几乎相同模型大概率崩了。如果方差在 0.02 到 0.08 之间说明有一定的多样性再结合实际图像确认风格是否符合预期。这个技巧在整个训练过程中可以每 10 个 epoch 运行一次比单纯看 loss 曲线更直接。本文还有配套的精品资源点击获取