尧图精选

用MATLAB实现GAN生成手写数字:从零搭建训练循环与避坑指南

🕒 发布时间:2026/10/1 23:29:39 📁 来源:尧图网络
简介压缩包提供了一份基于MATLAB的生成对抗网络GAN实现针对MNIST手写数字数据集展开实验适合希望借助MATLAB快速上手GAN的深度学习初学者、相关课程学习者及对生成模型感兴趣的开发者。资源包括2个文件共14.03MB其中mat文件为MNIST数据集的MATLAB格式存储m文件为主脚本可完成从数据加载、网络构建到训练生成的全流程操作。通过生成器与判别器的交替训练读者可以直观看到模型如何从随机噪声逐步生成逼真数字理解GAN的博弈机制与损失函数作用。已有313人学习下载说明该实现具备一定的参考价值。代码结构简洁注释清晰便于对照论文原理进行阅读和修改同时演示了MATLAB深度学习工具箱中网络层定义、训练循环及数据预处理等常用操作。压缩包内附完整数据下载后可直接运行观察效果省去单独准备数据集的麻烦是入门GAN和MATLAB深度学习实践的一份实用参考。1. 一个 rar 包里的 GAN为什么用 MATLAB 训手写数字生成器反而是最快路径在代码分享站看到GAN.rar这种压缩包先别急着划走。标题里 MINIST 是 MNIST 的手滑拼写GaN 是氮化镓的缩写、跟生成网络没半点关系nestsw 基本就是归档时夹带的尾巴词——拆开看这套资源讲的是「在 MATLAB 里用 GAN 生成手写数字」。一个反直觉的事实是这个任务用 MATLAB 跑恰恰是最省事的路径。没有 Python 依赖不用配 CUDA一台只有 CPU 的笔记本也能在几分钟内跑完 50 个 epoch。适合谁想在不换环境的前提下把生成模型跑通的研究生、做图像处理的工程师以及所有好奇对抗训练在 MATLAB 里到底长什么样的人。这篇文章不假装见过包里源码按这个方向给你拆一套能复现、能改写的落地实现顺带把最容易翻车的几个坑标出来。2. 对抗训练的最小单元损失函数怎么写以及为什么 trainNetwork 用不了你从 rar 里期望得到的无非一个能跑的脚本和一组不让损失跑飞的超参数。这两样东西都建立在一个问题上GAN 到底在更新什么。2.1 生成器和判别器各管一件事从 100 维噪声到 784 像素的输出链路生成器 G 把 100 维噪声向量 z 映射成 784 个像素也就是一张 28×28 的灰度图判别器 D 把一张 28×28 图像压成一个概率。真实图像来自 MNIST假图像来自 G 的输出重新排列成图。D 想让真实图像的得分接近 1、生成图像接近 0G 想让自己生成的图像被 D 判成 1。目标函数写成min_G max_D E[log D(x)] E[log(1 - D(G(z)))]看起来是两条网络在打架其实是一个 minimax 博弈。等号右边第一项是真实图像 x 的期望第二项是噪声 z 经 G 生成后的期望。D 要最大化整个式子G 要最小化它。G 的训练信号完全来自 D 的梯度——如果 D 一眼识破所有假图G 的梯度就会很小甚至消失如果 D 太弱G 又没有压力去生成更真实的图。这就是 GAN 训练的平衡木。z的维度取 100 不是随手定的。DCGAN 原论文把 100 维标准正态噪声作为默认配置这个规模对 MNIST 足够表达数字的类内变化。设得太低比如 10 维容易模式崩塌生成出来的数字永远就那几种设得太高比如 1000 维训练会变慢但对 MNIST 这种小图收益不明显。我一般固定zDim 100只在生成多样性出问题时才往上调到 128。生成器输出层用tanh把像素压到 [-1, 1]这是所有经典 GAN 实现的共同选择。判别器本质是一个二分类器它的输出经过sigmoid变成 0 到 1 的概率。注意MATLAB 里 I go 把判别器的最后一层设置成fullyConnectedLayer(1)然后在损失函数里手动调sigmoid而不是在层图里加一个sigmoidLayer。原因是这个层的可用性在不同 MATLAB 版本里有差异手动调用sigmoid(dlX)在 R2020a 之后都能跑。2.2 自定义训练循环才是 GAN 的主场dlgradient 与为什么 trainNetwork 用不了很多人习惯拿到数据就trainNetwork到 GAN 这里必须改掉这个惯性。trainNetwork是为监督学习设计的一个网络、一组标签、一个损失。GAN 需要两个网络交替更新生成器的损失还要穿过判别器反向传播并且没有「正确答案」标签只有「真/假」这种对抗信号。虽然理论上能把两个网络拼成一个多输出大网络再用trainNetwork但你需要写自定义 output layer 和 loss layer调试成本远高于直接写循环。MATLAB 官方示例在 GAN 这个任务上同样走自定义训练循环。训练循环的主角是dlgradient和dlfeval。dlgradient只能在函数内部调用而且必须由dlfeval触发否则直接报错。损失函数的离散化写法如下function [lossG, lossD] ganLoss(dReal, dFake) epsVal 1e-7; lossD -mean(log(dReal epsVal)) - mean(log(1 - dFake epsVal)); lossG -mean(log(dFake epsVal)); enddReal是判别器对真实图像输出的概率dFake是对生成图像输出的概率。判别器损失由两项组成真实图像预测越接近 1第一项越小生成图像预测越接近 0第二项越小。生成器损失只有一项生成图像被 D 判得越接近 1损失越小。mean是对一个 batch 内所有样本取平均这是标准做法。epsVal 1e-7是防 NaN 的关键。如果 D 完美碾压 GdFake会非常接近 0log(1 - dFake)里的1 - dFake接近 1 没问题但log(dFake)会变成负无穷。给dFake加一个极小值梯度爆炸的几率会小很多。这里的反向传播是黑匣子但你可以掌控梯度流向。G 的梯度要穿过 D 的网络计算所以dlgradient(lossG, dlnetG.Learnables)必须在 lossG 和 dlnetG 之间存在可微路径而这条路径是通过forward(dlnetD, dlXG)搭起来的。把 G 的输出直接当成图像喂给 D损失才能传回去。这也是为什么生成器输出的维度必须和判别器输入对得上下一章就处理这个对接。3. 把 MNIST 数据送进 dlnetwork归一化、维度约定与两种加载方式数据是第一个容易卡住的地方。MNIST 在 MATLAB 里有不止一种接法选错了会在数据预处理上浪费大量时间。3.1 先拿到数据MATLAB 内置数据集一行加载IDX 解析当备胎最省事的方式是用 Deep Learning Toolbox 附带的数据集加载函数它直接返回 4 维数组不需要从网上下载任何文件% 方案 AMATLAB 自带 MNIST返回训练图像和标签 [xTrain, tTrain] digitTrain4DArrayData; % 统一转成 single避免后续和 randn 生成的数据精度不一致 xTrain single(xTrain); % 像素从 [0, 255] 映射到 [-1, 1]和生成器 tanh 输出范围匹配 xTrain (xTrain - 127.5) / 127.5; % MNIST 只有 6 万张但 single 化后仍占约 188MB内存紧就取前 2 万张 % xTrain xTrain(:, :, :, 1:20000);如果你的 MATLAB 版本里没有digitTrain4DArrayData或者 rar 包自带的是原始 IDX 格式文件写一个解析函数也不难。IDX 文件头是四个 32 位大端整数分别表示魔数、样本数、行数、列数后面是像素字节function im loadMNISTImages(fn) fid fopen(fn, rb); magic fread(fid, 1, uint32, 0, ieee-be); assert(magic 2051, 不是 MNIST 图像文件); numImages fread(fid, 1, uint32, 0, ieee-be); rows fread(fid, 1, uint32, 0, ieee-be); cols fread(fid, 1, uint32, 0, ieee-be); raw fread(fid, inf, uint8uint8); fclose(fid); im reshape(raw, rows, cols, 1, numImages); im permute(im, [2 1 3 4]); % IDX 是行优先存储需要转置回正常图像方向 endfread的ieee-be指定大端读取这是 IDX 格式的固定要求漏掉它会读出完全错乱的维度。MNIST 长宽都是 28转不转置看不出区别但这个习惯必须保留以后换到 CIFAR 或者其他非对称数据集时permute错了图像就是横着的。GAN 训练不需要标签所以只读图像文件即可标签文件可以跳过。我一般把方案 A 当首选IDX 解析当备胎。原因很直接digitTrain4DArrayData返回的数组已经排成h×w×c×n不用管文件路径和字节序。压缩包里如果自带train-images.idx3-ubyte之类的文件再用方案 B。3.2 网络定义与维度约定生成器输出 784判别器要 28×28×1GAN 的层图比分类网络简单但维度对接必须严格。生成器每一层输出的行数要能一路乘到 784。下面这套是全连接版本的常见做法在 MNIST 上足够稳定function lgraph buildGenerator(zDim) layers [ featureInputLayer(zDim, Name, z) fullyConnectedLayer(256, Name, fc1) reluLayer(Name, relu1) fullyConnectedLayer(512, Name, fc2) reluLayer(Name, relu2) fullyConnectedLayer(28*28, Name, fc3) tanhLayer(Name, tanh_out) ]; lgraph layerGraph(layers); end function lgraph buildDiscriminator() layers [ imageInputLayer([28 28 1], Normalization, none, Name, input) convolution2dLayer(5, 16, Padding, same, Name, conv1) leakyReluLayer(0.2, Name, lrelu1) convolution2dLayer(5, 32, Padding, same, Stride, 2, Name, conv2) leakyReluLayer(0.2, Name, lrelu2) fullyConnectedLayer(1, Name, fc_out) ]; lgraph layerGraph(layers); end卷积版判别器比全连接版稳定得多因为卷积层保留了图像的空间结构D 更容易从局部纹理判断真假。convolution2dLayer(5, 16, Padding, same)的含义是 5×5 卷积核、16 个输出通道、same 填充保持空间尺寸不变。第二层加了Stride, 2把特征图尺寸减半相当于下采样。最后一层是fullyConnectedLayer(1)输出一个标量 logit真正的概率在损失函数里用sigmoid算。这里有一个最容易踩的维度坑生成器最后一层是全连接输出维度是784 × batchSize格式是CBchannel-batch而判别器的图像输入层要求SSCBspatial-spatial-channel-batch也就是28×28×1×batchSize。中间必须 reshape而且 reshape 之后要重新声明格式dlXG forward(dlnetG, z); % 784 × batch格式 CB dlXG reshape(dlXG, 28, 28, 1, []); % 先变出空间维度 dlXG dlarray(dlXG, SSCB); % 重新声明 dlarray 格式如果漏了最后一行forward(dlnetD, dlXG)会报维度不匹配的错误而且报错信息往往含糊到让你怀疑是卷积核写错了。reshape会丢掉dlarray的格式标签所以必须重新给一次。这也是为什么modelGradients函数里前向传播的顺序要固定成「生成器先跑、reshape、再进判别器」。归一化也要在这个阶段对齐。xTrain已经缩放到 [-1, 1]判别器的imageInputLayer必须写Normalization, none否则 MATLAB 会再叠一层默认归一化把 [-1, 1] 的数据二次变换导致训练集和生成器输出的分布根本不匹配。这个参数看着不起眼改错之后损失曲线会一直抖动生成图像全是灰蒙蒙的噪声。4. 训练循环骨架与超参取值0.0002、0.5 和 batch size 128 背后的调整逻辑网络定义好之后真正决定能不能出图的是训练循环和超参。这段代码是整个方案的心脏我会把完整循环写出来再拆开讲每个参数为什么是这个值。4.1 一个可直接改写的 modelGradients 函数modelGradients承担前向传播、损失计算、反向传播三件事。注意它必须通过dlfeval调用dlgradient才能正常工作function [gradG, gradD, lossG, lossD] modelGradients(dlnetG, dlnetD, dlX, z) % 生成器前向噪声 z - 784 维像素 dlXG forward(dlnetG, z); % reshape 成图像供判别器使用 dlXG reshape(dlXG, 28, 28, 1, []); dlXG dlarray(dlXG, SSCB); % 判别器对真实图像和生成图像分别打分 dReal sigmoid(forward(dlnetD, dlX)); dFake sigmoid(forward(dlnetD, dlXG)); % 损失函数加 eps 防 NaN epsVal 1e-7; lossD -mean(log(dReal epsVal)) - mean(log(1 - dFake epsVal)); lossG -mean(log(dFake epsVal)); % 反向传播G 的梯度穿过 D 网络 gradG dlgradient(lossG, dlnetG.Learnables); gradD dlgradient(lossD, dlnetD.Learnables); end用forward而不是predict是因为forward保留网络内部的中间状态在训练里是标准选择。predict更像推理模式如果你用了 batch normalization 之类的层两者行为会有差异。这里的网络没有 BN 层用forward是稳妥习惯。训练循环本体如下。数据被预先处理成 single 的 4D 数组每次迭代切片切出一个 batch避免minibatchqueue默认缓存带来的内存占用% 超参 zDim 100; batchSize 128; numEpochs 50; lr 2e-4; beta1 0.5; beta2 0.999; globalIter 0; % 构建网络 dlnetG dlnetwork(buildGenerator(zDim)); dlnetD dlnetwork(buildDiscriminator()); % Adam 状态初始化 avgG []; avgSqG []; avgD []; avgSqD []; % 固定一组噪声用于观察训练进度 zSample dlarray(randn(zDim, 25, single), CB); numIterationsPerEpoch floor(size(xTrain, 4) / batchSize); for epoch 1:numEpochs idxShuffle randperm(size(xTrain, 4)); for iter 1:numIterationsPerEpoch idx idxShuffle((iter-1)*batchSize (1:batchSize)); % 图像数据和噪声数据都要包成 dlarray并声明格式 dlX dlarray(xTrain(:, :, :, idx), SSCB); z dlarray(randn(zDim, batchSize, single), CB); % 前向 损失 梯度 [gradG, gradD, lossG, lossD] dlfeval(modelGradients, ... dlnetG, dlnetD, dlX, z); % Adam 更新两个网络 [dlnetG, avgG, avgSqG] adamupdate(dlnetG, gradG, ... avgG, avgSqG, globalIter, lr, beta1, beta2); [dlnetD, avgD, avgSqD] adamupdate(dlnetD, gradD, ... avgD, avgSqD, globalIter, lr, beta1, beta2); globalIter globalIter 1; end % 每 5 个 epoch 看一次生成效果 if mod(epoch, 5) 0 dlXSample predict(dlnetG, zSample); imSample extractdata(dlXSample); imSample (imSample 1) / 2; montage(reshape(imSample, 28, 28, 1, [])); title(sprintf(Epoch %d, epoch)); drawnow; end enddlX dlarray(xTrain(:, :, :, idx), SSCB)直接从预处理好的数组里切片整个过程不产生额外的大数组副本。randn(zDim, batchSize, single)生成噪声时显式指定single避免和 double 混用导致 MATLAB 自动转换拖慢速度。如果你的 MATLAB 版本不支持adamupdate(dlnet, ...)这种直接传网络的写法改用adamupdate(dlnet.Learnables, ...)并把结果赋值回dlnet.Learnables。两种写法在 R2020b 之后都常见旧版本只认第二种。4.2 判别器两步、生成器一步更新节奏与超参表上面的循环是 D 和 G 各更新一次1:1 节奏。DCGAN 原论文和大多数 MNIST 实现都用这个比例配合较小的学习率可以保持稳定。判别器如果明显比生成器强可以把 D 的更新包在一个内层循环里每轮 D 更新两次、G 更新一次。判别器的强弱可以从损失值判断如果dReal长期接近 1、dFake长期接近 0说明 D 碾压 G需要降低 D 的学习率或调低更新频率如果dReal在 0.5 附近徘徊且生成图像还是噪声说明 D 太弱给不出有效梯度这时要给 D 加一层卷积或增大通道数。参数取值调整方向学习率 lr2e-4训练震荡先降一半不要调大beta10.5保持不动0.9 会让训练不稳beta20.999默认值即可batchSize128内存不足改 64多样性差改 256zDim100模式崩塌时可试 128numEpochs50CPU 上 30 个 epoch 已能看出轮廓lr 2e-4是 DCGAN 的原始取值比常见分类网络低了一个数量级。GAN 的博弈特性决定了过大的学习率会让两个网络轮流碾压对方损失曲线像心电图一样乱跳。beta1 0.5是刻意压低历史梯度的影响——Adam 默认的 0.9 在 GAN 里会让训练进度过平滑反而拖慢收敛。这两个参数是 GAN 训练里最不该乱改的默认值改之前先想清楚自己要解决什么现象。5. GAN 训练避坑五个让生成图像变成马赛克的真实故障与排查顺序损失函数和循环都有了训练还是会出幺蛾子。下面五条按现象、原因、解决的顺序写基本覆盖了 MNIST 上最常见的失败模式。5.1 损失停在 0.693判别器躺平生成器开始自嗨现象判别器损失稳定在 0.693 附近生成图像是一堆噪点训练很多轮也没变化。原因0.693 是 log(0.5) 的绝对值说明 D 对真实和生成图像都输出 0.5完全分不清。这时 G 的梯度约等于零训练停摆。常见诱因是 D 学得太快或者太慢更常见的是 G 太弱生成的图像噪声太大D 觉得没必要认真分辨。解决先确认损失不是 NaN。如果稳定 0.693把 D 的更新频率降到 1:2D 更新一次G 更新两次或者把 D 的卷积通道数从 16/32 提到 32/64让 D 有足够容量给出梯度。也可以把z换方差更大的分布randn而不是rand给 G 更强的初始信号。5.2 NaN 与梯度爆炸log(0) 前先加 eps现象某个迭代步之后损失变成 NaN之后再也回不来。控制台偶尔会蹦出矩阵维度或 Inf 相关的警告。原因D 对某张假图输出概率无限接近 0log(dFake)算出负无穷或者 G 的输出值太大经过 D 的卷积后产生极大 logitsoftmax 之后梯度爆炸。GAN 里这个现象比分类网络常见得多因为两个网络互相放大对方的极端输出。解决modelGradients里所有log调用前都要加epsVal并且对dFake做一次 clampdFake min(dFake, 1 - 1e-7)。这样即使 D 的输出饱和损失也不会越界。如果 NaN 出现在修改之后把学习率从 2e-4 降到 1e-4并检查xTrain是否混入了 NaN 值。5.3 模式崩塌十个数字只剩三个现象生成的 25 张图里只有少数几种数字反复出现其他数字完全见不到。这是 GAN 最出名的失败模式。原因G 找到了一条欺骗 D 的捷径——输出某一类稳定图像能稳定骗过 D于是放弃探索其他数字类别。z 空间里大部分区域塌缩到同一个输出多样性丢失。解决先试增大 zDim 到 128给 G 更多表达空间。还可以在 D 的输入上加少量高斯噪声标准差 0.1 左右打破 D 的完美判断逼 G 去生成更多样的样本。改 batch size 也有用我见过把 batch 从 128 换到 256 之后模式崩塌直接消失的情况这里有点玄学但值得一试。5.4 生成的图像全黑或全白忘了把 tanh 输出映射回 [0, 1]现象训练正常损失正常下降montage显示的图像却是一团黑或者一团白完全看不到数字。原因生成器输出范围是 [-1, 1]montage和imshow默认把 0 当黑色、负值当纯黑正值大于 1 的全部裁成白色。你看到的全黑全白并不代表模型没学好只是显示范围和网络输出不匹配。解决显示或保存之前把像素从 [-1, 1] 映射回 [0, 1]imSample (imSample 1) / 2;然后再交给montage。保存图片时用im2uint8(imSample)转成 uint8imwrite才能写出正常 PNG。这个坑严格说不算训练 bug但几乎每个第一次在 MATLAB 里跑 GAN 的人都会遇到一次。5.5 内存溢出或训练慢到怀疑人生小心 double 和 minibatchqueue 缓存现象训练到中途报Out of Memory或者 CPU 训练速度越来越慢。原因两个常见来源。一是xTrain保持默认 double 类型6 万张图就是 376MB再叠加dlarray的自动微分图内存会翻好几倍二是minibatchqueue默认缓存多个 mini-batch虽然它能提前并行准备数据但在 MNIST 这种小数据集上反而容易把内存打到上限。解决数据加载后立刻single(xTrain)转单精度噪声生成也带single。训练循环用手动切片不经过minibatchqueue。如果你确实要换大图数据集再回到minibatchqueue并设置MiniBatchCacheSize, 4这类小缓存值。另外训练过程中每 5 个 epoch 的montage会保留上一次图像窗口关掉旧窗口也能省一点内存。6. 一个验证技巧latent 空间的线性插值判断生成器有没有死记硬背模型训完怎么知道生成器是真的学到了数字的流形还是只是把训练集里的某几张图背了下来一个简单而可靠的技巧是做 latent space 插值。取两个随机噪声向量在它们之间线性插值生成一串中间图像然后观察过渡是否平滑。% 生成两个随机噪声点 z1 randn(zDim, 1, single); z2 randn(zDim, 1, single); % 等间距插值 10 步 alphas linspace(0, 1, 10); imgs []; for a alphas z (1 - a) * z1 a * z2; z dlarray(z, CB); dlXG predict(dlnetG, z); im extractdata(dlXG); im (im 1) / 2; % 映射回 [0, 1] imgs cat(4, imgs, im); end % 排成一行显示过渡帧 montage(reshape(imgs, 28, 28, 1, []));如果中间帧是从一个数字平滑过渡到另一个数字说明生成器学会了连续的潜在表示如果中间帧是两张清晰图像的生硬切换说明模型只是记住了训练样本插值路径上没有有效的中间表达。我一般还会多看几组不同的 z1、z2并检查中间帧有没有新样式的数字出现——这比单看一组损失曲线更能说明生成质量。同一个技巧也能迁移到其他生成任务上。把噪声输入换成条件向量GAN 就可以往图像修复、图像翻译方向走插值验证逻辑完全一样。我个人的习惯是每 10 个 epoch 存一张 5×5 的网格 PNG文件名带上 epoch 号。这比依赖控制台输出更直观也让你在训练崩溃之后有机会翻出早期的生成快照定位到底是哪一步开始变坏的。这个后悔药成本极低值得养成习惯。希望帮到你。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联 返回资讯列表 →