☰
MATLAB实现生成对抗网络DCGAN:自定义训练循环与图像生成实战
2026/10/1 1:42:14 网站建设 项目流程

简介:以MATLAB为平台的生成对抗网络(GANs)代码与文档资源包,面向深度学习学习者与研究人员,围绕生成器、判别器、损失函数、优化器等核心模块,提供从理论原理到工程实现的完整链路,帮助用户解决在MATLAB环境中搭建、训练和调整GAN模型的问题。压缩包内含544个文件,其中超过360张png/jpg图片用于呈现训练过程、结果对比与可视化分析,另配有Python脚本、Markdown笔记、PDF文档、HTML辅助页面及ZIP压缩包,整体大小177.21MB,便于查阅文档、查看图示或解压复现。目前已吸引269人学习浏览。通过实际代码与配套资料,可重点理解两步梯度下降训练方式,掌握生成器与判别器的网络构建、优化器选择及损失计算等关键操作,并针对模式崩溃、收敛不稳定等常见难题积累调参经验,体会不同优化器与正则化手段对收敛效果的影响,为图像生成、超分辨率等GAN应用奠定实践基础。

1. 当生成对抗网络遇上 MATLAB:为什么这条路值得走下去

生成对抗网络(GANs)在 Python 生态里已经被聊烂了,但真正到了 MATLAB 用户手里,事情往往变成另一种局面:论文里全是 TensorFlow 代码,Github 上能找到的资源大多依赖特定 Python 版本,而你的实验环境偏偏是 MATLAB 为主,图像处理流程、数据标注、矩阵分析全都堆在这套环境里。用一句不太客气的话说,很多工程师不是不想做 GANs,而是被语言生态劝退了。这个标题直接点出了另一条路:用 MATLAB 写 GANs,配套代码、文档和调试思路一并整理成可复用的资源包,让对抗生成网络在 MATLAB 里不再是黑匣子。

这篇文章就是围绕这个需求写的。我会从判别器和生成器的对抗训练回路讲起,用一套完整可跑的 DCGAN 最小实现,带你走通数据准备、网络构建、自定义训练循环、超参数调优和常见踩坑全过程。适合的读者有两类:一类是刚接触 GANs、想在 MATLAB 里跑通第一个生成模型的初学者;另一类是已经在 Python 里调过 GANs、但被迫迁移到 MATLAB 环境完成项目交付的工程师。读完你应该能独立做出生成图像的效果,并且知道模型崩了之后该从哪里下手修。

2. 把 GANs 拆成能在 MATLAB 里建模的对抗训练回路

2.1 从生成模型到对抗训练:G 和 D 的职责边界

GANs 的核心思路是让两个网络互相博弈。生成器 G 负责把随机噪声映射成看起来像真实样本的图像,判别器 D 负责判断输入图像是来自真实数据集还是来自生成器。训练过程中,D 不断变得更敏锐,G 也不断变得更狡猾,最终达到一个纳什均衡点:G 生成的图像足以以假乱真,D 无法可靠区分真假。

这个对抗过程在 MATLAB 里建模时,要明确两件事:一是损失函数怎么定义,二是参数更新时谁先谁后。标准的做法是交替更新,一个 iteration 里先固定 G、更新 D,再固定 D、更新 G。D 的损失函数是真实图像和生成图像上的二分类交叉熵之和,G 的损失函数则是让 D 对生成图像误判为真实的交叉熵。用公式表示就是:

L_D = -E_x[log D(x)] - E_z[log(1 - D(G(z)))]

L_G = -E_z[log D(G(z))]

在 MATLAB 中实现时,不需要自己手写反向传播。自 R2021a 起,深度学习工具箱全面支持dlnetwork、dlarray、dlfeval和dlgradient,你可以直接在一个自定义函数里算出损失和梯度,剩下的交给自动微分。这个机制是后面所有代码的基础,理解它比背住任何 API 都重要。

2.2 用 dlnetwork 而不是 trainNetwork 的自定义训练路线

很多人第一次在 MATLAB 里接触深度学习,用的是trainNetwork,它会自动帮你完成前向传播、损失计算、反向传播和参数更新。但到了 GANs 这里,这套封装反而成了限制。GANs 需要两个网络交替训练,每一轮都有不同的梯度更新目标,且生成器和判别器的学习率、优化器状态必须独立维护。trainNetwork做不到这一点,所以要用dlnetwork加自定义训练循环。

下表是两条路线的典型差异,方便你在选型时做判断:

对比维度trainNetworkdlnetwork + 自定义循环
网络定义方式layerGraph 直接传给训练函数layerGraph 包成 dlnetwork
损失函数内置交叉熵,或设置自定义损失层任意 MATLAB 函数,任何数学表达式
梯度获取不可见,自动完成dlgradient 显式求出,可查看可修改
多网络交替不支持完全可控,G 和 D 分开更新
调试难度黑盒,出错定位慢可以在每个 iteration 打印中间变量

我一般建议只要是 GANs 变体,包括 DCGAN、WGAN、CycleGAN,都直接走 dlnetwork 路线。前期多写几十行胶水代码,后期调参会省非常多时间。使用dlnetwork还有一个额外好处:它允许你把网络参数作为结构体提取出来,做参数初始化、部分层冻结、梯度裁剪这些操作都更方便。

2.3 生成器与判别器的卷积结构如何设计

网络结构设计直接决定训练能不能收敛。DCGAN 的经典设计原则在 MATLAB 里依然适用,核心约束有四个:不用池化层,用带步长的卷积做下采样、用转置卷积做上采样;生成器除了最后一层外全部使用 ReLU 激活;判别器使用 LeakyReLU,斜率通常设 0.2;每个卷积层后都接批归一化,但生成器的输出层和判别器的输入层除外。

给出一份面向 28x28 灰度图像的具体结构设计,这是 DCGAN 类模型最常见的验证尺寸,训练速度快,CPU 也能在十几分钟内出效果。

生成器网络结构:

层编号层类型输出尺寸关键参数
1featureInput100x1噪声维度,不用归一化
2fullyConnected7x7x128展开后 reshape 为 4D
3transposedConv14x14x645x5 卷积核,stride 2,same
4transposedConv28x28x15x5 卷积核,stride 2,same
5tanhLayer28x28x1输出范围 -1 到 1

判别器网络结构:

层编号层类型输出尺寸关键参数
1imageInput28x28x1归一化设为 none
2convolution14x14x165x5,stride 2,same
3convolution7x7x325x5,stride 2,same
4fullyConnected1输出 logits,不加 sigmoid

判别器最后一层不加 sigmoid,而是在损失函数里通过 sigmoid 操作换算概率,这是训练稳定性的一个关键细节。如果直接在网络里加 sigmoid 层,会导致梯度更容易消失,尤其在判别器过强的时候,生成器几乎学不到东西。下一章会给出这套结构的完整可跑代码。

3. 用 DCGAN 在 MATLAB 上跑通图像生成:完整代码与参数

3.1 最简环境与数据准备:从内置手写数字开始

在写代码前,先确认环境。这个实现需要 MATLAB R2021a 以上版本,并且安装了 Deep Learning Toolbox。如果使用的是 R2023b 或更新版本,体验会更好,部分错误提示更友好。GPU 不是必需的,28x28 灰度图像的 DCGAN 在 CPU 上也能完成训练,但如果你有 NVIDIA 显卡且装了 Parallel Computing Toolbox,训练速度会提升至少一个数量级。

数据准备这里用 MATLAB 自带的手写数字数据集,避免你先花时间找数据。digitTrain4DArrayData会返回 28x28 的灰度图像和对应标签,直接加载即可。

% 加载内置手写数字数据,包含 5000 张 28x28 灰度图 [XTrain, YTrain] = digitTrain4DArrayData; % 转换为 single 类型,并归一化到 [-1, 1],对应生成器 tanh 输出范围 XTrain = single(XTrain); XTrain = (XTrain - 0.5) / 0.5; % 查看数据维度,确认是 HWC 排列 disp(size(XTrain));

这段代码里有一个容易忽略的细节:digitTrain4DArrayData返回的原始像素值范围是[0, 255]或归一化后的[0, 1]取决于版本。把它映射到[-1, 1]是因为生成器最后用了tanh激活,当真实图像和生成图像的取值范围一致时,判别器的任务才不会被截断效应干扰。如果你换用自己的图片数据集,这里要改成imageDatastore加自定义预处理函数的方式,把每张图 resize 到 28x28 并做同样的归一化。

3.2 构建生成器网络

生成器网络用layerGraph定义,然后包成dlnetwork。注意全连接层后的 reshape 操作在 MATLAB 里通过functionLayer实现,这是官方 DCGAN 示例里也在用的做法,能够保留梯度追踪链路。

function dlnetG = buildGenerator() % 生成器将 100 维随机噪声映射为 28x28x1 图像 layers = [ featureInputLayer(100, 'Normalization', 'none', 'Name', 'noise_in') fullyConnectedLayer(7*7*128, 'Name', 'fc1') reluLayer('Name', 'relu1') functionLayer(@(X) reshape(X, 7, 7, 128, []), 'Name', 'reshape') transposedConv2dLayer(5, 64, 'Stride', 2, 'Cropping', 'same', 'Name', 'deconv1') reluLayer('Name', 'relu2') transposedConv2dLayer(5, 1, 'Stride', 2, 'Cropping', 'same', 'Name', 'deconv2') tanhLayer('Name', 'tanh_out') ]; lgraph = layerGraph(layers); dlnetG = dlnetwork(lgraph); end

参数设计的一个关键点是全连接层的输出维度为7*7*128,这个数字不是随便定的。输入噪声经过第一个转置卷积后, stride 为 2 的卷积会把 7x7 放大到 14x14,第二个转置卷积再放大到 28x28,正好覆盖目标图像分辨率。如果你要生成 32x32 或 64x64 的图像,需要重新计算第一层全连接的输出尺寸,并增加相应数量的转置卷积层。transposedConv2dLayer的卷积核大小设为 5,能有效减少反卷积产生的块状伪影,比 3x3 卷积核表现更好。

3.3 构建判别器网络

判别器的结构相比生成器更简单,但有几处对稳定性影响很大的细节:使用leakyReluLayer而不是普通 ReLU,斜率 0.2;中间加入dropoutLayer,防止判别器过拟合真实样本;最后一层是全连接层输出单值,不加激活函数。

function dlnetD = buildDiscriminator() % 判别器输入 28x28x1 图像,输出未归一化的 logit layers = [ imageInputLayer([28 28 1], 'Normalization', 'none', 'Name', 'img_in') convolution2dLayer(5, 16, 'Stride', 2, 'Padding', 'same', 'Name', 'conv1') leakyReluLayer(0.2, 'Name', 'lrelu1') dropoutLayer(0.3, 'Name', 'drop1') convolution2dLayer(5, 32, 'Stride', 2, 'Padding', 'same', 'Name', 'conv2') leakyReluLayer(0.2, 'Name', 'lrelu2') dropoutLayer(0.3, 'Name', 'drop2') fullyConnectedLayer(1, 'Name', 'fc_out') ]; lgraph = layerGraph(layers); dlnetD = dlnetwork(lgraph); end

dropoutLayer的两个位置值得专门说明。判别器如果太强,损失会迅速降到接近零,生成器的梯度跟着消失,这是 GANs 训练中最常见的翻车现场。Dropout 在这里不是做正则化,而是主动削弱判别器的瞬时判断能力,给生成器留出追赶空间。另外注意,imageInputLayer的Normalization设置为'none',因为数据已经被手动归一化过,再让网络内部做一次归一化反而会改变分布。

3.4 自定义训练循环:损失函数、梯度更新与可视化

训练循环是整个项目的主干。先实现一个计算损失和梯度的函数,这个函数会被dlfeval调用,MATLAB 自动完成反向传播。核心要点是对生成器预测和判别器真实输出分别计算交叉熵。

function [lossG, lossD, gradsG, gradsD] = ganLoss(dlnetG, dlnetD, dlX, dlZ) % 前向传播,注意使用 forward 而不是 predict dlXGenerated = forward(dlnetG, dlZ); dlYPredGenerated = forward(dlnetD, dlXGenerated); dlYPredReal = forward(dlnetD, dlX); % 判别器损失:真实图像判真 + 生成图像判假 lossD = -mean(log(sigmoid(dlYPredReal))) - mean(log(1 - sigmoid(dlYPredGenerated))); % 生成器损失:让生成图像被判为真实 lossG = -mean(log(sigmoid(dlYPredGenerated))); % 分别求梯度 gradsG = dlgradient(lossG, dlnetG.Learnables); gradsD = dlgradient(lossD, dlnetD.Learnables); end

这里为什么要用forward而不是predict?原因是predict会把网络切换到推理模式,Dropout 层会被关闭,批归一化层会使用累计均值。训练过程中 Dropout 必须生效,批归一化必须使用当前 batch 的统计量,所以必须用forward。这是新手最容易踩的隐性错误之一,模型训练时看着 loss 不降,折腾半天才发现问题出在这。

主训练循环代码如下:

% 初始化网络 dlnetG = buildGenerator(); dlnetD = buildDiscriminator(); % 超参数 numEpochs = 50; miniBatchSize = 64; numLatentInputs = 100; lr = 2e-4; beta1 = 0.5; beta2 = 0.999; numIterationsPerEpoch = floor(size(XTrain, 4) / miniBatchSize); % Adam 优化器状态 avgGradG = []; avgSqGradG = []; avgGradD = []; avgSqGradD = []; iteration = 0; for epoch = 1:numEpochs % 每个 epoch 打乱数据 idx = randperm(size(XTrain, 4)); XTrain = XTrain(:,:,:,idx); for i = 1:numIterationsPerEpoch iteration = iteration + 1; % 取一个 batch 的真实图像 idxRange = (i-1)*miniBatchSize + 1 : i*miniBatchSize; dlX = dlarray(XTrain(:,:,:,idxRange), 'SSCB'); % 采样随机噪声 dlZ = dlarray(randn(numLatentInputs, miniBatchSize, 'single'), 'CB'); % 计算梯度和损失 [lossG, lossD, gradsG, gradsD] = dlfeval(@ganLoss, dlnetG, dlnetD, dlX, dlZ); % Adam 更新生成器 [dlnetG, avgGradG, avgSqGradG] = adamupdate(dlnetG, gradsG, avgGradG, avgSqGradG, iteration, lr, beta1, beta2); % Adam 更新判别器 [dlnetD, avgGradD, avgSqGradD] = adamupdate(dlnetD, gradsD, avgGradD, avgSqGradD, iteration, lr, beta1, beta2); end % 每轮 epoch 结束展示一组生成结果 dlZSample = dlarray(randn(numLatentInputs, 16, 'single'), 'CB'); dlGenerated = predict(dlnetG, dlZSample); im = extractdata(dlGenerated); im = (im + 1) / 2; montage(im, 'Size', [4 4]); title(sprintf('Epoch %d, G Loss %.4f, D Loss %.4f', epoch, extractdata(lossG), extractdata(lossD))); drawnow; end

训练循环中有三个参数值得细说。学习率lr = 2e-4是 DCGAN 论文给出的经典配置,比常规分类任务的 0.001 低五倍,目的是让两个网络更新更慢,避免一方快速压制另一方。beta1 = 0.5与默认 Adam 的 0.9 不同,这是对抗训练里降低动量影响的常见调法,防止优化器记住过去的梯度方向导致震荡。可视化时使用predict是合理的,因为此时不再需要 Dropout,用推理模式可以拿到更干净的结果。

4. 训练 GANs 的必调参数与稳定化技巧

4.1 学习率与优化器:为什么 DCGAN 的默认值会成为标准

前面代码里用了lr = 2e-4、beta1 = 0.5、beta2 = 0.999这套参数组合。这不是随手写的,而是 DCGAN 作者在大量实验后给出的经验组合,后来被各种 GANs 变体沿用。相比常规深度学习的 Adam 设置,beta1的下调是这里最反直觉的部分。常规任务中,Adam 的动量项起到加速收敛的作用;但在 GANs 中,告诉优化器“记住过去的梯度方向”会让两个网络的更新产生惯性,一旦某一方开始占优,惯性会让它优势扩大,训练崩溃的风险成倍上升。

如果你使用的是 WGAN 或 WGAN-GP 这类变体,优化器设置又会不同。WGAN 通常建议用 RMSProp 或 SGD,学习率继续维持在1e-4到5e-4之间,因为它的损失函数不再是交叉熵,而是 Wasserstein 距离估计,Adam 的动量机制在这个框架下反而容易制造不稳定的梯度。我的建议是:先从 DCGAN 的参数组合跑通,再根据 loss 曲线做微调,而不是一开始就相信网上各种“最好参数”。

4.2 标签平滑与噪声注入:让判别器不透支

训练不稳定的本质往往是判别器学会区分真假样本的速度远快于生成器学会伪造的速度。此时判别器的损失趋近于零,生成器的梯度也变得极小,整个训练过程停滞。两个简单技巧可以有效缓解这个问题。

第一个是标签平滑。把真实图像的标签从 1 换成0.9,让判别器即使判断正确也拿不到满分,留出梯度余量。实现方式不是修改数据,而是在损失函数里把真实部分的log(sigmoid(dlYPredReal))乘上0.9的系数,或在采样标签时直接生成0.9。代码改动如下:

% 原始:真实标签为 1 lossD = -mean(log(sigmoid(dlYPredReal))) - mean(log(1 - sigmoid(dlYPredGenerated))); % 平滑后:真实标签为 0.9 lossD = -0.9 * mean(log(sigmoid(dlYPredReal))) - mean(log(1 - sigmoid(dlYPredGenerated)));

第二个技巧是给判别器的输入注入高斯噪声。每轮训练时,在输入图像上叠加标准差为0.05左右的随机噪声,相当于给判别器增加分类难度。这个做法在 WGAN-GP 里被形式化为梯度惩罚,但在标准 DCGAN 里直接加噪声也是一种“穷人版”稳定化方案。注意噪声只在训练时加,生成可视化结果时不要加。

4.3 梯度裁剪与批归一化的交互陷阱

在训练循环里加梯度裁剪是防止梯度爆炸最后的手段。adamupdate本身不包含裁剪逻辑,需要手动实现。可以使用 MATLAB 的dlupdate或直接遍历梯度结构体做范数裁剪:

% 定义梯度裁剪函数,阈值设为 1 function gradients = clipGradients(gradients, threshold) layers = gradients.Layers; for i = 1:numel(layers) if isfield(layers(i), 'Weights') g = layers(i).Weights; normVal = sqrt(sum(g.^2, 'all')); if normVal > threshold layers(i).Weights = g * (threshold / normVal); end end if isfield(layers(i), 'Bias') g = layers(i).Bias; normVal = sqrt(sum(g.^2, 'all')); if normVal > threshold layers(i).Bias = g * (threshold / normVal); end end end gradients.Layers = layers; end

批归一化和梯度裁剪的交互要特别小心。批归一化层在training模式下会维护滑动平均,但如果你的网络结构把批归一化放在生成器输出层之后,裁剪梯度很容易破坏均值和方差的估计,导致生成图像的色彩饱和度出现周期性闪烁。经验法则是:生成器输出层前不要放批归一化,判别器第一层后不要放批归一化,这已经写在前面的网络构建代码里了。

4.4 何时提升 batch size:稳定性与显存的跷跷板

在 GANs 的训练中,batch size的选择往往是玄学,但对训练稳定性的影响非常直接。过小的 batch,比如 16,会导致每个 batch 的真实样本分布差异很大,判别器在不同 batch 之间来回横跳,生成器学到的特征忽明忽暗。过大 batch,比如 256,会让训练慢得不可接受,且生成器的更新步长相对判别器变得过小。

以 28x28 灰度图为例,miniBatchSize = 64是一个稳妥的起点。如果改用 128x128 甚至 256x256 的 RGB 图像,显存占用会指数级增长,此时不要强行加大 batch,而是通过梯度累积模拟大 batch 效果。具体做法是把几个小 batch 的梯度累加之后再做一次 Adam 更新,相当于用时间换空间。

5. 常见问题与避坑:DCGAN 训练现场的五次翻车记录

5.1 模式坍塌:生成器永远输出同一张图像

现象是训练到中期时,生成的 16 张图几乎全是一样的,只是像素级微小差别。判别器的 loss 没有明显异常,但生成器的多样性消失了。

这是 GANs 训练里最经典的失败模式。根因是生成器找到了一个能稳定骗过当前判别器的单一输出,而判别器没能迫使它探索其他模式。在数据分布是多峰的情况下,比如手写数字有 10 个类别,生成器完全可以只学其中一个类别并产生较低损失。

解决办法按优先级依次尝试:把判别器的学习率再降一半,比如1e-4,让判别器慢点适应;开启标签平滑,减弱判别器的置信度;在生成器的 latent code 中加入服从均匀分布的噪声,而不是纯高斯噪声。如果以上都没有效果,检查是不是生成器网络容量太低了,试着把全连接层的隐藏单元从 128 增加到 256。

5.2 loss 突然变成 NaN:梯度爆炸的前兆

现象是训练若干个 epoch 后,G loss 和 D loss 同时变成 NaN,且无法恢复。日志里看不到任何警告,代码也没有报错。

在数值层面,NaN 一般来自梯度范数爆炸,通常是某些权重更新幅度过大导致激活值溢出。也可能是学习率偏大,遇到数据中某些极端样本后参数瞬间飞走。另一种可能性是自定义损失函数里用了不稳定的数值操作,比如log(0)或log(1 - sigmoid(very_large_positive))。

解决分两步走。第一步加梯度裁剪,阈值设 1,这能在数值层面兜底;第二步把学习率降到1e-4以下,同时将生成器和判别器的参数初始化方式改为glorot或he,避免初始参数过大。做过这两步后,NaN 出现的概率会大幅降低。

5.3 生成图像出现棋盘格伪影

现象是生成的数字图像表面布满规则的网格状纹路,尤其在灰度过渡区域特别明显。图像整体轮廓是对的,但纹理很不自然。

棋盘伪影来自转置卷积的重叠区域。当卷积核大小不能被 stride 整除时,转置卷积的输出会在某些像素位置不均匀叠加,形成周期性的亮暗条纹。当 stride = 2、kernel = 5 时,重叠模式是 2x2 的棋盘。

解决的常见做法有三种。第一种最直接:把转置卷积的卷积核从 5x5 改成 4x4,会明显减少重叠,但可能丢失高频细节。第二种是在每个转置卷积后接一个普通的 3x3 卷积,让网络自己学习矫正伪影。第三种是用resize加普通卷积代替转置卷积,先双线性插值放大再卷积,彻底消除重叠问题,缺点是计算量变大。对于快速实验,我建议先试方案一,因为只需要改一行代码。

5.4 训练不动,loss 恒定在 0.69 附近

现象是最开始训练时,G loss 和 D loss 都停留在 0.69 附近,几十个 epoch 过去几乎没有变化。0.69 这个数字恰好是ln(2)的值,说明两个网络都没有学到任何东西,模型处于“随机猜测”状态。

原因通常是三个中的一个:数据归一化范围错误,生成器输出是[-1, 1],而真实图像还在[0, 1],判别器可以轻易区分两者;生成器和判别器的网络结构过于简单,没有足够的拟合能力;输入噪声维度太低,比如只有 10 维,不足以表达图像的完整变化。

排查方法是先单独检查数据预处理。把XTrain和生成器输出同时打印出来,确认取值范围一致。再看网络最后一层激活函数,生成器是tanh、判别器输出层无激活,这个组合不能错。最后把噪声维度提高到 100,这是 DCGAN 论文的标准配置。

5.5 MATLAB 中文注释乱码导致脚本无法正常保存运行

现象是代码里写了中文注释,保存后重新打开变成乱码,有时候整个文件直接无法运行。这在 MATLAB 2023a 之前的版本里非常常见,和操作系统默认编码以及 MATLAB 的编码设置有关。

原因是 MATLAB 脚本默认编码是系统区域设置决定的,中文环境下通常是 GBK,而 MATLAB 2023a 以后部分版本默认改用 UTF-8,两者不一致就会导致乱码。表现最明显的场景是git协作或跨平台拷贝代码文件。

解决方法是统一编码。在 MATLAB 主页的预设项里,将文件编码改为 UTF-8;如果已有乱码文件,用记事本打开后另存为 UTF-8 编码,再回到 MATLAB 中操作。写代码时的习惯建议是:所有脚本内注释用英文写,中文说明放在单独的 README 文档里,避免团队协作时编码冲突。

6. 用 FID 验证生成质量与 checkpoint 续训练

loss 曲线下降并不是生成质量的可靠指标,我现在的习惯是每训练完一轮,额外算一次 FID。FID 是 Fréchet Inception Distance 的缩写,核心思想是把真实图像和生成图像分别送入一个预训练分类网络,在某个特征层提取特征,然后计算两组特征分布之间的距离。FID 越小,说明生成分布与真实分布越接近。MATLAB 里可以直接用inceptionv3网络,去掉最后的分类层,用activations函数提取pool3层的特征:

net = inceptionv3; layerName = 'pool3'; featReal = activations(net, dlXReal, layerName, 'OutputAs', 'columns'); featFake = activations(net, dlXGenerated, layerName, 'OutputAs', 'columns'); muReal = mean(featReal, 2); muFake = mean(featFake, 2); sigmaReal = cov(double(featReal')); sigmaFake = cov(double(featFake')); fid = sum((muReal - muFake).^2) + trace(sigmaReal + sigmaFake - 2 * sqrtm(sigmaReal * sigmaFake));

注意inceptionv3需要的输入尺寸是 299x299,而你的数据是 28x28,所以计算前要做imresize或使用activations的自动 resize 功能。另外 FID 至少需要几百张图像才能稳定,建议使用全部的测试集而不是一个 batch。

训练中断是另一件常事,尤其当你在 CPU 上跑长时间训练时。我的做法是每隔几个 epoch 保存一次 checkpoint,保存内容包括两个网络、优化器状态和当前迭代数。续训练时要把优化器状态一并恢复,否则 Adam 的动量信息丢失,学习率等于被重置,前期训练效果前功尽弃。

save(sprintf('gan_checkpoint_epoch%d.mat', epoch), 'dlnetG', 'dlnetD', ... 'avgGradG', 'avgSqGradG', 'avgGradD', 'avgSqGradD', 'iteration');

还有一个小技巧:续训练时把学习率手动降到原来的一半,给已经接近收敛的模型一个更小的更新步长,通常能让 FID 继续下降几个点。如果用 GPU 训练,记得用gpuDevice确认显存充足,MATLAB 不会像 Python 那样提前报错,经常是训练到一半突然卡死,最后才发现是显存溢出。希望这些经验能帮你少踩几个坑。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询