简介:GANs生成对抗网络MATLAB实现资料包,面向深度学习研究者与初学者,解决在MATLAB环境中从零搭建和训练GAN模型的问题。资源以生成器与判别器的对抗训练为主线,详细阐述了生成器将低维随机向量映射为高维样本、判别器区分真实与生成数据的机制,并覆盖网络层构建、交叉熵损失、Adam/SGD优化器选择、两步梯度下降训练流程,以及模式崩溃、不稳定收敛等经典难点,适合课程设计、毕设和原理复现。压缩包共544个文件,大小177.21MB,以png/jpg图片、py脚本、md/pdf文档及html页面为主,其中图片数量占大多数,便于对照网络结构与训练曲线;代码脚本和文档可直接运行、二次修改,另有若干zip压缩包与文档,便于打包部署或离线阅读。目前已有269人学习,内容包含可运行的MATLAB代码示例,也整理了图像生成、超分辨率、数据增强等应用说明,可作为GAN入门与进阶的实操参考。
1. 从同名资源包说起:为什么MATLAB也能把GAN跑起来
把对抗生成网络搬进MATLAB,最现实的起点就是这些以 gansmatlab 命名的代码文档资源包。它们解决的问题很具体:你不想为了一个验证性实验去装一整套Python环境,也不想把自动微分、梯度回传这类底层细节重新踩一遍。这类包通常会把生成器、判别器、训练循环和数据集加载按脚本拆开,让你把注意力放在“改网络结构和调损失”上,而不是放在环境配置上。适合三类人:做图像增强但主力语言是MATLAB的工程师、需要合成样本扩充数据集的机械/自动化方向研究者、以及想在答辩里现场演示训练过程的同学。后面所有内容都围绕这类资源最常见的组织方式展开:手写数字生成、自定义训练循环、以及一看就懂的踩坑记录。
2. 换到MATLAB这个战场:选型理由与GAN结构先立住
2.1 Python生态再好,我也劝你先想清楚这三类场景
做GAN方向的人第一反应往往是“为什么不用PyTorch”。这个反问在绝大多数情况下是对的,Python的GAN生态、预训练模型和教程密度确实远超MATLAB,但选择工具不是选信仰,而是选落地成本。我见过三类场景,适合直接用MATLAB跑GAN:第一类是验证性实验,领导只要求一周内给出“生成样本能否辅助后续分类任务”的结论,这时候MATLAB的深度学习工具箱加上自带数据集的组合,半天就能看到第一版结果。第二类是与现有Simulink/图像处理链路打通,比如你手里已经有一套基于MATLAB的图像预处理流程,要插入一个生成模块,换成Python反而要在两种语言之间做文件中转。第三类是教学和答辩演示,MATLAB脚本的断点调试和变量可视化对“讲清楚每一步在干嘛”非常有帮助。
反过来,如果你的目标是训练一个高分辨率人脸生成模型,或者要在千万级数据集上刷SOTA指标,就不要指望MATLAB。这类资源包的定位从来不是生产级训练框架,而是让你低成本理解GAN、跑通GAN、验证GAN。装环境时记得把Deep Learning Toolbox、Parallel Computing Toolbox和Statistics and Machine Learning Toolbox一起勾上,少一个后面都会报错。
2.2 生成器、判别器与对抗损失:一个式子看懂全部结构
GAN的原始结构并不复杂,生成器G负责把随机噪声映射成伪造样本,判别器D负责判断一张图是真实数据还是生成数据。两者对抗的数学目标是:
min max V(D,G) = E[log D(x)] + E[log(1 - D(G(z)))]
判别器想让这个式子变大,也就是尽量认清真实样本;生成器想让这个式子变小,也就是尽量骗过判别器。实际训练时,生成器通常不直接用log(1-D(G(z))),而是转而最小化-log(D(G(z))),因为前者在D太强时梯度会趋近于零,生成器就学不动了。这是几乎所有GAN实现里默认替换的一步,也是你在阅读资源包代码时最先要确认的地方。
对应到MATLAB的dlgradient机制里,你要做的就是把这两个损失写进两个独立的损失函数,分别对netD.Learnables和netG.Learnables求梯度。判别器最后一层不要接sigmoid,让损失函数内部自己处理logit,这样数值上更稳。原因很简单:如果你先把输出过一遍sigmoid,再在外面套log,当logit很大时sigmoid会饱和到1或0,log(0)直接给你一个Inf,梯度也跟着爆炸。这是新手最容易踩的坑,后面避坑章节会专门说。
2.3 资源包里的代码文档长什么样:常见组织方式与读包顺序
这类以gansmatlab命名的资源包,压缩包里通常不是一堆散乱脚本,而是按“代码、文档、数据”三层组织的。拿到包以后,我不建议直接右键跑主脚本,先把目录结构看清楚。
| 目录/文件 | 里面常见内容 | 你该怎么用 |
|---|---|---|
| 代码目录 | 网络定义脚本、训练主脚本、评估脚本、工具函数 | 先改数据路径和输出目录,再逐步运行 |
| 文档目录 | README、参数说明、复现步骤、版本要求 | 先读参数说明再跑,不要跳过 |
| 数据目录 | MNIST、CIFAR或自己准备的图片文件夹 | 换成自己的数据时,先检查尺寸、通道数和取值范围 |
读包顺序我一般是这样:第一步打开README,确认它要求的MATLAB版本和工具箱,很多报错根因不是代码错了,而是版本太老。第二步打开训练主脚本,看它数据是怎么读进来的,是imageDatastore还是直接load一个.mat文件,这一步决定了你换数据时改哪里。第三步直接搜损失函数部分,看它用的是log(sigmoid())写法还是softplus写法,两种都对,但改参时候的表现不同。第四步才点运行。
新一点的资源包会用classdef把生成器和判别器封装成两个类,这是基于MATLAB自身OOP架构的写法,好处是换网络结构时不用改动训练主循环,坏处是初读代码时多了一层跳转。遇到这种结构,先看构造函数里定义的网络层,再看predict方法里的前向计算顺序,训练主循环基本不用细看。老一点的全脚本结构反而更适合第一遍跑通,变量都在工作区里,随时可以双击查看shape。
3. 把“生成手写数字”跑通:全套可复现代码与参数
3.1 数据准备:用内置手写数字库还是自建图片文件夹
第一个实验我强烈建议用手写数字集,不要一上来就上自己的业务数据。原因有三个:图像小训练快、类别结构简单、GAN是否学到了结构一眼就能看出来。MATLAB的深度学习工具箱自带一个DigitDataset,是一个按文件夹存放的图片集,用imageDatastore可以直接读,省去下载和格式转换。
% 数据准备:加载工具箱自带手写数字图片集,构造minibatch队列 digitFolder = fullfile(matlabroot, 'toolbox', 'nnet', 'nndemos', 'nndatasets', 'DigitDataset'); imds = imageDatastore(digitFolder, 'IncludeSubfolders', true, 'LabelSource', 'foldernames'); % minibatchqueue是自定义训练循环的标准数据接口 mbq = minibatchqueue(imds, ... 'MiniBatchSize', 128, ... 'MiniBatchFormat', {'SSCB'}, ... 'OutputEnvironment', 'gpu');这里MiniBatchFormat的SSCB指的是四个维度依次是空间高、空间宽、通道数、批量数,这是dlnetwork自定义训练循环里最常用的数据布局。OutputEnvironment设成gpu后,minibatchqueue会把数据自动搬到显卡上,前提是Parallel Computing Toolbox可用并且你确实有GPU。没有GPU也没关系,改成cpu就行,只是训练会慢一些。批量大小128对28x28这样的小图完全够用,如果显存紧张可以降到64。
训练循环里拿到一个batch后,记得把像素值从[0,255]归一化到[-1,1],这要和生成器最后一层tanh的输出范围匹配。后面生成图片时再映射回[0,255]做显示。
3.2 生成器与判别器:先用轻量全连接版跑通,再换卷积
这个资源包对应的最常见做法,是先给一个轻量级的全连接GAN,让你关注训练循环本身,而不是被卷积层的维度计算绊住。生成器输入一个100维噪声,输出784个像素值,对应28x28的图像。判别器正好反过来,输入784维像素,输出一个标量logit。
% 生成器:100维噪声 -> 784维像素 layersG = [ featureInputLayer(100, 'Name', 'in') fullyConnectedLayer(256, 'Name', 'fc1') reluLayer('Name', 'relu1') fullyConnectedLayer(784, 'Name', 'fc2') tanhLayer('Name', 'tanh')]; netG = dlnetwork(layersG); % 判别器:784维像素 -> 1个logit,最后一层不接sigmoid layersD = [ featureInputLayer(784, 'Name', 'in') fullyConnectedLayer(256, 'Name', 'fc1') leakyReluLayer(0.2, 'Name', 'lrelu1') fullyConnectedLayer(1, 'Name', 'fc2')]; netD = dlnetwork(layersD);featureInputLayer在这里相当于把每个像素当成一个特征通道,配合CB格式使用。全连接版的好处是代码里没有reshape层的位置争议,不会因为网络中间层的维度推导报错。判别器用leakyReLU而不是普通ReLU,是为了避免负区间梯度归零导致判别器过快地收敛到完美状态,这在GAN里是常规操作。alpha参数0.2是原版DCGAN里传下来的经验值,不用动。
等你跑通这个结构,再换成卷积版时,常见做法是生成器用fullyConnected加reshape再接transposedConv2dLayer升采样,判别器用convolution2dLayer加leakyReLU做下采样。那时要注意的是转置卷积的Cropping参数,same保证输出尺寸是输入的两倍,这一步才是真正会花时间的地方。
3.3 对抗损失与自定义训练循环:dlfeval与adamupdate的配合
训练主循环是资源包的核心,也是最值得逐行读的部分。MATLAB自定义训练循环不用手写反向传播,你要做的是把损失写清楚,然后交给dlgradient去算梯度。
% 训练主循环:30个epoch,Adam优化器 numEpochs = 30; latentDim = 100; learnRateD = 2e-4; learnRateG = 2e-4; iter = 0; trailingAvgD = []; trailingAvgSqD = []; trailingAvgG = []; trailingAvgSqG = []; for epoch = 1:numEpochs reset(mbq); while hasdata(mbq) % 读取一个batch,并映射到[-1,1] X = next(mbq); X = single(extractdata(X)); X = dlarray(reshape((X / 127.5) - 1, 784, []), 'CB'); % 随机噪声:latentDim行,batchSize列 Z = dlarray(randn(latentDim, mbq.MiniBatchSize, 'single'), 'CB'); % 分别计算两个网络的梯度和损失 [gradD, lossD] = dlfeval(@discriminatorLoss, netD, netG, X, Z); [gradG, lossG] = dlfeval(@generatorLoss, netD, netG, Z); % Adam更新:维护一阶和二阶动量 iter = iter + 1; [netD, trailingAvgD, trailingAvgSqD] = adamupdate(netD, gradD, ... trailingAvgD, trailingAvgSqD, learnRateD, iter); [netG, trailingAvgG, trailingAvgSqG] = adamupdate(netG, gradG, ... trailingAvgG, trailingAvgSqG, learnRateG, iter); end fprintf('epoch %d: D loss %.4f, G loss %.4f\n', epoch, lossD, lossG); end整个循环的核心是dlfeval。它会对内部的dlgradient调用建立计算图,计算出网络Learnables的梯度,而不是我们自己写反向传播。adamupdate是深度学习工具箱自带的优化器更新函数,你只需要传入网络、梯度、上一轮的一阶/二阶动量、学习率和当前迭代次数,它会自动完成带偏置校正的Adam更新。这里两个网络分开更新,学习率也分开定,是因为GAN训练时判别器和生成器的收敛节奏经常不一致,分开调参是必要的。训练过程中如果loss打印值出现NaN,第一个要查的就是数据里有没有NaN,第二个再怀疑损失函数数值溢出。
判别器和生成器的损失函数定义在后面,这里单独列出来。判别器要同时看真实样本和生成样本,生成器只看自己被判别的结果。
function [gradD, lossD] = discriminatorLoss(netD, netG, X, Z) % 真实样本的判别输出 YReal = forward(netD, X); % 用生成器制造假样本 XFake = forward(netG, Z); YFake = forward(netD, XFake); % 判别器损失:真实样本判真、假样本判假 lossD = -mean(log(sigmoid(YReal)) + log(1 - sigmoid(YFake))); gradD = dlgradient(lossD, netD.Learnables); end function [gradG, lossG] = generatorLoss(netD, netG, Z) % 生成假样本并送入判别器 XFake = forward(netG, Z); YFake = forward(netD, XFake); % 生成器要让假样本被判真,所以损失是 -log(D(G(z))) lossG = -mean(log(sigmoid(YFake))); gradG = dlgradient(lossG, netG.Learnables); end这里最关键的是判别器里面forward(netD, XFake)重用了同一个netD,而gradD只对netD.Learnables求导,gradG只对netG.Learnables求导,两者互不干扰。dlgradient的求导对象是函数的第二个参数,写错的话会得到空梯度,训练根本不动。另外,你可以看到判别器损失里同时出现了真实样本和生成样本,而生成器损失只依赖生成样本,这对应了2.2节里那个对抗目标式的拆解方式。sigmoid不写进网络结构里,而是放在损失计算时临时用,这就是我前面说的logit做法。
跑完30个epoch以后,生成图片的代码很简单:固定16个随机噪声,过一遍netG,再把输出reshape成图片网格。
% 生成16张图并显示 Znew = dlarray(randn(latentDim, 16, 'single'), 'CB'); Xgen = forward(netG, Znew); Xgen = extractdata(Xgen); Xgen = reshape(Xgen, 28, 28, 1, []); Xgen = (Xgen + 1) / 2; % tanh输出从[-1,1]映射回[0,1] imshow(imtile(Xgen, 'ThumbnailSize', [28 28]));extractdata把dlarray里的数值取出来变成普通数组,reshape之后用imtile拼成4x4网格。这时的图像可能还是模糊的,但应该能看到不同数字的大致轮廓。如果全是噪点或者全都是同一个图案,按参数表先排查一轮。
3.4 参数表:先说结论,再说怎么调
资源包里一般会给一组默认参数,这组参数适合小图、轻量网络。我整理了一张表,训练前把每个值在代码里对应位置找一遍,比你全部跑完了再回头查要快得多。
| 参数 | 推荐值 | 说明 |
|---|---|---|
| latentDim | 100 | 噪声维度,64到128区间都常见,太小时生成样本容易重复 |
| miniBatchSize | 128 | 显存够就128,够稳;32以下容易让判别器一步学太强 |
| learnRateD | 2e-4 | 判别器学习率,不收敛时优先调这个 |
| learnRateG | 2e-4 | 生成器学习率,一般和判别器同量级 |
| numEpochs | 30 | 全连接轻量版30到50个epoch能看出趋势 |
| 优化器 | Adam,默认beta1=0.9,beta2=0.999 | 不要轻易改成SGD |
最容易被忽略的是latentDim。它控制的是生成器的输入自由度,太小的时候生成器没有足够的信息去区分不同类别,所有输出都会往同一个模式靠。128x128以上的大图生成任务里latentDim往往还要再翻倍,这个参数在你换了网络结构以后要重新试,不是一直固定100就行的。
4. 训练曲线与参数调优:从马赛克到清晰数字的几条路
4.1 损失曲线是显示器,不是温度计
我第一次跑GAN时盯着打印出来的两个loss看,发现判别器的loss在0.69附近反复横跳,以为训练失败了,因为分类任务里0.69是很差的指标。后来才反应过来,在GAN里loss值大小本身没有绝对意义,它只反映两个网络当前的相对状态。判别器loss从高走低说明它在变强,生成器loss从低走高说明它在变弱;两个loss都不动才是最可疑的,往往意味着梯度根本没传进去。
有一种非常容易误判的情况是loss看起来“完美收敛”:判别器稳定在0.5,生成器也稳定在0.5。这个数值其实意味着判别器已经完全分不出真假了,或者更糟糕的是判别器放弃了,开始随机猜测。所以要盯训练过程,光看loss是远远不够的。我的做法是每个epoch结束都把当轮生成的16张图存到本地文件夹,翻一下历史图像就看得出生成质量是在变好还是停滞。这也是为什么4.4节的checkpoint那么重要,不存图的话,等30个epoch跑完你根本没有中间过程可回看。
4.2 学习率与Batch Size:最值得动的两个旋钮
训练GAN的绝大多数问题都出在学习率和batch size这两个参数上。学习率太高的时候,判别器会快速收敛到完美分类状态,梯度回传到生成器时直接消失,生成器彻底学不动;学习率太低的时候两个网络都在原地磨蹭,几十个epoch像没跑一样。我一般先把学习率都设成2e-4跑一轮,看输出图像有没有粗糙的数字轮廓。如果完全是噪声,把学习率降到1e-4再试;如果出现图像重复,把判别器学习率单独降到1e-4或5e-5。
Batch size的影响更隐蔽。batch太大,比如1024,判别器在一个batch内就能看到大量真实样本和假样本,它学得太快太稳,生成器根本骗不过它,最终陷入模式崩塌。batch太小,比如16,判别器每次看到的样本太少,损失曲线会剧烈震荡,训练不稳定。做手写数字这种简单数据集,128是很好的中间值。如果你在16G显存的GPU上跑,batch 128完全没有压力。
调这个方向的参数有一个很玄学的地方:GAN两个网络的学习率比值比绝对值更影响最终效果。很多资源包的代码里两个学习率写在同一行,你要刻意把它们拆开,因为改成判别器5e-5、生成器2e-4这种不对称组合,往往比两个都相同的组合收敛得更好。
4.3 标签平滑与BatchNorm位置:稳定训练的常规手法
当你发现训练慢慢变得不稳,loss曲线出现周期性尖峰时,第一个值得做的小改动是标签平滑。原版GAN的判别器目标是真实样本为1、生成样本为0,标签平滑把这改成真实样本为0.9、生成样本为0.1,相当于告诉判别器“不要对自己太自信”。
在3.3节的损失函数里,标签平滑只需要改一行:
% 标签平滑版本的判别器损失:把真实标签从1改成0.9,假标签从0改成0.1 lossD = -mean(0.9 * log(sigmoid(YReal)) + 0.1 * log(1 - sigmoid(YFake)));这个改动对训练稳定性有立竿见影的效果,尤其是全连接网络这种判别器容易收敛太快的结构。生成器损失不用改,只动判别器侧。代价是判别器的判别精度会略降,但现在GAN训练的目标本来就不是让判别器变强,而是让生成器变好,所以这点代价完全值得。
BatchNorm在GAN里的位置也是个坑。生成器里用BatchNorm基本没问题,它可以缓解初始化带来的分布偏移。判别器里加BatchNorm要小心,如果你的判别器是全连接层堆起来的,加BatchNorm反而容易造成同一batch内部样本相互影响,判别器开始依赖batch统计量而不是单样本特征。我见过不少资源包在判别器里不加BatchNorm,只靠leakyReLU和Dropout来稳定,就是这个原因。如果你换卷积结构,可以把BatchNorm放在判别器中间层,但最后一层之前务必去掉。
4.4 Checkpoint是后悔药:中断训练也能续跑
训练十几个小时后发现生成质量开始退化,你会非常需要后悔药。GAN训练的特点是前半段图像逐渐变好,后半段可能突然崩掉,但你还想回到中途那个最佳状态。所以从第一个epoch开始就存checkpoint,这是我在所有GAN项目里的固定习惯。
% 每个epoch结束后保存:网络参数 + Adam动量 save(fullfile('checkpoints', sprintf('gan_epoch_%02d.mat', epoch)), ... 'netG', 'netD', 'trailingAvgD', 'trailingAvgSqD', ... 'trailingAvgG', 'trailingAvgSqG', 'iter');恢复训练时不是只load网络权重,一定要把Adam的一阶二阶动量一起load回来,否则迭代次数归零、动量清空,优化器会认为这是一次全新的训练,梯度方向突变可能导致几个epoch内训练崩掉。迭代次数iter也要接着之前的数值继续加,adamupdate内部依赖它做偏置校正,iter重置后前几百步的学习率会被放大,表现就是恢复训练的瞬间loss暴跳一下。
保存频率一般每个epoch一次,磁盘占用很小。如果你想省事,可以只保留当前最佳版本和最终版本两份,但我还是建议每轮都存,因为“最佳”这个判断标准往往要等后期回顾才能定下来,你当时觉得差的结果,可能后面会变成重要对照样本。
5. 避坑排故:MATLAB里跑GAN最容易翻车的五个现场
5.1 Loss在几步之内变成NaN
现象是训练前几个iteration一切正常,第10或者第20个iteration开始loss打印成NaN,后面全是NaN。生成图像也随之变成一片纯色。
原因是判别器损失里的log(sigmoid(y))在y绝对值很大时出现log(0)或log(Inf)。当判别器太强、输出logit跑到20以上,sigmoid(20)约等于1,log(1-0.9999...)会得到一个很大的负数,mean之后再加负号就是Inf。另一种可能相对少见:数据里混入了NaN,比如图片读取失败后写入了坏值。
解决分两步。第一步在训练循环开头加一行assert,检查X里有没有NaN:assert(all(isfinite(extractdata(X)), 'all'))。第二步把损失从log(sigmoid)形式改成softplus形式,它天然不会产生log(0):lossD = mean(softplus(-YReal)) + mean(softplus(YFake))。这个数学变换等价于原始损失,但数值上安全得多。我做手写数字这个例子时,遇到NaN基本都是判别器先崩,先把判别器学习率从2e-4降到1e-4往往也能避开。
5.2 生成图像永远是同一张脸,或者同一个扭曲图案
现象是训练到了后期,生成器输出的16张图几乎完全一样,只是细微抖动。这叫模式崩塌,是GAN最典型的失败方式。
原因是生成器找到了一个能让判别器“及格”的捷径,把某个最容易骗过判别器的模式反复输出,不再尝试覆盖整个数据集分布。判别器在这个模式下也进入了局部最优,对这类样本的判别力不足,于是两者达成了一种不健康的平衡。
解决的优先级我建议这样排:先把判别器学习率降一个量级,让生成器有喘息空间;然后把Dropout加进判别器,让它的判断没那么稳定;如果还没改善,把latentDim从100加到128,给生成器更多自由度去表达不同类别。还有一个笨但有用的招:每次训练换一个随机种子,因为模式崩塌有时和初始化直接相关,换一种初始化可能就跑出完全不同的结果。
5.3 换电脑或换MATLAB版本后维度报错
现象是同一份代码,在你电脑上跑得好好的,换到同事的机器上就报“输入数据维度不匹配”或“Unable to infer dimensions”,有时还会指向dlnetwork的初始化过程。
原因是不同MATLAB版本对featureInputLayer和fullyConnectedLayer的维度推断规则有细微差异,特别是在R2023b前后,深度学习工具箱对format的处理行为改过一版。老的版本可能允许隐式推断,新版本严格要求显式指定。
解决是在所有关键数据上显式标注格式,不要依赖默认。训练循环里的X和Z都用dlarray(X, 'CB')这种写法带上format,网络定义里featureInputLayer也写上'Name'和'Normalization'。这两个小动作看着繁琐,但能让你从“换一次环境调两小时”里解脱出来。
5.4 GPU占用率上不去,训练像龟爬
现象是任务管理器里GPU利用率只有10%左右,风扇不转,训练速度比CPU快不了多少。
原因是数据预处理卡在CPU上。minibatchqueue的OutputEnvironment设为gpu只管数据搬运,但你的M和单次读图、resize这些步骤还是CPU在做,如果batch里图片数量太少,CPU喂数据的速度跟不上GPU计算速度,GPU就在空转等数据。
解决方法是三条线同时走:把MiniBatchSize从128加到256或512,让单次搬运的数据量变大;确认augmentedImageDatastore里没有做过于复杂的在线增强,每个epoch做随机平移缩放会大量消耗CPU;再就是检查代码里是不是每个iteration都调用了extractdata和dlarray来回转换,这些转换也会打断GPU上的计算链路。正常跑手写数字时GPU占用应该在70%以上,如果达不到,问题基本都在数据管线上。
5.5 保存图片时imwrite全黑或发灰
现象是imwrite保存生成图片,打开以后要么全黑、要么灰蒙蒙一片完全看不清数字,但imshow显示时看起来正常。
原因是生成器输出经过tanh,像素范围在[-1,1],而imwrite期望的输入是uint8或double在[0,1]范围。你把[-1,1]的数据直接传进去,负值被截断成0,正值大于1的部分又被截断成1,等于把所有细节都丢掉了。
解决是在保存前做一次映射,把[-1,1]映射回[0,255]:Xsave = uint8((Xgen + 1) * 127.5)。如果Xgen是gpuArray,还要先gather回CPU再imwrite。这个坑非常隐蔽,因为imshow不会报错,它内部会做数据范围适配,而imwrite不会,所以“看着正常、保存变黑”是典型的GAN图像输出问题。
6. 进阶验证:用FID和Inception Score给生成质量打分
6.1 FID:求真实图像与生成图像的分布距离
当你不再满足于“看起来像”时,就要引入量化指标。FID,也就是Frechet Inception Distance,是目前最常用的生成质量指标。它的思路是:把真实图像和生成图像分别送入一个预训练分类网络,取倒数第二层的特征向量,然后假设这两组特征各自服从高斯分布,计算两个分布之间的Frechet距离。FID越低,说明两组图像在特征分布上越接近,生成质量越好。
在MATLAB里做这个实验不需要从零写Inception网络,直接用深度学习工具箱里的alexnet或vgg16当特征提取器,替换最后一层后取fc7层输出作为特征向量。注意:特征提取器要用在ImageNet上预训练过的权重,而不是你自己训练的模型,因为FID要的是一个稳定、通用的特征空间。算完两组特征以后,用MATLAB的mean和cov求均值和协方差,最后套FID公式。计算中在协方差乘积后加一个极小值eps防止开根号时出现数值问题。
6.2 Inception Score:只看生成图像多样性
Inception Score衡量的是另一个维度:生成图像是否既清晰又多样。做法是把每张生成图送入分类网络,得到条件分布p(y|x),然后计算这个条件分布与边缘分布p(y)之间的KL散度。如果每张图都清晰可分类,且类别分布比较均匀,IS就高;如果生成器只会输出同一个类别的样本,即使单张很清晰,IS也会很低。
这个指标在类别数较少时容易虚高,比如手写数字只有10类,轻量GAN生成的图像类别可能集中在四五个数字上,IS算出来也不会太低。所以我会同时看FID和IS两个指标,FID偏重分布距离,IS偏重多样性和清晰度。如果FID降了但IS也降了,往往意味着生成器在牺牲多样性换取单张质量,这时候要回头检查判别器是否被Dropout省掉了。
6.3 训练收尾时给自己留的三个习惯
最后一个建议,是我做这类项目积累下来的习惯。第一,训练脚本开头固定随机种子rng(0),否则每次跑的结果不可比,你将无法判断改进是来自你的参数调整还是纯粹运气。第二,把每次实验的超参数组合写进checkpoint文件名里,比如gan_ep30_lrD1e4_lrG2e4_b128.mat,这样一个月后翻到这批文件,你不用打开代码就知道当时试了什么。第三,每轮epoch的生成图都保留,训练完以后按时间顺序翻一遍图集,哪个epoch开始变清晰、哪个epoch开始崩掉,一眼就能定下来,这比任何loss曲线都直白。
我最早跑通这版代码的时候,loss看起来收敛得很好,后来翻生成图才发现判别器只是认输了,生成器还在原地打转。从那以后我再也不只看数值,而是让图像和指标互相佐证。希望这个思路能帮你少走一段弯路,也希望你跑出来的第一版GAN能让你觉得“原来这东西真的能教会一台机器去造假”。
本文还有配套的精品资源,点击获取