ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

用Python和TensorFlow从零实现GAN:生成手写数字实战

用Python和TensorFlow从零实现GAN:生成手写数字实战 第一次跑通GAN是在一个周末的晚上。屏幕上慢慢出现一排模糊的数字从刚开始的一团噪声到逐渐有了笔画轮廓最后竟然能隐约辨认出是手写数字0到9。那一刻你就明白了为什么生成对抗网络GAN会成为这些年生成式AI浪潮里绕不开的名字。这篇文章用Python加TensorFlow从零搭一个能跑通的生成对抗网络目标是让屏幕前的你也能看到同样的“进化”过程。我会讲清楚GAN的核心原理给出可直接复制的完整代码也会分享一些我在实际调试中踩过的坑和验证过的调参经验。内容更适合刚接触深度学习的同学以及那些理论看了不少、但一直没亲手跑通代码的朋友。文中的所有代码组合起来就是一个完整脚本装好环境就能运行。1. GAN的核心设计思路造假者与验钞者的博弈1.1 用一个“造假”故事理解生成对抗网络生成对抗网络听起来很高深核心思想却非常朴素。想象一下有两个角色一个造假钞的一个验钞的。造假者不断改进工艺让假钞越来越像真的验钞者也不断修炼技术让假钞越来越难蒙混过关。两者互相博弈最终达到一种奇妙的平衡——造假者造出来的东西已经足以乱真验钞者真假难分。在这个故事里造假者就是生成器Generator验钞者就是判别器Discriminator。生成器输入一个随机噪声向量输出一张图像判别器输入一张图像输出一个0到1之间的概率值表示这张图有多“真”。训练过程就是两个网络交替进化生成器想尽办法骗过判别器判别器则拼了命分辨真假。这就是“对抗”的含义也是GAN和普通分类模型最大的区别。普通模型有明确的标签和唯一目标GAN里的两个网络却互为对手没有一个固定的“正确答案”一切都从博弈中涌现出来。理解这一点你就能明白为什么下文里的损失函数、训练循环都不是标准套路。1.2 为什么选择Python加TensorFlow先回答一个很多人纠结的问题深度学习框架那么多PyTorch现在也很火为什么这篇用TensorFlow答案很简单——对第一个GAN来说TensorFlow的Keras高层API是上手成本最低的方案之一。定义网络就像搭积木几行Sequential就能把层串起来训练循环用GradientTape逻辑直观好理解。TensorFlow 2.x之后API风格统一文档齐全遇到问题搜一下能找到大量案例这对新手极其友好。PyTorch当然也很好它的动态图机制在研究和实现复杂模型时更灵活这几年在学术界和工业界的势头确实很猛。但如果你目标是“在最短时间内用尽可能少的代码跑通一个GAN”Keras能帮你省下不少折腾样板代码的时间。先理解GAN的底层逻辑以后再切其他框架并不难底层思想完全一致。1.3 项目目标让模型从噪声中写出手写数字这篇文章要做的是训练一个GAN让生成器学会伪造MNIST数据集里的手写数字。MNIST是机器学习界的“Hello World”6万张28x28像素的灰度图内容是0到9的手写数字。图像小、模型简单、训练快对硬件要求极低普通笔记本CPU就能跑起来非常适合作为第一个GAN的训练场。当你看到生成器输出的图像从一整片噪声慢慢出现笔画、轮廓最后变成能辨认的数字那种“见证进化”的体验是其他入门项目很难给的。这也是我推荐新手从MNIST开始的原因——正反馈来得快调试成本也低。现在的图像生成模型已经能产出以假乱真的高清图片但底层那套“生成器与判别器对抗”的框架很多都能追溯到这篇文章将要讲的原理。2. 环境准备把Python和TensorFlow安置好2.1 Python版本怎么选、怎么装提到环境好多新手第一关就卡在装环境上。Python本身不难装坑都在版本兼容上。TensorFlow的更新会滞后于Python新版本所以不建议一上来就装最新的Python比较稳妥的选择是3.8到3.11之间的版本。我自己用的是3.10和TensorFlow 2.x配合很顺畅。Windows下安装时记得勾选“Add Python to PATH”不然后续在命令行里输入python会提示找不到命令。装好后打开命令行输入python --version能正常显示版本号就说明装好了。Linux和macOS下一般自带Python但版本可能旧了。macOS可以用Homebrew装新版本Linux建议用apt、yum或源码包安装也可以直接用Anaconda管理Python版本。如果不想在这些细节上花时间Anaconda是省心之选它会附带conda管理器和一套科学计算常用包。总之一句话选一个明确的版本配好PATH然后用虚拟环境隔离项目依赖。2.2 TensorFlow的两种安装方式要用TensorFlow先创建虚拟环境是个好习惯。虚拟环境可以理解成一个独立的Python小房间里面装的包互不干扰避免不同项目依赖互相打架。创建和激活的流程很简单python -m venv gan_env # Windows激活 gan_env\Scripts\activate # macOS/Linux激活 source gan_env/bin/activate激活后命令行前面会出现(gan_env)字样这时候再装TensorFlow就装进了这个小房间pip install tensorflow如果机器没有独立显卡或者不想折腾CUDA也可以只装CPU版本pip install tensorflow-cpu值得提醒的是TensorFlow在Windows上的GPU支持有一个“分水岭”2.10版本是最后一个在Windows原生环境下支持CUDA的版本之后的版本在Windows上要用GPU就得借助WSL2。所以如果打算用GPU训练最好先查清楚自己的TensorFlow版本和系统环境是否匹配。对本文这个MNIST项目来说CPU训练完全够用不必为此焦虑。如果网络环境不理想直接pip安装可能比较慢可以加国内镜像源pip install tensorflow -i https://pypi.tuna.tsinghua.edu.cn/simple这是很小但很实用的操作尤其在小水管网络下装包速度能快上好几倍。2.3 开发环境VSCode还是PyCharm编辑器这块我的建议是喜欢轻量选VSCode喜欢一站式选PyCharm。VSCode需要手动配置Python环境安装Microsoft官方Python扩展然后按CtrlShiftP调出命令面板选择“Python: Select Interpreter”指定之前创建的虚拟环境里的Python解释器。配置好之后代码补全、语法提示、调试都能正常使用对训练脚本来说完全够用。PyCharm的好处是开箱即用新建项目时可以直接选择已有的虚拟环境也可以在设置里为当前项目指定解释器。社区版免费做Python开发没有任何限制。我个人的习惯是小脚本用VSCode完整项目用PyCharm各有各的顺手之处。真正的重点是让编辑器和你当前的虚拟环境对齐很多“离奇错误”其实都是解释器选错了导致的。2.4 用一段代码验证环境是否就绪环境装好后不要急着开写先跑一段验证脚本确认TensorFlow能正常导入并检测到硬件import tensorflow as tf print(TensorFlow版本:, tf.__version__) print(可见GPU设备:, tf.config.list_physical_devices(GPU))如果能打印出版本号就说明核心环境没问题。GPU设备列表为空也没关系CPU一样能完成今天的项目只是慢一点而已。到这里环境这块就绪了可以放心进入原理和代码环节。3. 核心原理损失函数、交叉熵和对抗博弈3.1 生成器与判别器的结构设计开始写代码之前先把两个网络的结构和职责讲清楚。判别器的任务很简单输入一张784维向量28x28的MNIST图像展开成一维输出一个标量概率。它本质上就是一个二分类器我实现里用三层全连接加Dropout最后一层sigmoid输出概率值。Dropout是关键它防止判别器训练得过拟合也在训练早期削弱一点判别器的“火力”让生成器有机会学到有效梯度。生成器的任务正好相反输入一个100维的标准正态分布噪声向量输出一个784维向量经过reshape之后就是一张28x28的图像。中间层用了BatchNormalization它让每一层的输入分布更稳定能极大改善GAN训练初期梯度不稳定的情况。最后一层用tanh激活函数把输出压缩到[-1, 1]这是因为数据预处理时也把图像像素归一化到了这个范围两边对齐网络才容易学。如果你的数据范围是[0, 1]最后一层用sigmoid才匹配数据范围变了激活函数不跟着变训练基本就会陷入泥潭。3.2 原始GAN的最小最大博弈GAN的理论根基是下面这个目标函数V(D, G) E_x[log D(x)] E_z[log(1 - D(G(z)))]判别器要做的是让这个值尽量大也就是把真实图像判断为真D(x)接近1把生成图像判断为假D(G(z))接近0生成器要做的是让这个值尽量小也就是让D(G(z))接近1骗过判别器。一方的“损失”就是另一方的“机会”这就是min-max博弈。实际操作中原始的log(1-D(G(z)))在生成器训练早期会遇到梯度消失问题判别器太强时D(G(z))趋近于0log(1-D(G(z)))趋近于0但梯度极小生成器根本学不动。所以几乎所有的实战代码都会用一个变通方案——让生成器的目标变成最大化log D(G(z))也就是让生成器努力让判别器对假图输出高概率。即使D(G(z))很小时log函数的导数仍然能提供足够梯度。别小看这个改动它决定了你的GAN能不能在有限时间内训出结果。后续的WGAN、LSGAN这些改进很大一部分精力也花在解决这个“梯度不健康”的问题上。3.3 为什么原始公式的交叉熵“没有负号”很多朋友看代码时会有个著名困惑论文公式里明明是log D(x)和log(1-D(G(z)))怎么TensorFlow代码里用的是BinaryCrossentropy而且还要传ones和zeros这中间是不是少了负号、出了bug答案是公式里写的是“最大化目标”TensorFlow的优化器默认只做“最小化损失”。而BinaryCrossentropy的定义本身就是L -(y * log(ŷ) (1-y) * log(1-ŷ))它自带了一个负号。所以对于判别器真实图像的损失-log(D(x))生成图像的损失-log(1-D(G(z)))两者相加再取负就恢复了论文里“最大化”的形式最小化这个带负号的交叉熵和最大化论文里的V(D, G)在数学上是完全等价的。代码里用tf.ones_like(real_output)作为真实图像的标签、tf.zeros_like(fake_output)作为生成图像的标签就是这个思路的直接落地。搞清楚这一点再看任何框架下的GAN代码你都不会再被负号问题绕晕。3.4 GAN为什么难训没有裁判的平衡游戏关于GAN听得最多的抱怨大概就是“训练不稳定”。同样一套代码换个随机种子结果可能天差地别。这不是玄学根源在于GAN的训练目标不是简单的最小化某个损失而是要找到一个纳什均衡点。在这个点上生成器已经无法通过单方面改变来降低自己的损失判别器也一样。但实际上两个网络交替梯度更新的过程很难正好落在均衡点上。更常见的情况是判别器太强——生成图像被瞬间识破生成器梯度消失或者生成器太强——它发现某种图像最容易骗过判别器于是只生成那一类这就是大家熟知的“模式崩塌”。后面第5节我会重点讲怎么识别和处理这些情况。理解训练不稳定的根源你才能在接受“GAN本来就难训”的前提下减少试错成本。4. 完整代码实现从数据到训练的每一步4.1 导入依赖与加载MNIST数据现在正式进入代码环节。先导入需要的库并加载MNIST数据集。TensorFlow内置了MNIST的下载和读取非常方便import tensorflow as tf from tensorflow.keras import layers import numpy as np import matplotlib.pyplot as plt import time # 超参数 BUFFER_SIZE 60000 BATCH_SIZE 256 EPOCHS 100 NOISE_DIM 100 LEARNING_RATE 2e-4 BETA_1 0.5 # 加载MNIST数据集 (train_images, _), (_, _) tf.keras.datasets.mnist.load_data() train_images train_images.reshape(train_images.shape[0], 784).astype(float32) train_images (train_images - 127.5) / 127.5 train_dataset tf.data.Dataset.from_tensor_slices(train_images) train_dataset train_dataset.shuffle(BUFFER_SIZE).batch(BATCH_SIZE)这里有个细节值得多说一句为什么把像素从[0, 255]归一化到[-1, 1]因为生成器最后一层用的是tanh输出范围正好是[-1, 1]数据范围如果对不上生成器永远学不会输出正确的像素分布。减127.5再除以127.5本质就是把数据分布和网络输出的分布对齐。这种“数据分布与激活函数匹配”的意识在GAN训练里非常关键。4.2 判别器与生成器的完整定义判别器我用LeakyReLU作为激活函数。LeakyReLU在负数区间保留了一个很小的坡度不会让神经元完全死掉这对对抗训练的稳定性有帮助def make_discriminator(): model tf.keras.Sequential([ layers.Dense(512, input_shape(784,)), layers.LeakyReLU(alpha0.2), layers.Dropout(0.3), layers.Dense(256), layers.LeakyReLU(alpha0.2), layers.Dropout(0.3), layers.Dense(1, activationsigmoid) ]) return model生成器这边Dense后面接BatchNormalization和LeakyReLU。注意第一层的input_shape是NOISE_DIM也就是100维噪声def make_generator(): model tf.keras.Sequential([ layers.Dense(256, input_shape(NOISE_DIM,)), layers.BatchNormalization(), layers.LeakyReLU(alpha0.2), layers.Dense(512), layers.BatchNormalization(), layers.LeakyReLU(alpha0.2), layers.Dense(784, activationtanh) ]) return model全连接层实现GAN在MNIST上足够用了。如果数据集换成人脸或自然图像这个结构就不行必须得用卷积架构。先把全连接版本跑通理解整个训练的节奏后面再往DCGAN方向升级会顺畅很多。4.3 损失函数与优化器的选择GAN的损失函数代码量很少却是整个训练的“引擎”。一个常见误区是认为生成器和判别器共用同一个损失其实不是。判别器要区分真假所以它的损失包含真实图像和生成图像两部分生成器只需要考虑如何骗过判别器所以只计算生成图像的损失loss_object tf.keras.losses.BinaryCrossentropy() def discriminator_loss(real_output, fake_output): real_loss loss_object(tf.ones_like(real_output), real_output) fake_loss loss_object(tf.zeros_like(fake_output), fake_output) return real_loss fake_loss def generator_loss(fake_output): return loss_object(tf.ones_like(fake_output), fake_output)优化器这里有一个GAN特有的经验值——Adam的beta_1要用0.5而不是默认的0.9。这个经验来自DCGAN论文作者发现默认的Adam参数会导致训练振荡调低beta_1能显著提升稳定性。建议新手先原样沿用这个配置等跑通了再尝试修改感受一下不同超参带来的变化。generator make_generator() discriminator make_discriminator() generator_optimizer tf.keras.optimizers.Adam(learning_rateLEARNING_RATE, beta_1BETA_1) discriminator_optimizer tf.keras.optimizers.Adam(learning_rateLEARNING_RATE, beta_1BETA_1)4.4 训练循环交替更新两个网络GAN训练的核心循环每一次迭代做四件事用随机噪声生成一批假图、把真实图像和假图都送入判别器、分别计算两个网络的损失、交替更新两个网络的参数。TensorFlow 2.x里用GradientTape实现梯度记录和反向传播。两个网络各自维护自己的梯度带互不干扰然后各自的优化器更新各自的变量。这是GAN训练和普通分类网络最大的不同点不是一个模型端到端地更新而是两个模型在同一个循环里各学各的tf.function def train_step(images): noise tf.random.normal([BATCH_SIZE, NOISE_DIM]) with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape: generated_images generator(noise, trainingTrue) real_output discriminator(images, trainingTrue) fake_output discriminator(generated_images, trainingTrue) gen_loss generator_loss(fake_output) disc_loss discriminator_loss(real_output, fake_output) gradients_of_generator gen_tape.gradient(gen_loss, generator.trainable_variables) gradients_of_discriminator disc_tape.gradient(disc_loss, discriminator.trainable_variables) generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables)) discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables)) return gen_loss, disc_loss初学时建议先不要加tf.function装饰器这样方便在train_step里设置断点排查问题。等调试一切顺畅了再加上它来提速。TensorFlow会把整个函数编译成计算图训练速度快不少但出错时的报错信息也会变得不直观。有一个容易忽略的细节调用生成器和判别器时都传了trainingTrue。这个参数不是摆设——它对BatchNormalization和Dropout层的行为影响很大。训练阶段BN用的是当前batch的统计量Dropout会随机丢弃节点推理阶段BN用的是训练时的滑动均值Dropout全部保留。如果传错了训练效果会大打折扣。4.5 外层的epoch循环与结果保存有了train_step外层训练循环就不复杂了。遍历每个batch累计loss每隔若干个epoch把生成器当前生成的图像保存下来这样可以直观看到模型进步过程def generate_and_save_images(model, epoch, test_input): predictions model(test_input, trainingFalse) fig plt.figure(figsize(4, 4)) for i in range(16): plt.subplot(4, 4, i 1) plt.imshow(np.reshape(predictions[i], (28, 28)), cmapgray) plt.axis(off) plt.savefig(generated_epoch_{:03d}.png.format(epoch)) plt.close(fig) def train(dataset, epochs): test_input tf.random.normal([16, NOISE_DIM]) for epoch in range(epochs): start time.time() total_gen_loss, total_disc_loss, num_batches 0, 0, 0 for image_batch in dataset: gen_loss, disc_loss train_step(image_batch) total_gen_loss gen_loss total_disc_loss disc_loss num_batches 1 avg_gen_loss total_gen_loss / num_batches avg_disc_loss total_disc_loss / num_batches print(fEpoch {epoch1}, Gen Loss: {avg_gen_loss:.4f}, Disc Loss: {avg_disc_loss:.4f}, Time: {time.time()-start:.2f}s) if (epoch 1) % 10 0: generate_and_save_images(generator, epoch 1, test_input)把4.1到4.5节的代码整理到一个.py文件里最后加上train(train_dataset, EPOCHS)一行就是完整可运行的脚本。代码里的test_input在训练开始时就固定了这样每次保存的图像都来自同一组噪声向量前后对比起来非常直观一看就知道模型有没有进步。4.6 第一次跑通时的预期如果你直接运行前几个epoch不会出现任何数字屏幕上只会是一堆噪点。这个阶段判别器很容易就能区分真假损失会快速下降。大概十几二十个epoch之后生成的图像里会出现一些形状像数字的色块。到了30到50个epoch数字的笔画开始连贯。100个epoch下来至少能生成不少能看懂的7、1、3之类的数字。用CPU跑100个epoch大概需要半小时到一小时GPU则只需要几分钟。所以完全不必一开始就追求完美效果先把流程跑通比什么都重要。等脚本跑顺了那些“为什么loss不走”“为什么图像不清晰”的问题反而更容易在调整中理解清楚。5. 训练效果与调参实战5.1 怎么判断训练是正常还是失败GAN没有类似准确率这样的直观指标判断训练状态主要靠loss数值和生成图像质量。根据我的经验训练正常时生成器和判别器的loss通常都在0.5到2之间的区间里来回震荡不会持续下探到接近0也不会一路飙升。如果出现判别器loss一路跌到0.001以下生成器loss走高到5以上那多半是判别器太强了。这时候生成的图像会非常糟糕因为生成器根本学不到有效梯度。反过来如果判别器loss一直在2附近居高不下说明它已经分不清真假了这时生成器输出可能已经开始骗过判别器但图像未必真实。我整理了一个快速对照表方便大家判断训练状态观察到的现象可能的结论建议动作Gen和Disc loss都在0.5~2震荡训练基本正常继续训练观察图像Disc loss极低Gen loss很高判别器过强增大Disc的Dropout或减小网络层数Gen loss极低Disc loss很高生成器过强/判别器太弱适当增强Disc调整学习率两者都极高且不稳定学习率过大或结构有误降低学习率、检查归一化这些都只是经验参考不能当成严格标准。GAN的loss数值不像分类任务的accuracy那样有明确含义最终还是要落到生成图像画质上。所以我在训练时一定会定期保存图像眼见为实。5.2 模式崩塌生成器“偷懒”的经典问题模式崩塌是GAN训练中最典型的问题表现是生成器输出的图像高度重复比如16张图里全是同一个数字或者不同的噪声输入对应几乎相同的图像。原因是生成器发现某个类型的样本最容易骗过判别器于是就走捷径只输出这类样本丧失了多样性。处理办法从简单到复杂都有。先看网络结构——增加生成器的容量、把噪声维度调大一些让生成器有更多表达空间再看训练节奏——每轮迭代多更新几次生成器、少更新几次判别器给生成器更多追赶机会。更进阶的做法是引入MiniBatch Discrimination让判别器一次性看一个batch识别出整个batch高度相似的情况或者直接换用WGAN这类改良模型从损失函数层面缓解崩塌。对刚入门的朋友最简单有效的是先调整噪声维度和判别器的Dropout比例往往就能缓解。调查MONSTER问题如果你发现生成图像高度集中在一两个类别不要急着上很复杂的方案先把这些简单的参数动一动。5.3 学习率、batch size等关键参数怎么调先说结论GAN的超参数敏感度远高于普通分类网络。我踩过的坑包括学习率太大loss直接发散成NaNbatch size太小训练振荡剧烈学习率太小生成器更新太慢几十个epoch都在原地踏步实践中最常用的组合是Adam学习率2e-4beta_10.5batch size选择64或256。如果希望加速收敛可以适当提高batch size到512如果图像质量迟迟提不上去降低学习率到1e-4再慢慢磨也是一个稳妥的路径。还有一个很实用的调试技巧固定验证噪声向量。每次训练时都使用同一组随机噪声作为生成的固定输入这样每个epoch保存的图片来自同一个“起点”能非常直观地看出模型在不同阶段的差异。我有一个项目就因为坚持这个习惯很快发现了某次调参后生成效果反而倒退的问题。5.4 从全连接到DCGAN图像生成的标准进化方向本文的全连接GAN能跑通但要说生成质量距离现代GAN还差得远。跑通之后建议立刻向DCGAN升级。DCGAN把全连接层替换成卷积层用转置卷积做上采样生成图像的空间结构感会强很多。结构上大致如下def make_generator_dcgan(): model tf.keras.Sequential([ layers.Dense(7*7*128, input_shape(100,), use_biasFalse), layers.BatchNormalization(), layers.LeakyReLU(alpha0.2), layers.Reshape((7, 7, 128)), layers.Conv2DTranspose(64, (5, 5), strides(2, 2), paddingsame, use_biasFalse), layers.BatchNormalization(), layers.LeakyReLU(alpha0.2), layers.Conv2DTranspose(1, (5, 5), strides(2, 2), paddingsame, use_biasFalse, activationtanh) ]) return model def make_discriminator_dcgan(): model tf.keras.Sequential([ layers.Conv2D(64, (5, 5), strides(2, 2), paddingsame, input_shape(28, 28, 1)), layers.LeakyReLU(alpha0.2), layers.Dropout(0.3), layers.Conv2D(128, (5, 5), strides(2, 2), paddingsame), layers.BatchNormalization(), layers.LeakyReLU(alpha0.2), layers.Dropout(0.3), layers.Flatten(), layers.Dense(1, activationsigmoid) ]) return model使用DCGAN版本时数据要reshape成(-1, 28, 28, 1)是一个四维张量代表batch大小、高度、宽度、通道数。训练循环和损失函数都不用改直接把两个make函数换掉就能跑。这也是Keras高层API的好处模型结构替换起来非常方便。6. 常见问题与排错实录把藏在角落里的坑挖出来6.1 ModuleNotFoundError和各种安装问题“No module named tensorflow”是我见过最多的问题。照例先检查环境当前命令行里激活的是不是之前创建的虚拟环境激活后pip list里有没有tensorflow如果用VSCode是否在右下角状态栏选中了正确的解释器这三个环节环环相扣任何一环错了都会导致模块找不到。Windows上还有个高频报错是pip不是内部或外部命令。这基本就是安装Python时忘了勾选Add to PATH。最快的修复办法是重新运行Python安装包选择Modify并勾选添加PATH或者手动把Python安装目录加到环境变量里。这类问题很基础但几乎每周都能在社区里看到这里统一整理一下。6.2 训练时报CUDA错误如果装了GPU版TensorFlow常见报错是找不到cudart64_*.dll或者CUDA版本不匹配。最直接的解决方式要么严格按照TensorFlow官方文档安装指定版本的CUDA和cuDNN要么干脆使用CPU版本。对MNIST这个量级的任务CPU和GPU的训练时间差距远没有想象中离谱先跑通流程才是首要目标。我自己的经验是初学阶段用CPU版不碰CUDA等后面做更大的项目再来解决GPU加速会省掉很多因为环境问题产生的挫败感。毕竟GAN本身已经够难训了没必要让环境问题再添乱。6.3 loss变成NaN怎么办NaN基本等于训练“炸了”。最常见的原因是学习率过高导致梯度爆炸其次可能有数据归一化错误比如图像没有缩放到[-1,1]甚至出现无穷值。可以按顺序排查先把学习率降到1e-4甚至5e-5试试确认train_images经过(images - 127.5) / 127.5处理检查数据里有没有NaN如果都检查过还炸给梯度加一个clip操作比如在Adam里设置clipnorm1.0往往能兜底梯度裁剪不改变整体训练方向只不过把每一步更新的步长限制在安全范围是处理NaN最实用的急救手段。6.4 生成图像全黑或全白这个问题的根源几乎都在数据范围和激活函数不匹配。生成器最后一层用了tanh输出范围是[-1,1]那训练数据也必须归一化到[-1,1]。如果忘了数据归一化保留了[0,255]的原始像素生成器无论在tanh输出空间里怎么变换都无法精确落到像素值上结果就是图像一片黑或一片白。另外保存图像时记得用cmapgray否则matplotlib默认的彩色映射会得到一张诡异的红蓝渐变图看起来像是模型出了问题实际上只是可视化的问题。这类“假异常”最容易让人误判训练状态遇到颜色诡异的情况先看看是不是可视化时没指定灰度映射。6.5 常见问题速查表把以上排查思路整理成一张表方便快速定位问题现象主要原因快速解法ModuleNotFoundError解释器或虚拟环境不对确认当前解释器和pip listCUDA相关报错CUDA/cuDNN版本不匹配按官方文档匹配版本或改CPU版loss为NaN学习率过高或梯度爆炸降低学习率检查归一化图像全黑/全白数据范围与激活函数不一致数据归一化到[-1,1]生成器用tanh判别器太强Dropout不足或网络容量过大增加Dropout、降低网络层数模式崩塌生成器发现捷径提高噪声维度、调整训练节奏图像是诡异配色可视化没指定灰度映射保存时加cmapgray最后再分享一个我自己的习惯。跑GAN的时候我从来不会只盯着最终的结果图而是每隔几个epoch就把中间过程保存下来。这些连续进化的图片比任何指标都更能说明模型到底在干什么。当你看到噪声慢慢有了结构、结构慢慢变成数字你对训练过程的直觉也会越来越准。这篇文章里所有代码组合起来就是一个完整的、可以直接运行的GAN训练脚本。把环境装好把代码敲进去先跑出第一张能辨认的图片再去调参、换结构、引入更先进的方法。把基本功扎扎实实过一遍生成模型的这扇门你就已经迈进去了。
返回列表