MATLAB实现GAN:MNIST手写数字入门代码实战

发布时间:2026/10/1 5:35:33
MATLAB实现GAN:MNIST手写数字入门代码实战 简介这是一份在MATLAB环境中实现生成对抗网络GAN并用于MNIST手写数字数据集训练与测试的源码资源适合深度学习初学者、生成模型研究者以及希望在MATLAB平台快速上手GAN的开发者。压缩包共2个文件包含一个可直接运行的.m主脚本和一个预处理好的.mat数据文件整体大小约14.03MB结构简洁便于对照学习。已有313人学习下载。MNIST数据集内含60000个训练样本和10000个测试样本均为28×28像素的灰度手写数字图像运行代码可直观看到生成器如何从随机噪声逐步生成逼真数字以及判别器如何区分真假样本。通过阅读和修改代码读者能掌握随机噪声采样、网络层参数设置、交替训练逻辑、损失计算与图像输出等关键环节该实现也可作为模板迁移到其他图像生成或数据增强任务中是一份兼顾原理与实战的参考资料。1. GAN 的 MATLAB 实现一份能跑通的 MNIST 入门代码手写数字识别是深度学习的 Hello World但生成对抗网络GAN在 MATLAB 环境里的入门资源远没有 Python 生态那么丰富。这份压缩包里的 GANtest.m 和 mnist_uint8.mat恰好是我一直在找的稀缺组合一个主训练脚本加一份 MATLAB 格式的 MNIST 数据没有冗余的工程目录也没有依赖冲突。它把 Goodfellow 在 2014 年提出的原始 GAN 思路直接翻译成了可运行的 MATLAB 脚本你不需要先搭一套完整的深度学习框架只要会跑脚本、会读循环就能看到生成图像从纯噪声逐步逼近手写数字的整个过程。适合刚接触 GAN 的在校学生也适合需要在 MATLAB 上快速验证生成式算法的工程师。2. 源码拆解GANtest.m 和 mnist_uint8.mat 里的关键细节2.1 两个文件组成的完整训练闭环压缩包里只有两个核心文件但这两者构成了一个自包含的 GAN 训练流程。GANtest.m 是主脚本负责网络初始化、数据加载、训练循环和可视化输出mnist_uint8.mat 是 MNIST 的 MATLAB 存储格式变量通常包含 train_x、train_y、test_x、test_y。MNIST 数据集本身有 60000 个训练样本和 10000 个测试样本每个样本是 28×28 的灰度图。在 mnist_uint8.mat 里train_x 的维度是 60000×784也就是把二维图像展平成了行向量。这是原始 GAN 的标准输入方式全连接层接受一维向量784 正好对应 28×28 的像素总数。先看数据加载这一步常见写法是load(mnist_uint8.mat); train_x double(train_x) / 255; % 归一化到 [0,1]否则梯度会震荡 train_y double(train_y);这里 double() 是必要的。mat 文件里存的是 uint8如果不转成 double后续矩阵运算里 MATLAB 会频繁做隐式类型转换拖慢训练速度除以 255 则是把像素范围从 0-255 压缩到 0-1让梯度更新更平稳。这个归一化的位置是很多生成图像全是黑块或灰块问题的根源第 5 章会专门展开。注意 train_x 是 60000×784 的稠密矩阵读取时内存占用大约 60000×784×1 字节转成 double 后约为 360 MB在内存小于 8GB 的机器上要留意。2.2 网络结构推断从变量维度反推设计虽然 GANtest.m 是单脚本但通过关键变量的维度可以反推出网络结构。MNIST 的 784 维输入决定了判别器的输入层是 784 个神经元生成器的输出同样是 784 维最后经 sigmoid 或 tanh 映射到像素区间。中间隐藏层的维度常见做法是 128 或 256激活函数用 ReLU 或 tanh。原始 GAN 的判别器输出的是一个标量概率所以最后一层是经过 sigmoid 激活的单个神经元。生成器的输入是一个低维噪声向量一般是 100 维从标准正态分布里采样。生成器和判别器互为对手每次迭代里两个网络交替更新。以最朴素的全连接实现为例网络初始化大概是这样的latent_dim 100; hidden_dim 256; % 生成器 G100 - 256 - 784 G_W1 0.01 * randn(latent_dim, hidden_dim); G_b1 zeros(1, hidden_dim); G_W2 0.01 * randn(hidden_dim, 784); G_b2 zeros(1, 784); % 判别器 D784 - 256 - 1 D_W1 0.01 * randn(784, hidden_dim); D_b1 zeros(1, hidden_dim); D_W2 0.01 * randn(hidden_dim, 1); D_b2 zeros(1, 1);权重用 0.01 缩放是为了避免初始化值过大导致 sigmoid 饱和。如果直接把 randn 的结果原样用训练初期梯度会非常小生成器几乎学不到东西。这一点在 MATLAB 里跑通代码后可以自己改数值对比现象非常明显。2.3 训练循环交替更新到底在更新谁整个训练流程可以抽象成四个阶段从真实数据里采样一个小批量 MNIST 图片。从噪声分布里采样一批 z 向量过生成器得到假图片。用真实图片和假图片一起算判别器损失反向传播更新判别器参数。固定判别器再生成一批假图片去骗判别器通过判别器回传的梯度更新生成器。这个交替更新的逻辑就是 GAN 的博弈核心。在脚本里核心循环会是类似下面的结构for epoch 1:num_epochs for batch 1:num_batches % 采样真实图片 idx randi(size(train_x, 1), batch_size, 1); real_data train_x(idx, :); % 生成假图片 z randn(latent_dim, batch_size); fake_data forward_G(z); % 训练判别器真实图片给 1假图片给 0 d_real forward_D(real_data); d_fake forward_D(fake_data); d_loss -mean(log(d_real eps) log(1 - d_fake eps)); % 这里对 D 的参数做反向传播更新 % 训练生成器目标让 D 认为假图片是真实的 z randn(latent_dim, batch_size); fake_data forward_G(z); g_loss -mean(log(forward_D(fake_data) eps)); % 这里对 G 的参数做反向传播更新 end end参数说明latent_dim 一般取 100batch_size 常见取 64 或 128num_epochs 在 CPU 上可以先跑 5-10 轮做验证GPU 上可以跑到 50 轮以上。代码里的 forward_G 和 forward_D 是前向传播的函数封装实际脚本里可能直接写矩阵乘法加激活函数逻辑相同。eps 加在 log 里是为了防止概率接近 0 时出现 -Inf 的数值溢出。有个容易困惑的点判别器训练时真实样本的标签是 1生成样本的标签是 0但生成器训练时它的目标是让判别器对生成样本输出 1也就是骗过判别器。这两种损失函数看着相似但反向传播的参数范围完全不同不能混在一起更新。2.4 这份实现和现代 GAN 的差距从标签里的 nestsw 和文件名结构推测这份代码更像是教学用途的原始 GAN。它没有梯度惩罚、没有谱归一化、没有 EMA 参数平均所以训练稳定性会差一些。这是意料之中的2014 年的原始 GAN 本来就不稳定后来研究者花了好几年才把训练稳定下来。理解这一点很重要如果你发现训练曲线反复横跳不代表代码写错了而可能是这种原始 GAN 架构本身的天性。在 MATLAB 里复现时不要用 PyTorch 的成熟 GAN 实现的标准去苛责它——它要展示的是博弈学习的基本原理而不是工业级生成质量。3. 把训练跑起来从环境检查到损失曲线解读3.1 启动前的环境检查清单在双击运行 GANtest.m 之前有几个环境问题先确认清楚。首先MATLAB 版本。R2021b 以上对深度学习支持比较完整trainNetwork、dlarray 这些接口都可用R2019b 或更老版本的话部分新版 API 会报未定义函数错误需要降级写法。其次工具箱。需要 Deep Learning Toolbox。如果没安装运行时会直接报未定义函数或变量 trainNetwork或类似的错误。判断方法是在命令行敲 ver查看已安装的工具箱列表。最后路径。把解压后的文件夹设为当前路径或用 addpath 添加否则 load(mnist_uint8.mat) 会找不到文件。路径添加的写法addpath(你的解压路径/GAN文件夹);但更省事的方法是在 MATLAB 当前文件夹窗口里右键目标文件夹选择添加到路径再选选定文件夹及子文件夹。这样即使下次重启 MATLAB路径设置也会保留在预设里。有一个检查技巧值得养成习惯在运行前先敲一句exist(mnist_uint8.mat, file)如果返回 2 说明文件找得到返回 0 就说明路径没配对。这个检查在 MATLAB 里比在 Python 里更重要因为 MATLAB 的路径管理对新手是第一个拦路虎。3.2 训练脚本的核心运行逻辑如果 GANtest.m 用的是自实现的反向传播而不是 trainNetwork 高层接口它的训练循环会比 dlarray 版本更直白但也更容易出错。自实现版本里前向传播和反向传播需要手动维护每一层的缓存变量比如function [out, cache] fc_forward(X, W, b) out X * W b; % 全连接前向 cache.X X; cache.W W; % 缓存用于反向传播 end function [dX, dW, db] fc_backward(dout, cache) dW cache.X * dout; db sum(dout, 1); dX dout * cache.W; end这种写法在课程代码里很常见它清晰地展示了每一层梯度的流动方向。如果你下载的 GANtest.m 用的是这种风格那么训练主循环里会在每个 batch 结束后手动做参数更新% 参数更新lr 为学习率 G.W1 G.W1 - lr * dG_W1; G.b1 G.b1 - lr * dG_b1; G.W2 G.W2 - lr * dG_W2; G.b2 G.b2 - lr * dG_b2; D.W1 D.W1 - lr * dD_W1; D.b1 D.b1 - lr * dD_b1; D.W2 D.W2 - lr * dD_W2; D.b2 D.b2 - lr * dD_b2;学习和调参时我建议你先把这两个网络的前向输出打印出来确认维度匹配再跑完整循环。自实现代码最常见的翻车点不是数学推导错而是矩阵维度没对齐。比如生成器输出是 batch_size×784判别器输入预期是 784×batch_size中间差一个转置训练直接崩。3.3 损失曲线怎么读训练开始后你会看到两张曲线判别器损失和生成器损失。这两条线不是用来追求越低越好的它们的关系更像一场拔河。判别器损失 d_loss 持续下降说明判别器越来越能分清真假但这也意味着生成器的梯度信号越来越弱生成器会学不动。反过来生成器损失 g_loss 下降而判别器损失上升说明生成器在逐步骗过判别器。理想的情况是两条线在震荡中缓慢变化而不是一方彻底压倒另一方。如果在运行 GANtest.m 时发现 d_loss 快速降到接近 0而 g_loss 飙升到几十甚至上百几乎可以断定判别器把假样本全识别出来了生成器拿不到有效梯度这种状态在原始 GAN 里很难自动恢复。一个简单粗暴的调整是降低判别器学习率、提高生成器学习率或者给判别器做截断训练——每训一次判别器就训两次生成器。3.4 训练中途的可视化检查在 GAN 训练里判断模型有没有在学最直观的方式是定期保存生成图像。每跑完一个 epoch用当前生成器对同一组固定噪声向量生成一次图片用 subplot 排列成网格显示% 固定一组噪声每次用同一个 z 对比生成效果 z_fixed randn(latent_dim, 16); gen_imgs forward_G(z_fixed); gen_imgs reshape(gen_imgs, [28, 28, 1, 16]); figure(epoch); for i 1:16 subplot(4, 4, i); imshow(gen_imgs(:, :, 1, i), []); end drawnow;注意 reshape 前的转置因为生成器输出的每一行是一个样本转置后每列对应一个样本再按 28×28 还原图像。imshow 的第二个参数 [] 是自动拉伸灰度范围不然图像会偏暗。这个技巧几乎是 GAN 训练的刚需。只看损失曲线你很难判断生成器到底在学什么因为损失函数只是标量但每轮看一眼生成图像就能直观看到数字轮廓从模糊到清晰的演变。我一般习惯把每个 epoch 的图像用 print 或 saveas 存成 png训练结束后按时间顺序翻一遍比盯着损失曲线有效得多。4. 参数调优让生成图像从噪点变成清晰数字4.1 学习率为什么总在 1e-3 附近原始 GAN 对学习率非常敏感。在普通分类网络里学习率可以从 0.1 开始慢慢衰减但 GAN 不能这么干因为存在两个网络互相博弈学习率过高直接导致数值震荡。常见做法是用 1e-3 量级的固定学习率甚至为生成器和判别器设置不同的学习率。MATLAB 里如果用 Adam 优化器初始学习率设置可以这样genOpt adamupdate(genParams, genGrads, genState, 1e-3); disOpt adamupdate(disParams, disGrads, disState, 1e-4);在这里我特意给判别器设了更低的 1e-4给生成器设 1e-3。原因是判别器任务本质上简单得多它只需要区分真假容易先收敛一旦判别器太强生成器的梯度就会被压制。这种判别器慢、生成器快的配置是 GAN 训练的常见经验值。如果你只想改一个参数优先改这个比例。4.2 隐藏层维度与网络容量的关系隐藏层维度决定网络的表示能力。在原始 GAN 里生成器隐藏层从 128 改到 256生成图像的清晰度会有可感知的提升但训练时间也几乎翻倍。判别器的隐藏层推荐和生成器保持一致或略大一点。隐藏层维度也与数据量有关。MNIST 足够简单60000 个样本对 256 维隐藏层来说非常充裕但如果你换成更复杂的数据集比如 CIFAR-10256 维全连接网络根本不可能生成有意义的图像那时候必须上卷积架构。在调参时留意一个现象隐藏层从 128 升到 256训练中 g_loss 下降速度明显变快说明生成器容量确实提升了但再升到 512收益就不明显了反而更容易出现模式坍塌。4.3 批大小与训练轮数的取舍batch_size 的选择影响梯度估计的稳定性。64 是最常用的起点128 会让梯度更平滑但每轮更新次数减少32 则梯度噪声大容易让 GAN 训练早期不稳定。在 MATLAB 里跑这份代码时batch_size 直接决定单次矩阵运算的尺寸对于 CPU 上的小规模实验我建议先用 64。训练轮数 num_epochs 不能一概而论。这份 GANtest.m 的脚本在 CPU 上跑 10 个 epoch 大约需要十几分钟到半小时能看出数字的大致轮廓要想生成肉眼可辨的手写数字通常需要 50 轮以上。MATLAB 的循环效率不如 Python 里的 PyTorch所以如果你对每轮训练时长没概念先用 5 个 epoch 验证整个流程能跑通再逐步加轮数。4.4 判别器太强怎么办这是原始 GAN 里最常遇到的问题。判别器损失降到极低生成器损失居高不下生成图像完全是一团噪声。有几个应急手段一是降低判别器容量把隐藏层从 256 降到 128。二是给判别器加 dropout强制它不依赖固定特征。三是调整更新频率判别器每训练 1 次生成器训练 2 到 3 次用 labrador 策略也被称为拉布拉多策略平衡双方。提到 dropoutMATLAB 里可以用dropoutLayer(0.5)插入网络但如果你用的是自实现反向传播版本手动加 dropout 会比较繁琐。更快的验证方式是直接把判别器的隐藏层删一层降低它的学习能力让生成器能追上。在我自己跑类似代码的经验里调大生成器学习率比加 dropout 更直接适合新手先试。5. 避坑指南MATLAB 跑 GAN 的五个高频问题5.1 报错 未定义函数或变量 或 无法识别类现象运行 GANtest.m 时MATLAB 直接报未定义函数或提示找不到某个类的方法。原因大部分情况是 Deep Learning Toolbox 没有安装或者用到了某个特定工具箱提供的函数。偶尔也会遇到 MATLAB 路径没把当前文件夹加入搜索路径导致脚本里的辅助函数找不到。解决先用ver命令确认工具箱列表。在命令行窗口输入ver(deep)如果返回空结构说明 Deep Learning Toolbox 不存在需要重新安装或换用免工具箱的纯矩阵实现版本。如果工具箱没问题就检查当前文件夹是否在 MATLAB 路径里可以用addpath(pwd)强制添加。这个报错在新装 MATLAB 时出现概率极高几乎每个新手都会遇到一次。5.2 生成图像全是均匀灰色或纯黑块现象训练能跑完损失曲线也正常下降但 imshow 显示的生成图像是一片均匀灰色什么轮廓都看不到。原因最常见的是激活函数与数据范围不匹配。如果生成器最后一层用 sigmoid输出范围是 [0,1]而训练数据被归一化到了 [0,1]那显示时应该没问题但如果你把数据归一化成了 [-1,1]同时又用了 sigmoid 输出那么生成器无论如何输出也只能落在 [0,1]等于天然把一半的表示空间舍弃了。解决检查 mnist_uint8.mat 的归一化方式。如果训练时用了train_x double(train_x) / 255生成器输出层用 sigmoid 就合适如果用了2 * (double(train_x) / 255 - 1)的归一化输出层要换成 tanh。记住一条原则生成器输出范围必须与训练数据范围一致这是生成图像发灰或发黑的第一排查方向。5.3 训练到一半内存溢出或 MATLAB 无响应现象训练在前面几个 epoch 正常到某一步突然报内存不足或 MATLAB 窗口长时间无响应。原因常见的不是数据集本身太大而是训练过程中不断累积的变量。比如每次迭代都在保存前向传播的缓存没有在反向传播后清理或者 figure 窗口每轮都新建句柄变量越来越多。另外如果 train_x 在归一化前被反复复制double 转换后内存峰值会明显升高。解决养成两个习惯。第一在反向传播完成后用clear cache或让缓存变量在下一轮循环自然覆盖不要保存每一轮的历史值。第二不要每轮都新建 figure而是固定一个 figure 句柄用cla清空后重绘。如果内存仍然吃紧只加载 train_x 的一部分比如train_x train_x(1:10000, :);在资源有限时可以先跑小规模实验验证代码逻辑。5.4 模式坍塌所有生成图片都是同一个数字现象训练后期生成图像看起来都差不多比如全是 7 或全是 1多样性完全丢失。这是 GAN 训练里最令人头大的问题。原因模式坍塌的本质是生成器找到了一个能稳定骗过判别器的捷径只生成某一种容易以假乱真的样本而不是真正学习整个数据分布。判别器对那种样本的判断能力下降后生成器就更愿意待在舒适区。解决先降低判别器学习率限制它的能力让生成器有更多空间探索不同模式。其次增大噪声向量维度从 100 提到 128给生成器更多创作自由。还有一个在 MATLAB 里容易实现的偏方在每次训练生成器时重新采样噪声不要连续多个 batch 用同一批 z这能缓解生成器背答案式地拟合特定噪声。如果这些都不奏效就得考虑给损失函数加梯度惩罚但那已经超出原始 GAN 的改动范围了。5.5 不同 MATLAB 版本下结果不一致甚至 R2026b 跑不了现象同一份 GANtest.m在 R2021a 上能跑通在 R2023b 上报错或者换个版本后训练过程完全不一样。原因深度学习工具箱在不同版本间 API 变化较大。比如trainNetwork在 R2021b 之后对输入数据格式的要求更严格dlarray部分函数的维度排序规则也在演进新版还加强了对未定义变量和隐式扩展的警告。教育类资源代码常常基于某个特定版本编写对新版本反而兼容性不佳。解决先看报错信息是不是Error using ...加一个函数名再用which 函数名定位这个函数来自哪个工具箱或哪个版本。如果是 API 名称变了用新版对应函数替换如果是维度问题检查循环里每个矩阵的 size。对于只想跑通实验的人来说装一个和代码编写年代相近的 MATLAB 版本比如 R2020b 到 R2022b 之间是最省心的做法我在本地就同时保留了新旧两个版本专门应对这种教学代码。6. 进阶改造从原始 GAN 到 DCGAN 与条件 GAN6.1 把全连接层换成卷积层原始 GAN 用全连接处理 784 维图像信息瓶颈明显。DCGAN 的核心改动是把生成器里的全连接层换成转置卷积判别器用普通卷积替代全连接。在 MATLAB 里用网络层组合的方式写会比较清晰genLayers [ featureInputLayer(latent_dim, Normalization, none) fullyConnectedLayer(7*7*64) % 先扩展成特征图 reluLayer functionLayer((x) reshape(x, [7, 7, 64, size(x,1)])) transposedConv2dLayer(4, 64, Stride, 2, Cropping, 1) reluLayer transposedConv2dLayer(4, 1, Stride, 2, Cropping, 1) sigmoidLayer ];这组层把 100 维噪声先映射到 7×7×64 的特征图再用两个转置卷积逐步上采样到 28×28。改完之后训练稳定性通常比全连接版高一个台阶生成图像的纹理细节也更丰富值得花时间过渡到这个版本。6.2 条件 GAN 的标签注入条件 GAN 让生成器接受额外的标签信息从而控制生成的数字类别。做法是把标签做一个 one-hot 编码拼接到噪声向量后面同时喂给生成器判别器也相应的把真实样本和标签一起作为输入。MNIST 有 10 个类别所以噪声维度从 100 变成 110。在 MATLAB 里拼接操作可以放在训练循环内部% label 是当前 batch 的真实标签onehot 是 10 维向量 z randn(latent_dim, batch_size); z_cond [z; onehot]; % 拼接后维度为 110 x batch_size fake_data forward_G(z_cond);这个改动虽然不大但能让你在训练结束后指定任意数字生成对应手写样本。条件 GAN 也是后续很多应用如图像修复、超分辨率的起点在 MNIST 上把标签拼接练熟迁移到其他数据集基本是无痛的。6.3 用固定噪声模板验证生成效果一个我在多轮实战后反复验证的实用技巧固定一组噪声向量在整个训练过程中反复使用同一组 z 生成图像。因为噪声相同生成图像的每一处变化都只能归因于模型权重的更新你能非常清楚地看到数字轮廓从模糊到锐化的过程。% 记录一个观察日志方便复盘 z_probe randn(latent_dim, 16); probe_hist zeros(16, 784, num_epochs); for epoch 1:num_epochs gen_imgs forward_G(z_probe); probe_hist(:, :, epoch) gen_imgs; end % 训练完成后按 epoch 回放直观对比生成质量演进这个做法比随机采样后看单张图像更可靠因为随机采样的连续性差某一次生成效果好可能只是运气。固定噪声模板让每一步改进都具备可比较性也是论文里常用可视化方式。从这个资源出发顺着这条路径改造你对 GAN 的理解会扎实很多。从那以后我每次跑 GAN 实验都会先固定一组探针噪声再谈训练效果这个习惯帮我避开过不少看似收敛实则过拟合的假象希望帮到你。本文还有配套的精品资源点击获取