ARTICLE DETAIL

资讯详情

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

MATLAB实现生成对抗网络(GAN):从搭建到训练的全流程代码解析

MATLAB实现生成对抗网络(GAN):从搭建到训练的全流程代码解析 简介面向深度学习初学者与GAN研究者的MATLAB实现生成对抗网络可直接运行代码包解决了从零搭建GAN训练流程的入门门槛问题。压缩包含14个文件其中13个.M脚本覆盖了网络初始化、前向传播、反向传播、梯度更新、激活函数及交叉熵损失等核心模块另附一张lena.bmp测试图用于直观验证生成效果。代码结构清晰包含生成器与判别器定义、交替训练循环、参数配置与可视化接口配合代码说明可快速理解生成对抗网络的运行机制。已有6516人学习适合希望在MATLAB环境下动手实践GAN、进行数据增强或图像合成实验的读者。 如果你在网上搜“GAN MATLAB代码”大概率会看到两类结果一类是很久之前的老代码函数名和现在的版本对不上跑之前要改一堆报错另一类是某个商业工具箱封装好的示例看起来能出图但想改网络结构时完全无从下手。我前段时间因为项目需要在MATLAB环境下完整跑通一个生成对抗网络把这个问题从头啃了一遍最后整理出一份能在R2022a及以上版本直接运行、结构清楚、方便二次修改的代码。这篇文章就当是踩坑记录加实现笔记把每一步为什么这么做讲明白也把那些最容易卡住的地方一并列出来。这篇文章适合三类人看一是课程或课题把环境限定在MATLAB又偏偏要做GAN相关实验的同学二是已经会用Python跑GAN想在MATLAB里做对照组或者复现结果的研究人员三是只想快速拿到一段能跑通的代码再逐步改造成自己网络结构的初学者。如果你属于其中之一按照下面的步骤走基本能避开大部分坑。1. 为什么选MATLAB做GAN适用场景与工具箱检查先说结论MATLAB跑GAN不是主流选择但在某些场景下确实有它不可替代的优势。一个最典型的场景是课程设计或者论文复现。我见过不少控制、通信、图像处理方向的课题整个实验框架都搭在MATLAB里数据预处理、指标计算、图表导出全是一套流程。这时候如果你为了一个GAN模块单独去搭Python环境不仅要额外管理一堆依赖库还要处理两种环境之间的数据传递。与其这样不如直接在MATLAB里把GAN实现掉整个实验链路保持统一。另一个优势是可视化。MATLAB里imshow、montage这些函数对图像类结果展示非常友好训练过程中随时可以看一眼生成效果。相比之下用Python写matplotlib虽然也不难但总归要多写几行代码。MATLAB的调试体验和矩阵操作习惯对很多工程背景的人来说也更亲切。不过MATLAB版GAN的劣势也很明显。最大的问题就是社区生态远不如Python活跃很多经典模型的官方实现都是PyTorch/TensorFlow版本MATLAB里经常需要自己从论文公式一点点翻译过来。所以如果你完全没接触过GAN的底层原理我建议还是先在Python里跑通一个Demo理解了输入输出关系以后再迁移到MATLAB会轻松很多。动手之前先检查环境是否满足条件。我用的是MATLAB R2022a完整代码里用到了以下几个关键能力Deep Learning Toolboxdlnetwork、trainNetwork之外的底层训练API都在这个工具箱里dlarray和dlgradient这是实现自定义训练循环的核心利用自动微分计算梯度adamupdate内置的Adam优化器更新函数省去手写动量计算的麻烦如果你不确定自己有没有这些工具箱在命令行输入ver查看已安装的工具箱列表或者运行下面这行代码直接检查[installed, toolboxList] ismember(... {Deep Learning Toolbox}, {matlab.addons.installedAddons().Name});如果没有Deep Learning Toolbox后面代码基本跑不起来。请通过学校或单位的正版授权渠道补充工具箱MATHWORKS官方支持试用版申请这块自己解决我不展开。2. 网络搭建生成器与判别器的层级设计这里用MNIST手写数字作为实验数据图像尺寸28×28单通道这也是GAN最经典的入门场景。网络设计没有直接照搬原始GAN论文里的全连接版本而是参考了DCGAN的思路把卷积和转置卷积用上。原因后面会讲。2.1 生成器从100维噪声到28×28图像生成器的作用是接收一个随机噪声向量输出一张逼真的图片。噪声维度定为100这是从原始GAN论文沿用的惯例100维足够提供生成多样性又不是特别大。经过一个全连接层后先映射到7×7×128的特征图然后连续做两次转置卷积每次把空间尺寸放大一倍最终从7×7升到28×28。function dlnetGen createGenerator() layers [ featureInputLayer(100, Normalization, none, Name, noise) fullyConnectedLayer(7*7*128, Name, fc1) reluLayer(Name, relu1) functionLayer((X) reshape(X, 7, 7, 128, []), ... Formatted, false, Name, reshape7x7) transposedConv2dLayer(5, 64, Stride, 2, Cropping, same, Name, tconv1) reluLayer(Name, relu2) transposedConv2dLayer(5, 1, Stride, 2, Cropping, same, Name, tconv2) tanhLayer(Name, tanh_out) ]; dlnetGen dlnetwork(layers); end为什么最后一层用tanh而不是sigmoid因为tanh输出范围是[-1, 1]而MNIST数据在输入网络前也需要归一化到[-1, 1]。生成器和真实数据的数值范围保持一致判别器就不会单凭输出区间就轻松分辨真假。这是很多新手第一次写GAN容易忽略的细节数据归一化和激活函数的选择必须配套。中间为什么拆了一步reshapetransposedConv2dLayer需要输入是H×W×C×N的四维格式而全连接层输出是一维向量所以必须先把向量reshape成7×7×128的特征图。这里用functionLayer写了一个匿名函数完成reshape这个写法在R2022a里实测可行。如果你的版本提示functionLayer有问题也可以在训练循环里对predict的输出手动做reshape效果一样。2.2 判别器卷积下采样与真假二分类判别器的任务相对简单输入一张图输出一个0到1之间的分数越接近1表示越像真实图片。结构上采用卷积下采样逐步提取高层特征最后通过全连接层输出单值。function dlnetDis createDiscriminator() layers [ imageInputLayer([28 28 1], Normalization, none, Name, img) convolution2dLayer(5, 16, Stride, 2, Padding, 2, Name, conv1) leakyReluLayer(0.2, Name, lrelu1) convolution2dLayer(5, 32, Stride, 2, Padding, 2, Name, conv2) leakyReluLayer(0.2, Name, lrelu2) fullyConnectedLayer(1, Name, fc_score) sigmoidLayer(Name, sigmoid_out) ]; dlnetDis dlnetwork(layers); end这里有个设计细节值得单独说明原始GAN论文里判别器用的是ReLU但DCGAN之后的主流实践都换成了LeakyReLU。原因是ReLU在负区间的梯度恒为0当判别器发现输入是假图时某些神经元的输出会变成负数反向传播时梯度直接被截断参数得不到更新导致判别器“躺平”。LeakyReLU在负区间保留了一个0.2的小斜率保证梯度始终能流动。这个小改动对训练稳定性有非常直接的影响。判别器设计的另一个重点是步幅卷积代替池化。每个卷积层的Stride都设为2相当于每经过一层特征图尺寸缩小一半。这样做的好处是下采样过程可以通过卷积参数学习而不是像池化那样直接丢弃信息。28×28经过两次步幅2卷积变成7×7最后接入全连接层时参数数量也合理。2.3 训练循环损失函数、梯度更新与可视化监控网络结构搭好之后训练才是GAN真正考验人的地方。GAN的训练本质是一个二人零和博弈判别器努力分辨真假生成器努力骗过判别器。两者相互对抗、共同进化。如果其中一个太强另一个就学不到东西。先看损失函数。标准GAN采用二分类交叉熵。判别器的目标是最小化下面这个值[ L_D -\frac{1}{m}\sum_{i1}^{m} \left[ \log D(x^{(i)}) \log(1 - D(G(z^{(i)}))) \right] ]生成器的目标则是最大化判别器的出错率等价于最小化[ L_G -\frac{1}{m}\sum_{i1}^{m} \log D(G(z^{(i)})) ]在MATLAB自定义训练循环里需要把这两个损失写进一个函数并在dlfeval中调用这样dlgradient才能自动求梯度。核心代码如下function [lossGen, lossDis, gradsGen, gradsDis] modelGradients(... dlnetGen, dlnetDis, XReal, dlZ) % 生成假图 XFake predict(dlnetGen, dlZ); XFake reshape(XFake, 28, 28, 1, []); % 判别器对真实图与假图的评分 YReal forward(dlnetDis, XReal); YFake forward(dlnetDis, XFake); % 判别器损失 lossDis -mean(log(YReal eps) log(1 - YFake eps)); % 生成器损失 YFakeGen forward(dlnetDis, XFake); lossGen -mean(log(YFakeGen eps)); % 自动求梯度 gradsGen dlgradient(lossGen, dlnetGen.Learnables); gradsDis dlgradient(lossDis, dlnetDis.Learnables); end训练主循环里对两套网络分别用Adam优化器更新参数。Adam是一个自带动量与自适应学习率的优化器对GAN这种非凸博弈问题有较好的稳定性。值得注意的是我把学习率设成了2e-4而不是常用的1e-3。原因是GAN训练中学习率过大容易导致判别器损失迅速下降到接近0生成器后续完全学不到梯度。2e-4是DCGAN论文里经过调参验证的经验值实测在MNIST上非常稳定。numIterations 5000; batchSize 128; learningRate 2e-4; beta1 0.5; trailingAvgGen []; trailingAvgSqGen []; trailingAvgDis []; trailingAvgSqDis []; for iter 1:numIterations % 从真实数据中随机采样一个batch idx randi(size(XTrain, 4), batchSize, 1); XReal XTrain(:, :, :, idx); XReal dlarray(single(XReal), SSCB); % 生成随机噪声 Z randn(100, batchSize, single); dlZ dlarray(Z, CB); % 计算损失和梯度 [lossGen, lossDis, gradsGen, gradsDis] dlfeval(... modelGradients, dlnetGen, dlnetDis, XReal, dlZ); % Adam更新生成器 [dlnetGen, trailingAvgGen, trailingAvgSqGen] adamupdate(... dlnetGen, gradsGen, trailingAvgGen, trailingAvgSqGen, iter, ... learningRate, beta1); % Adam更新判别器 [dlnetDis, trailingAvgDis, trailingAvgSqDis] adamupdate(... dlnetDis, gradsDis, trailingAvgDis, trailingAvgSqDis, iter, ... learningRate, beta1); % 每100轮可视化一次 if mod(iter, 100) 0 ZShow dlarray(randn(100, 16, single), CB); XShow predict(dlnetGen, ZShow); XShow reshape(XShow, 28, 28, 1, []); imshow(imtile(extractdata(XShow), ThumbnailSize, [28 28])); title(sprintf(Iter %d, G Loss %.4f, D Loss %.4f, ... iter, extractdata(lossGen), extractdata(lossDis))); drawnow; end end有两个细节必须提一下。第一dlgradient不能直接在普通脚本里调用必须被包在dlfeval中否则会报“Must be called within a function”的错误。第二生成器更新的梯度方向看的是YFakeGen也就是把假图重新送入判别器目标是让判别器给出非常接近1的分数。很多初学者会误以为生成器直接用判别器上一轮的假图评分即可但那样梯度信息与当前生成器参数已经产生了脱节需要在更新前重新forward一次。可视化监控这部分我的建议是每训练一段时间就看看生成图不要只盯着损失曲线。GAN的损失曲线下降并不等于生成质量变好很多时候两者是此消彼长的。生成图片的实际观感才是最直观的判断依据。3. 数据准备MNIST读取与归一化处理MNIST数据集在MATLAB里没有内置需要自己准备。我使用的方式是去网上找一个已经转成.mat格式的MNIST版本这类文件通常包含trainX和trainY两个变量。如果你的数据来源是官方IDX二进制格式需要先用脚本转成MATLAB矩阵。这一步虽然只是在开头执行一次但对后面的训练影响很大。读取之后要做两个处理一是归一化到[-1, 1]。原始MNIST像素值是0到255之间的整数如果不处理直接送入生成器对比会产生问题。因为生成器最后一层是tanh输出落在[-1, 1]如果真实数据是[0, 255]两者的分布中心完全不同判别器只要看像素均值就能轻松判断真假生成器无论怎么优化都很难学到正确的映射。归一化代码很简单XTrain double(trainX) / 127.5 - 1;除以127.5再减1正好把[0, 255]映射到[-1, 1]。上学期的朋友可能已经注意到这里没有除以255再乘2减1效果是等价的但写在一起更简洁也不容易出错。二是调整维度顺序。MATLAB深度学习工具箱默认的数据格式是H×W×C×N即高、宽、通道、样本数。MNIST的.mat文件如果存储维度和这个不一致需要在送入网络之前用permute转一下。比如原始的trainX是N×784的矩阵需要先reshape成28×28再转置XTrain reshape(XTrain, 28, 28, 1, []);注意这个转置很关键。MNIST的常见存储格式是每行一个样本、每列一个像素如果不转置reshape出来的图像会是倒置的。我在这里踩过一次坑当时生成的图像全部是旋转90度的数字排查了半天才发现是数据排列问题。数据准备好以后训练过程中直接按随机索引抽取batch即可idx randi(size(XTrain, 4), batchSize, 1); XReal XTrain(:, :, :, idx); XReal dlarray(single(XReal), SSCB);dlarray的第二个参数SSCB表示这个数组的四个维度分别是空间(S)、空间(S)、通道(C)、批(B)。这个标记看起来很绕但对于dlgradient正确计算梯度非常重要。标记错了网络的前向传播不会报错但梯度方向会出现隐性问题。4. 实测结果与高频报错排查用上面这套配置在普通CPU笔记本上训练5000轮大约需要40到60分钟。前1000轮基本看不出形状全是灰蒙蒙的噪点到2000轮左右开始出现明暗分界隐约能看出笔画的痕迹4000轮以后数字轮廓变得比较清晰部分数字比如0和1已经比较像样了。如果电脑有NVIDIA GPU并且安装了 Parallel Computing ToolboxMATLAB会自动调用GPU加速训练时间能缩短到10分钟以内。下面把我在调试过程中遇到的高频问题逐一列出来这些问题是网上问得最多的也是初学者最容易卡住的。4.1 维度不匹配报错最常见的报错信息长这样Error using dlarray/reshape Number of elements must not change.这通常发生在生成器输出的reshape步骤。fullyConnectedLayer(7*7*128)的输出元素总数是6272但如果你在createGenerator里写的transposedConv2dLayer输入通道数不是128reshape后的元素总数就对不上。检查方法和解决思路很简单计算一下每一层的输出尺寸是否按预期变化尤其注意Padding, same和Stride, 2组合时偶数尺寸输入经过转置卷积会得到正好两倍尺寸奇数尺寸会有出入。4.2dlarray格式标签出错Error using dlarray/forward Input data must have trailing singleton dimensions.这类问题基本都是dlarray的格式标签写错了。生成器输入的噪声应该是CB图像数据应该是SSCB。如果忘记给数据包装成dlarray或者标签和实际维度顺序不一致都有可能触发类似报错。我的习惯是在每个网络的入口处先disp(size(X))打印一下输入尺寸和标签确认无误再往下写。4.3 判别器损失直接降到0训练刚开始几百轮判别器的损失就快速掉到接近0生成器的损失反而一直升高。这说明判别器太强生成器发出来的噪音一眼就能被识破。碰到这种情况有几个常用对策把判别器的学习率调低比如从2e-4降到1e-4给判别器加Dropout层降低它的过拟合能力换用更浅的判别器结构比如把第二个卷积层的通道数从32降到16修改训练节奏每更新两次生成器再更新一次判别器反过来如果生成器损失迅速归零而判别器损失一直很高那就是生成器太强要反向操作。4.4 生成图像全是噪声或者全黑全白这种情况要分两种可能。如果生成图像是纯噪声大概率是训练根本没收敛数值震荡导致生成器输出不稳定优先检查学习率是否过大。如果生成图像是全黑或全白需要检查数据归一化是否正确以及imshow显示时是否把[-1, 1]范围的数据直接显示成了全黑。imshow默认把输入数值按0到1范围解释小于0的值会被当成0处理。显示之前用extractdata取出数据后先(X 1)/2转换回[0, 1]范围imshow((extractdata(XShow) 1) / 2);5. 从基础版到高质量生成改进路线参考如果这份基础代码已经顺利跑通下一步可以尝试几个经典的改进方向每一步都对应GAN发展历史上一个重要的突破点。5.1 引入Batch Normalization在生成器和判别器的卷积层后各加一个batchNormalizationLayer这是DCGAN的典型改动。BatchNorm能把每层输入拉回到一个较为稳定的分布解决GAN训练中常见的内部协变量偏移问题。实测加上之后训练的稳定性明显提升对学习率的敏感度也降低了。注意判别器的输入层和生成器的输出层不要加BatchNorm否则会引入不必要的随机性。5.2 用WGAN-GP代替标准交叉熵损失标准GAN的交叉熵损失在判别器训练得过于充分时容易产生梯度消失。WGAN把损失函数换成了Wasserstein距离从根上改善了这个现象。简单说WGAN的判别器critic不再输出概率值而是输出一个实数值通过限制critic的Lipschitz约束来保证训练的平稳性。MATLAB里改起来不算太复杂主要是把sigmoidLayer去掉损失函数改为lossDis mean(YFake) - mean(YReal); lossGen -mean(YFakeGen);同时要对critic的权重做梯度惩罚gradient penalty这部分代码稍长但稳定效果立竿见影。5.3 条件生成从数字生成到指定数字生成如果你想让生成器能指定生成某个数字就需要引入条件GANConditional GAN的思想。做法是在生成器输入时把100维噪声和标签的embedding拼接在一起判别器输入时也把标签信息以某种方式叠加进去。这样生成器就能学会按类别生成图像。这个扩展方向在上一篇代码基础上改动最小但能让你对“生成对抗网络还能做什么”有更直观的理解。我个人在实际操作中的体会是跑通基础版GAN只是第一步真正理解GAN的博弈原理必须亲手改结构、调损失、看结果变化。这套MATLAB代码的价值就在于结构足够清晰每一部分都能独立替换适合用来做各种对比实验。如果只是照着一份代码跑通就结束最多只能收获一个“我跑过GAN”的结论远不如花一个下午把BatchNorm加进去、观察训练曲线的变化来得有用。希望你拿到这份代码之后不只是复制运行而是每一行都亲手敲一遍配合本文把损失函数和梯度流理解透。本文还有配套的精品资源点击获取
返回列表