ARTICLE DETAIL

资讯详情

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

基于Matlab的生成对抗网络仿真:从数据预处理到自定义训练循环

基于Matlab的生成对抗网络仿真:从数据预处理到自定义训练循环 简介一份基于Matlab的生成对抗性网络GAN仿真资料包重点覆盖图像数据生成、生成器与判别器对抗训练等核心环节面向计算机、电子信息工程、数学等专业的学生可用于课程设计、期末大作业或毕业设计的参考资料。包内共45个文件包含43个m脚本负责模型定义、训练循环、可视化等、1个mat数据文件和1份说明文档压缩后大小约13.98MB。mat数据文件提供标准手写数字数据集说明文档梳理实现思路与使用流程方便在Matlab环境下快速复现并二次开发。已有249人学习下载适合需要结合源码理解GAN原理、搭建仿真实验的初学者和具有一定基础的开发者。1. 为什么用Matlab做GAN仿真而不是Python很多人拿到一个带源码的生成对抗性网络仿真项目第一反应是找Python和PyTorch。但对于需要与信号处理、通信仿真或现有Matlab代码对接的人Deep Learning Toolbox从R2019a开始已经支持dlnetwork自定义训练循环生成的图像可以直接imshow、打印和存入results目录不需要额外装GPU版本就能跑通小规模GAN。这个标题里的“源码数据说明文档”三件套是典型的课程设计或科研交付物结构。如果直接打开压缩包就跑main_train.m大概率会卡在维度不匹配、数据没归一化、不收敛这几个环节。我会围绕这个结构讲清楚网络层怎么搭、数据怎么喂、训练循环怎么写、说明文档里最该记录哪些参数以及遇到“仿真发散”时先改哪里。2. 从零搭一个GAN仿真工程源码目录与数据准备2.1 项目文件结构源码、数据、说明文档怎么放拿到一个名为“基于Matlab实现生成对抗性网络仿真源码数据说明文档.rar”的压缩包最常见的目录划分是GAN_Simulation/ ├── code/ │ ├── main_train.m │ ├── createGenerator.m │ ├── createDiscriminator.m │ ├── ganLoss.m │ └── projectAndReshapeLayer.m ├── data/ │ └── images.mat ├── results/ │ └── training_progress.gif └── README.md我一般会这样安排code只放脚本和函数data放原始数据和预处理后的mat文件results放训练过程中的图像和checkpointREADME.md是说明文档。这样做的原因是GAN训练会反复调整内容和实验参数如果不把结果单独放一个目录最后根本分不清哪张图是哪次实验产出的。说明文档不应该只写“怎么运行”至少要有三个部分环境要求、数据格式说明和参数表。环境要求里不要只写“Matlab 2023b”要写清楚需要Deep Learning Toolbox如果要用Parallel Computing Toolbox配合GPU训练也需要说明。数据格式说明要写出数据维度、归一化范围、训练集和验证集怎么划分。参数表则要列出latentDim、batchSize、learningRate等关键值否则换一台机器复现时对不上。2.1.1 数据契约先定形状再写网络GAN的网络定义和数据处理是强耦合的。生成器输出的图像尺寸必须和判别器输入尺寸一致而判别器输入的通道数又由数据决定。所以我的习惯是先写一段数据读取脚本打印出数据的真实维度再开始定义网络。下面是一段常见的数据加载代码。% 数据加载从data/images.mat读取并统一为double类型 load(fullfile(data, images.mat), images); % 原始数据可能是uint8 imgs im2double(images); % 归一化到[0,1]区间 % 如果数据是彩色图但想跑灰度GAN这里做单通道转换 if size(imgs, 3) 3 imgs rgb2gray(imgs); % 转为单通道灰度图 end imgs imgs(:, :, 1, :); % 确保维度是 H*W*1*N % 打印维度信息确认后面网络层能对齐 disp(size(imgs)); % 期望 [28 28 1 5000]这段代码的关键点有两个。第一im2double会把uint8的数据从0-255映射到0-1GAN的判别器一般都期望输入在0-1或-1到1之间直接输入uint8会让损失值本身很大梯度更新也变得不稳。第二rgb2gray之后必须保留(:, :, 1, :)这个四维切片因为Matlab的dlnetwork接受的数据格式是SSCB空间-空间-通道-批形状错了会直接在predict时报维度不匹配。2.2 在Matlab中加载和预处理数据集2.2.1 以MNIST为例的快速读取最省事的做法是直接用Matlab自带的digitTrain4DArrayData它返回5000张28x28的灰度图已经归一化到0-1非常适合验证GAN网络是否收敛。代码是这样写的。% 使用Matlab内置手写数字数据做快速验证 [trainImages, ~] digitTrain4DArrayData; % 只取图像分类标签在这里用不到 trainImages single(trainImages); % 转换为single精度加快GPU计算这里选MNIST而不是CIFAR-10的原因是生成对抗性网络的训练需要大量迭代MNIST的28x28分辨率很小CPU上也能在几分钟内看到生成效果。等网络结构调通后再换自己的高分辨率数据可以节省大量排错时间。如果手头只有文件夹里的图片也可以用imageDatastore读取但要注意imageDatastore返回的是cell数组需要手动cellfun转成四维数组不如load一个mat文件直接。2.2.2 数据归一化与批处理GAN对数据分布非常敏感。常见做法是把数据归一化到-1到1之间对应生成器最后一层用tanh激活函数这样生成器和判别器的数值范围是对称的。im2double得到的是0-1区间所以还需要再变换一下。% 将数据从[0,1]映射到[-1,1] imgs imgs * 2 - 1; % 按batch读取数据 numObs size(imgs, 4); shuffledIdx randperm(numObs); miniBatchIdx shuffledIdx(1:batchSize); XBatch imgs(:, :, :, miniBatchIdx); XBatch dlarray(XBatch, SSCB); % 封装成dlarray格式为SSCB到底用0-1还是-1到1取决于生成器最后一层激活函数。如果生成器最后是sigmoid就用0-1如果是tanh就用-1到1。混搭是导致训练发散最常见的原因之一后面第4章会再次遇到。dlarray的SSCB格式中第一个S是行第二个S是列C是通道B是batch这个顺序不能乱否则卷积层的Padding和Stride计算全都会错位。2.3 说明文档应该写什么2.3.1 复现步骤与运行环境我在交付说明文档时一定会把“从Matlab启动到看到生成图像”的命令按顺序抄进去。比如先写打开main_train.m再写按F5运行然后说明训练结束后在results目录下会生成什么文件。不要假设读者对Matlab的类定义和dlnetwork熟悉。环境部分要具体到工具箱版本例如“MATLAB R2023b Deep Learning Toolbox 23.2”因为dlnetwork的API在R2021a和R2023b之间有小幅变化。2.3.2 参数表示例参数名推荐值说明latentDim100随机噪声维度太小会模式崩溃batchSize128受显存限制小数据可降为64learningRate0.0002使用Adam时的常见设定beta10.5Adam的动量衰减0.5比默认0.9更稳maxEpochs50MNIST上一般30轮能看到清晰数字这张参数表直接拷贝进说明文档就行每一行都要写为什么。比如beta1设成0.5是DCGAN论文里被验证过的设定因为GAN训练中的梯度震荡大过大的历史动量会让更新方向失真。如果你的说明文档里只有命令没有参数表读者改了超参数后就没有参照物训练出问题也无从排查。3. 生成器与判别器的Matlab网络定义3.1 用dlnetwork定义生成器为什么不用trainNetwork因为GAN需要交替训练生成器和判别器两个网络共用一组数据却使用不同的损失函数trainNetwork只能自动训练单一损失函数模型。所以必须用dlnetwork配合dlgradient做自定义训练。生成器的作用是把一个latentDim维的随机噪声映射成28x28图像。3.1.1 转置卷积与上采样下面的生成器定义函数可以直接放在createGenerator.m里。核心思路是先用投影层把噪声向量拉长成特征图再用三个转置卷积逐层放大分辨率。function dlnetG createGenerator(latentDim, numFilters, imageSize) layers [ featureInputLayer(latentDim, Normalization, none, Name, noiseIn) projectAndReshapeLayer(latentDim, numFilters*8, imageSize/8, proj) transposedConv2dLayer(4, numFilters*4, Stride, 2, Cropping, same, Name, tconv1) batchNormalizationLayer(Name, bn1) reluLayer(Name, relu1) transposedConv2dLayer(4, numFilters*2, Stride, 2, Cropping, same, Name, tconv2) batchNormalizationLayer(Name, bn2) reluLayer(Name, relu2) transposedConv2dLayer(4, 1, Stride, 2, Cropping, same, Name, tconv3) tanhLayer(Name, genTanh) ]; dlnetG dlnetwork(layers); end这段代码里有两个必须解释的地方。第一projectAndReshapeLayer不是Matlab内置层它是官方DCGAN示例里的辅助自定义层作用是把latentDim * 1 * 1的向量reshape成(imageSize/8) x (imageSize/8) x (numFilters*8)的特征图。如果你从dlnetwork直接复制这段代码会报错需要在code目录下补一个同名类文件。第二生成器最后一层用的是tanh所以我们在第2章里把数据归一化到-1到1如果这里改成sigmoid前面的数据归一化也要跟着改回0-1。这是最容易忽略的“数据与网络契约”。3.1.2 自定义层的替代写法如果不想引入自定义类可以把reshape逻辑放到训练循环里。常见做法是让生成器第一层用fullyConnectedLayer(numFilters*8*(imageSize/8)^2)然后在训练循环里用reshape把输出变成特征图再传给后续卷积层。但这样的话dlnetwork的层图里会多出一个functionLayer代码反而更绕。所以我一般保留projectAndReshapeLayer这个类文件一个类文件也就二三十行维护成本不高。3.2 判别器网络设计3.2.1 卷积层与LeakyReLU判别器是二分类器输入图像输出一个logit。它的结构相对简单卷积层加LeakyReLU最后一层是fullyConnectedLayer(1)没有sigmoid因为损失函数里会自己计算sigmoid交叉熵。以下是一个标准的判别器定义。function dlnetD createDiscriminator(imageSize, numFilters) layers [ imageInputLayer([imageSize imageSize 1], Normalization, none, Name, imgIn) convolution2dLayer(4, numFilters, Stride, 2, Padding, same, Name, conv1) leakyReluLayer(0.2, Name, lrelu1) convolution2dLayer(4, numFilters*2, Stride, 2, Padding, same, Name, conv2) batchNormalizationLayer(Name, bn2) leakyReluLayer(0.2, Name, lrelu2) convolution2dLayer(4, numFilters*4, Stride, 2, Padding, same, Name, conv3) batchNormalizationLayer(Name, bn3) leakyReluLayer(0.2, Name, lrelu3) fullyConnectedLayer(1, Name, fcOut) ]; dlnetD dlnetwork(layers); end判别器的卷积层步长都是2相当于每次把图像分辨率减半28 - 14 - 7 - 4。Padding设为same保证输出尺寸是输入的一半而不是大小更乱。LeakyReLU的斜率设成0.2这是DCGAN的设定。为什么要用LeakyReLU而不是ReLU因为判别器如果输出大量负值ReLU会把它们全部截断成0梯度消失生成器就没法从判别器那里拿到有用信息。LeakyReLU允许小的负梯度流过训练更稳定。3.3 损失函数与梯度计算GAN的损失函数是两个交叉熵的组合。判别器要区分真图和假图生成器要让假图被判别为真。Matlab里用dlarray和sigmoid函数可以直接写不需要自己计算交叉熵。function [lossG, lossD] ganLoss(dlYPredReal, dlYPredFake) % 判别器损失真图标签为1假图标签为0 lossD -mean(log(sigmoid(dlYPredReal) eps)) ... -mean(log(1 - sigmoid(dlYPredFake) eps)); % 生成器损失希望假图被判定为真 lossG -mean(log(sigmoid(dlYPredFake) eps)); end这里eps的作用是防止log(0)出现NaN。sigmoid函数接收dlarray并返回dlarray整个计算图是自动可微的。在网络定义时没有给判别器加sigmoid就是为了在这个函数里统一计算避免中间层数值过大导致饱和。实际训练中判别器如果比生成器强太多dlYPredReal会很快接近很大的正数损失会趋于0这时需要看第4章里的“训练失衡”处理。4. 训练循环从手写数字到稳定收敛4.1 训练超参数设置4.1.1 用Adam还是SGDGAN训练没有银弹常见的选择是Adam优化器而不是SGD。Matlab的adamupdate函数支持trailingAvg和trailingAvgSq状态训练循环里需要自己维护这两个状态。以下是一组经过验证的初始参数。learningRate 0.0002; beta1 0.5; beta2 0.999; numEpochs 30; batchSize 128; latentDim 100; imageSize 28;为什么要用beta10.5默认的Adam动量是0.9会让梯度更新方向平滑但GAN的损失曲面是非凸且震荡大的过大的历史梯度会导致更新方向被旧梯度主导生成器迟迟学不会新的分布。0.5相当于缩短记忆让模型更“健忘”更适应快速变化的数据分布。learningRate从0.0002开始不要一上来就用0.01否则生成器会出现棋盘伪影甚至直接发散。4.2 自定义训练循环4.2.1 前向传播与损失计算自定义训练循环的骨架大致如下。注意每个epoch都要打乱数据顺序而且adamupdate的迭代计数不能重置。numIterPerEpoch floor(numObs / batchSize); monitor trainingProgressMonitor; % 训练进度可视化 avgG []; avgSqG []; avgD []; avgSqD []; iteration 0; for epoch 1:numEpochs shuffledIdx randperm(numObs); for iter 1:numIterPerEpoch idx shuffledIdx( (iter-1)*batchSize1 : iter*batchSize ); XReal imgs(:, :, :, idx); XReal dlarray(XReal, SSCB); % 生成随机噪声维度为 latentDim x batchSize Z randn(latentDim, batchSize, single); Z dlarray(Z, CB); % 生成器前向注意这里用predict而不是forward XFake predict(dlnetG, Z); % 判别器前向对真图和假图分别做 YReal forward(dlnetD, XReal); YFake forward(dlnetD, XFake); % 计算损失 [lossG, lossD] ganLoss(YReal, YFake); % 分别计算梯度 gradG dlgradient(lossG, dlnetG.Learnables); gradD dlgradient(lossD, dlnetD.Learnables); % 更新参数 iteration iteration 1; [dlnetG, avgG, avgSqG] adamupdate(dlnetG, gradG, avgG, avgSqG, iteration, learningRate, beta1, beta2); [dlnetD, avgD, avgSqD] adamupdate(dlnetD, gradD, avgD, avgSqD, iteration, learningRate, beta1, beta2); end end这段代码最关键的是forward和predict的使用。在判别器训练时forward会保留中间结果用于梯度计算在生成器生成图像时predict不保留前向的中间值节省内存。如果对生成器也用forward训练时会因为计算图占用过多显存而变慢甚至OOM。另外dlgradient必须在dlarray支持的自动微分上下文中调用所以这里没有把前向过程封装成函数而是全部写在训练循环里。gradG和gradD的结构分别与dlnetG.Learnables和dlnetD.Learnables一致adamupdate会按字段逐个更新。4.2.2 梯度更新与可视化adamupdate需要传入迭代次数iteration这个值要从1开始递增不能每个epoch重置。如果重置了学习率退化和动量状态都会错乱。可视化时我一般每50次迭代用imshow(extractdata(XFake(:,:,:,1)))显示一张生成图同时输出当前的lossD和lossG。注意extractdata会把dlarray还原成普通数组否则imshow会因为输入类型不对报错。训练到第10轮左右如果生成图仍然是纯噪声那大概率不是迭代次数不够而是某个参数设置完全错了。4.3 常见训练问题与对策4.3.1 模式崩溃与不收敛模式崩溃的表现是生成图里只有少数几种数字重复出现。常见对策有三个降低学习率到0.0001不要增加生成器的更新频率那样会破坏两个网络之间的平衡更有效的是给判别器加标签平滑即真图标签从1改成0.9。代码层面改动很小只改ganLoss.m一行。% 在ganLoss.m中使用标签平滑注意第一项0.9 labelSmooth 0.9; lossD -mean(labelSmooth * log(sigmoid(dlYPredReal) eps)) ... -mean(log(1 - sigmoid(dlYPredFake) eps));标签平滑的作用是让判别器的输出不要过于自信避免梯度消失。实际使用中我见过只靠这个改动就把训练从发散拉回来的情况。另一个很常见的原因是数据没有打乱每个epoch内按固定顺序取batch判别器会记住数据顺序而不是学习数据分布改成randperm后就恢复正常。4.3.2 训练状态对比表在说明文档里我建议放一张训练状态判断表方便排错时快速定位。现象最可能原因优先操作lossD一直降lossG完全不动判别器太强降低判别器学习率或加标签平滑lossG震荡生成图全是噪点learningRate过大降到0.0001生成图只有同一种数字模式崩溃增大latentDim或调节beta1loss为NaN数据里有NaN或学习率过大检查数据和eps4.3.3 如何用说明文档记录实验每次跑完一个实验我会在说明文档里追加一段结果记录格式是日期、latentDim、batchSize、学习率、损失曲线截图、生成的最终图像。不要只贴一张好看的图要贴训练过程中的失败样本这对后来人排错价值更大。标题里既然有“说明文档”就应该把参数实验当成正式内容来写而不是最后随便补两句话。5. 用说明文档做出来的三个调试技巧5.1 用predict范围检查中间张量训练发散时我第一件事是打印生成图像的范围。如果max(abs(XFake(:)))远大于1说明tanh输出的坐标已经超出合理范围。可以用下面代码快速检查。fakeData extractdata(XFake); fprintf(fake pixel range: [%f, %f]\n, min(fakeData(:)), max(fakeData(:)));正常训练中这个范围应该靠近-1到1。如果出现NaN或者长时间卡在-1或1附近问题一般出在batchNormalizationLayer或者学习率过大。5.2 用输出图像网格快速判断GAN状态不要只看一张生成图。我习惯把16张图拼成网格每50次迭代保存一帧这样可以通过动画看到模式崩溃的过程。Matlab的montage函数可以直接用四维数组生成网格。imgGrid imresize(extractdata(XFake), [64 64]); % 放大便于观察 montage(imgGrid, Size, [4 4]); drawnow;如果动画里前20轮图像形态一直在变第30轮突然变成同一种数字说明模式崩溃发生在前20轮到30轮之间这时候可以回头检查beta1和标签平滑的改动效果。5.3 保存checkpoint与恢复训练训练到一半中断是常事。需要在说明文档里写明checkpoint的存放位置。保存方式用save即可但要注意同时保存优化器状态和迭代序号。save(fullfile(results, checkpoint.mat), dlnetG, dlnetD, avgG, avgSqG, avgD, avgSqD, iteration);恢复训练时直接load回来继续跑adamupdate的动量状态不会丢失。以上三个技巧对任何基于dlnetwork的GAN变体都适用下次看到“源码数据说明文档”压缩包时先改的是数据契约再改的是标签平滑最后动的才是网络结构。本文还有配套的精品资源点击获取
返回列表