《深度学习》课件-第8章生成对抗网络_第1页
《深度学习》课件-第8章生成对抗网络_第2页
《深度学习》课件-第8章生成对抗网络_第3页
《深度学习》课件-第8章生成对抗网络_第4页
《深度学习》课件-第8章生成对抗网络_第5页
已阅读5页,还剩43页未读, 继续免费阅读

下载本文档

版权说明:本文档由用户提供并上传,收益归属内容提供方,若内容存在侵权,请进行举报或认领

文档简介

深度学习基础与实例教程生成对抗网络第八章目录02DCGAN网络DCGANnetwork01生成对抗网络GenerativeAdversarialNetworks03CGAN网络CGANnetwork生成对抗网络GenerativeAdversarialNetworks3018.1.1GAN网络模型4生成式对抗网络(GAN,GenerativeAdversarialNetworks)是一种深度学习模型,是近年来复杂分布上无监督学习最具前景的方法之一。模型通过框架中(至少)两个模块:生成模型(GenerativeModel,G)和判别模型(DiscriminativeModel,D)的互相博弈学习产生相对较好的输出。原始GAN理论中,并不要求G和D都是神经网络,只需要是能拟合相应生成和判别的函数即可。但实用中一般均使用深度神经网络作为G和D。一个优秀的GAN应用需要有良好的训练方法,否则可能因为神经网络模型的自由性而导致输出不理想。48.1.1GAN网络模型5GAN的主要结构包括一个生成器(Generator)和一个判别器(Discriminator)。如图所示。501GAN的基本结构8.1.1GAN网络模型6现在拥有大量的手写数字的数据集,希望通过GAN生成一些能够以假乱真的手写字图片。主要由如下两个部分组成:601GAN的基本结构判别器

判别器的任务是准确判断输入数据是真实的还是由生成器生成的假数据。它通常也是一个基于深度神经网络的分类器。判别器同样通过训练调整其参数,目的是提高其判断真伪的准确性。在判别器看来,真实数据应该被分类为“真”,而生成的数据则被分类为“假”。生成器的主要任务是创建尽可能逼真的假数据,以欺骗判别器。它通常是一个深度神经网络,通过随机噪声作为输入,输出与真实数据相似的数据。这个随机噪声向量通常称为“潜在空间”(latentspace)向量,它通过网络的层次结构逐渐转化成数据。生成器在训练过程中不断调整其参数(通过反向传播算法),以使产生的数据越来越难以被判别器识别。生成器8.1.1GAN网络模型7生成器和判别器在GAN中通过一个极具竞争性的互动过程进行训练,训练流程如下。702GAN的训练方式生成假数据02训练判别器03训练生成器04在训练开始之前,需要初始化生成器(Generator)和判别器(Discriminator)的参数。这些参数通常通过随机初始化的方式获得。初始化阶段0105迭代优化训练过程需要多次迭代,每一次迭代包括上述的判别器和生成器的训练步骤。通常,判别器训练的频率会高于生成器,以保持双方的平衡。生成器接收一个随机噪声向量(通常来自某种概率分布,如高斯分布)作为输入。生成器通过一个通常包含多层神经网络的模型,将这个噪声向量转化为与真实数据具有相同维度的假数据。首先将从真实数据集中抽取的样本和生成器产生的假数据进行混合,其中真样本标记为1,假样本标记为0。随后将数据输入至判别器中,判别器给出判断结果。接下来计算判别器的损失,通常使用二元交叉熵损失函数。随后使用梯度下降法(或其他优化算法)更新判别器的参数,以减少分类误差。生成器的训练过程旨在提高生成假数据的质量,以至于判别器不能轻易区分真假数据。生成器首先产生新的假数据,这些假数据再次被送入判别器进行评估,但这次,生成器的目标是让判别器将这些假数据判断为真实数据。8.1.1GAN网络模型8GAN训练过程中生成器判别器与样本可以通过图片进行展示,如图所示。802GAN的训练方式我们的目标是使用生成的样本分布(绿色实线)去拟合真实的样本分布(黑色虚线),来达到生成以假乱真样本的目的。可以看到在(a)状态处于最初始的状态的时候,生成器生成的分布和真实分布区别较大,并且判别器(蓝色虚线)判别出样本的概率不是很稳定,因此会先训练判别器来更好地分辨样本。通过多次训练判别器来达到(b)样本状态,此时判别样本区分得非常显著和良好。然后再对生成器进行训练。训练生成器之后达到(c)样本状态,此时生成器分布相比之前,逼近了真实样本分布。经过多次反复训练迭代之后,最终希望能够达到(d)状态,生成样本分布拟合于真实样本分布,并且判别器分辨不出样本是生成的还是真实的(判别概率均为0.5)。此时则表示生成器可以通过造成得到逼真的生成数据。8.1.2案例:基于GAN模型的手写数字识别实战9本节将会使用MNIST手写数字识别数据集进行GAN模型的代码实战。在机器学习中,尤其是深度学习,模型的性能依赖于大量的数据。手写数字生成可以用来扩展现有的数据集,增加样本的多样性,从而提升模型的泛化能力。生成手写数字的技术可以扩展到生成手写字母、符号甚至艺术字,应用于艺术创作和设计中。98.1.2案例:基于GAN模型的手写数字识别实战10引入依赖库1importnumpyasnpimportmatplotlib.pyplotaspltfromkeras.datasetsimportmnistfromkeras.modelsimportSequential,Modelfromkeras.layersimportDense,LeakyReLU,BatchNormalization,Reshape,Flatten,Inputfromkeras.optimizersimportAdam首先将GAN模型所需要的库引入到代码中。代码如下:8.1.2案例:基于GAN模型的手写数字识别实战11加载MNIST数据集2#加载MNIST数据集(X_train,_),(_,_)=mnist.load_data()#归一化到[-1,1]之间X_train=(X_train.astype(np.float32)-127.5)/127.5X_train=np.expand_dims(X_train,axis=3)加载代码实战所需要的MNIST手写数字数据集,并作归一化。代码如下:8.1.2案例:基于GAN模型的手写数字识别实战12构建模型3defbuild_generator():model=Sequential()model.add(Dense(256,input_dim=100))model.add(LeakyReLU(alpha=0.2))model.add(BatchNormalization(momentum=0.8))model.add(Dense(512))model.add(LeakyReLU(alpha=0.2))model.add(BatchNormalization(momentum=0.8))model.add(Dense(1024))model.add(LeakyReLU(alpha=0.2))model.add(BatchNormalization(momentum=0.8))model.add(Dense(28*28*1,activation='tanh'))model.add(Reshape((28,28,1)))returnmodel构建GAN模型中的生成器与鉴别器。代码如下:defbuild_discriminator():model=Sequential()

model.add(Flatten(input_shape=(28,28,1)))

model.add(Dense(512))

model.add(LeakyReLU(alpha=0.2))

model.add(Dense(512))

model.add(LeakyReLU(alpha=0.2))

model.add(Dense(256))

model.add(LeakyReLU(alpha=0.2))

model.add(Dense(1,activation='sigmoid'))returnmodel8.1.2案例:基于GAN模型的手写数字识别实战13编译模型4在完成生成器与鉴别器的构建后,将生成器与鉴别器组合到一起,生成器的输出是判别器的输入。对模型进行编译,在训练过程中,使用二元交叉熵函数与Adam优化器进行优化。代码如下:

#判别器discriminator=build_discriminator()pile(loss='binary_crossentropy',optimizer=Adam(0.0002,0.5),metrics=['accuracy'])#生成器generator=build_generator()#GAN模型z=Input(shape=(100,))img=generator(z)discriminator.trainable=Falsevalid=discriminator(img)combined=Model(z,valid)pile(loss='binary_crossentropy',optimizer=Adam(0.0002,0.5))8.1.2案例:基于GAN模型的手写数字识别实战14训练模型5最后对模型进行训练。每次迭代仅从数据集中随机采样一个batch作为训练鉴别器的真实样本,同时,生成同样批次的噪声作为生成器的输入,使其生成假样本。由于每次迭代中的数据量较少,对模型进行10000次迭代训练,并每隔200轮保存生成样本,以观察模型的训练情况。代码如下:

deftrain(epochs,batch_size=128,save_interval=50):#创建标签

valid=np.ones((batch_size,1))fake=np.zeros((batch_size,1))forepochinrange(epochs):#---------------------#训练判别器

#---------------------#选择一个随机批次的图像

idx=np.random.randint(0,X_train.shape[0],batch_size)

imgs=X_train[idx]#生成一个批次的噪声样本

noise=np.random.normal(0,1,(batch_size,100))

gen_imgs=generator.predict(noise)#训练判别器

d_loss_real=discriminator.train_on_batch(imgs,valid)8.1.2案例:基于GAN模型的手写数字识别实战15训练模型5d_loss_fake=discriminator.train_on_batch(gen_imgs,fake)d_loss=0.5*np.add(d_loss_real,d_loss_fake)#---------------------#训练生成器

#---------------------noise=np.random.normal(0,1,(batch_size,100))#训练生成器

g_loss=combined.train_on_batch(noise,valid)#打印进度

print(f"{epoch}[Dloss:{d_loss[0]},acc.:{100*d_loss[1]}%][Gloss:{g_loss}]")#保存生成的图像

ifepoch%save_interval==0:save_imgs(epoch)defsave_imgs(epoch):r,c=5,5noise=np.random.normal(0,1,(r*c,100))

gen_imgs=generator.predict(noise)#归一化到[0,1]之间

gen_imgs=0.5*gen_imgs+0.5fig,axs=plt.subplots(r,c)

fig.suptitle(f'Epoch{epoch}')

cnt=0foriinrange(r):forjinrange(c):

axs[i,j].imshow(gen_imgs[cnt,:,:,0],cmap='gray')

axs[i,j].axis('off')

cnt+=1

fig.savefig(f"images/mnist_{epoch}.png")

plt.close()#开始训练train(epochs=10000,batch_size=128,save_interval=200)8.1.2案例:基于GAN模型的手写数字识别实战16训练模型54/4[==============================]-0s2ms/step0[Dloss:0.7054079174995422,acc.:42.578125%][Gloss:0.7215986847877502]1/1[==============================]-0s55ms/step4/4[==============================]-0s2ms/step1[Dloss:0.35473743826150894,acc.:92.1875%][Gloss:0.7465304732322693]4/4[==============================]-0s3ms/step2[Dloss:0.3253946267068386,acc.:93.359375%][Gloss:0.8437875509262085]4/4[==============================]-0s3ms/step3[Dloss:0.30309223756194115,acc.:96.484375%][Gloss:1.011417031288147]4/4[==============================]-0s2ms/step4[Dloss:0.2710184617899358,acc.:98.828125%][Gloss:1.1895053386688232]4/4[==============================]-0s2ms/step5[Dloss:0.23247419437393546,acc.:99.609375%][Gloss:1.448891282081604]4/4[==============================]-0s3ms/step6[Dloss:0.17886821180582047,acc.:100.0%][Gloss:1.803950309753418]4/4[==============================]-0s2ms/step7[Dloss:0.1350019983947277,acc.:100.0%][Gloss:2.073486089706421]4/4[==============================]-0s2ms/step8[Dloss:0.10138186067342758,acc.:100.0%][Gloss:2.3666954040527344]4/4[==============================]-0s3ms/step9[Dloss:0.07878409419208765,acc.:100.0%][Gloss:2.6639771461486816]4/4[==============================]-0s2ms/step10[Dloss:0.059554457664489746,acc.:100.0%][Gloss:2.8840837478637695]……9990[Dloss:0.6262336075305939,acc.:68.75%][Gloss:1.1972607374191284]4/4[==============================]-0s7ms/step9991[Dloss:0.5867605209350586,acc.:69.140625%][Gloss:1.070681095123291]4/4[==============================]-0s6ms/step9992[Dloss:0.6060336828231812,acc.:63.28125%][Gloss:1.1831493377685547]控制台部分输出如下所示:8.1.2案例:基于GAN模型的手写数字识别实战17训练模型54/4[==============================]-0s6ms/step9993[Dloss:0.6228442192077637,acc.:62.5%][Gloss:1.1726912260055542]4/4[==============================]-0s6ms/step9994[Dloss:0.6101853251457214,acc.:62.109375%][Gloss:1.1897691488265991]4/4[==============================]-0s6ms/step9995[Dloss:0.617253452539444,acc.:66.40625%][Gloss:1.091201901435852]4/4[==============================]-0s6ms/step9996[Dloss:0.6438726782798767,acc.:60.15625%][Gloss:1.0807502269744873]4/4[==============================]-0s6ms/step9997[Dloss:0.5925973951816559,acc.:71.09375%][Gloss:1.1156954765319824]4/4[==============================]-0s6ms/step9998[Dloss:0.6090530753135681,acc.:66.015625%][Gloss:1.1978905200958252]4/4[==============================]-0s6ms/step9999[Dloss:0.6097795367240906,acc.:65.234375%][Gloss:1.1621670722961426]控制台部分输出如下所示:8.1.2案例:基于GAN模型的手写数字识别实战18训练模型5这里仅截取了训练前十轮与训练最后十轮的训练情况。可以发现,鉴别器仅用了5轮就实现了100%的准确率,而生成器的损失还在持续上升。可见鉴别器的训练要比生成器的训练简单得多。在训练的后十轮鉴别器的识别准确率仅有65%左右,生成器损失相较于初始阶段持续下降,并趋于稳定,此时鉴别器已经难以鉴别真实数据与生成数据,生成器生成的图像也更加逼近真实图像,迭代生成的图像如图所示。8.1.2案例:基于GAN模型的手写数字识别实战19训练模型5生成器在没有经过训练时,产生的图像为噪声图,随着循环迭代的进行,白色的噪声点逐渐聚合,随后出现了形状,生成的图像逐渐逼真。在迭代的9800轮,大部分图像已经与真实图像无异,但部分图像的生成效果依然不够逼真,存在边缘模糊等问题。DCGAN网络DCGANnetwork20028.2.1DCGAN网络模型21池化层的作用是缩小特征图的大小,但缩小特征图的代价是抛弃一些特征值。DCGAN的作者认为,直接抛弃特征值会造成特征的损失,不利于特征的学习。使用卷积替代池化仅需将卷积的步长stride设置为大于1的数值。用卷积层替代池化层使得下采样过程不再是固定的抛弃某些位置的像素值,而是可以让网络自己去学习下采样方式。用卷积层替代池化层DCGAN把GAN中的生成器G(Generator)和判别器D(Discriminator)换成了两个卷积神经网络。D可以理解为一个用于分类的卷积网络。G则是一个全卷积的生成网络。DCGAN不是简单的讲网络结构进行替换,而是对卷积神经网络的结构做了一些改进,以提高样本的质量和收敛的速度,DCGAN的改进如下:DCGAN的作者通过实验发现了全局均值池化有助于模型的稳定性。全局平均池化是通过对每个特征图进行全局平均操作,将每个特征图转化为一个单一的数值,从而减少参数数量,降低过拟合的风险,并简化模型。使用全局平均池化BN的全称是BatchNormalization,是一种用于常用于卷积层后面的归一化方法,起到帮助网络的收敛等作用。采用BN层8.2.1DCGAN网络模型22DCGAN生成器的网络结构如图所示。DCGAN在进行上采样时,使用了一个重要的卷积结构:转置卷积(TransposedConvolution)。转置卷积(TransposedConvolution),也被称为反卷积(Deconvolution)或上采样卷积(UpsamplingConvolution),是一种用于上采样的卷积操作。它的主要作用是增加特征图的空间维度(即高和宽),即增大特征图的尺寸。转置卷积的基本思想是通过填充零值的方式。8.2.1DCGAN网络模型23第一种形式是在像素值之间进行零填充,如图所示。这种方式可以均匀的放大原始的特征图,从而实现上采样的目的。8.2.1DCGAN网络模型24转置卷积的另一种形式是对特征图的边缘进行零填充,如图所示。这种方式是对原始特征图的边缘进行填充,这种方式容易造成中心区域的过采样,边缘区域的欠采样。自编码器可以用于信息检索。其主要作用在于通过学习数据的紧凑表示来提高检索的准确性和效率。自编码器通过无监督学习生成输入数据的低维表示,从而提取更具判别性和紧凑的特征表示。使用自编码器将数据编码为潜在表示后,可以在该低维空间中计算不同数据之间的相似性,提高检索速度。8.2.2案例:基于DCGAN的手写数字数据生成25本节将使用MNIST手写数字识别数据集进行DCGAN模型的代码实战。98.2.2案例:基于DCGAN的手写数字数据生成26引入依赖库1importnumpyasnpimportmatplotlib.pyplotaspltfromkeras.datasetsimportmnistfromkeras.modelsimportSequential,Modelfromkeras.layersimportDense,LeakyReLU,BatchNormalization,Reshape,Flatten,Input,Conv2DTranspose,Conv2D,Dropoutfromkeras.optimizersimportAdam首先将DCGAN模型所需要的库引入到代码中,由于DCGAN使用了卷积神经网络,因此需要导入DCGAN所需的Conv2D与Conv2DTranspose。代码如下:8.2.2案例:基于DCGAN的手写数字数据生成27加载MNIST数据集2#加载MNIST数据集(X_train,_),(_,_)=mnist.load_data()#归一化到[-1,1]之间X_train=(X_train.astype(np.float32)-127.5)/127.5X_train=np.expand_dims(X_train,axis=3)加载代码实战所需要的MNIST手写数字数据集,并作归一化,这里DCGAN与GAN的处理一致。代码如下:8.2.2案例:基于DCGAN的手写数字数据生成28构建模型3defbuild_generator():model=Sequential()model.add(Dense(7*7*256,activation="relu",input_dim=100))model.add(Reshape((7,7,256)))model.add(BatchNormalization(momentum=0.8))model.add(Conv2DTranspose(128,kernel_size=4,strides=2,padding='same',kernel_initializer='he_normal'))model.add(LeakyReLU(alpha=0.2))model.add(BatchNormalization(momentum=0.8))model.add(Conv2DTranspose(64,kernel_size=4,strides=2,padding='same',kernel_initializer='he_normal'))model.add(LeakyReLU(alpha=0.2))model.add(BatchNormalization(momentum=0.8))DCGAN中的生成器与鉴别器均为卷积神经网络,因此需要定义卷积生成器与卷积鉴别器。代码如下:

model.add(Conv2DTranspose(1,kernel_size=4,strides=1,padding='same',activation='tanh',kernel_initializer='he_normal'))returnmodeldefbuild_discriminator():model=Sequential()

model.add(Conv2D(64,kernel_size=4,strides=2,input_shape=(28,28,1),padding='same',kernel_initializer='he_normal'))

model.add(LeakyReLU(alpha=0.2))

model.add(Dropout(0.25))

model.add(Conv2D(128,kernel_size=4,strides=2,padding='same',kernel_initializer='he_normal'))

model.add(LeakyReLU(alpha=0.2))

model.add(Dropout(0.25))

model.add(Conv2D(256,kernel_size=4,strides=2,padding='same',kernel_initializer='he_normal'))

model.add(LeakyReLU(alpha=0.2))

model.add(Dropout(0.25))

model.add(Flatten())

model.add(Dense(1,activation='sigmoid'))returnmodel8.2.2案例:基于DCGAN的手写数字数据生成29编译模型4此处也可直接复用GAN中的代码。代码如下:

#判别器discriminator=build_discriminator()pile(loss='binary_crossentropy',optimizer=Adam(0.0002,0.5),metrics=['accuracy'])#生成器generator=build_generator()#GAN模型z=Input(shape=(100,))img=generator(z)discriminator.trainable=Falsevalid=discriminator(img)combined=Model(z,valid)pile(loss='binary_crossentropy',optimizer=Adam(0.0002,0.5))8.2.2案例:基于DCGAN的手写数字数据生成30训练模型5deftrain(epochs,batch_size=128,save_interval=50):#创建标签

valid=np.ones((batch_size,1))fake=np.zeros((batch_size,1))forepochinrange(epochs):#---------------------#训练判别器

#---------------------#选择一个随机批次的图像

idx=np.random.randint(0,X_train.shape[0],batch_size)imgs=X_train[idx]

#生成一个批次的噪声样本

noise=np.random.normal(0,1,(batch_size,100))gen_imgs=generator.predict(noise)#训练判别器

d_loss_real=discriminator.train_on_batch(imgs,valid)

d_loss_fake=discriminator.train_on_batch(gen_imgs,fake)

d_loss=0.5*np.add(d_loss_real,d_loss_fake)#---------------------#训练生成器

#---------------------noise=np.random.normal(0,1,(batch_size,100))#训练生成器

g_loss=combined.train_on_batch(noise,valid)#打印进度

print(f"{epoch}[Dloss:{d_loss[0]},acc.:{100*d_loss[1]}%][Gloss:{g_loss}]")#保存生成的图像

ifepoch%save_interval==0:

save_imgs(epoch)模型训练过程与GAN一致,可以直接复用。代码如下:8.2.2案例:基于DCGAN的手写数字数据生成31训练模型5defsave_imgs(epoch):r,c=5,5noise=np.random.normal(0,1,(r*c,100))gen_imgs=generator.predict(noise)#归一化到[0,1]之间

gen_imgs=0.5*gen_imgs+0.5fig,axs=plt.subplots(r,c)fig.suptitle(f'Epoch{epoch}')cnt=0foriinrange(r):forjinrange(c):axs[i,j].imshow(gen_imgs[cnt,:,:,0],cmap='gray')axs[i,j].axis('off')cnt+=1fig.savefig(f"images/mnist_{epoch}.png")plt.close()#开始训练train(epochs=10000,batch_size=128,save_interval=200)4/4[==============================]-0s21ms/step0[Dloss:1.334330976009369,acc.:14.0625%][Gloss:0.3234080672264099]4/4[==============================]-0s19ms/step1[Dloss:0.90601547062397,acc.:52.734375%][Gloss:0.5993558764457703]4/4[==============================]-0s18ms/step2[Dloss:0.1476823091506958,acc.:98.828125%][Gloss:0.9541223645210266]4/4[==============================]-0s19ms/step3[Dloss:0.035137902945280075,acc.:100.0%][Gloss:1.0736141204833984]4/4[==============================]-0s18ms/step4[Dloss:0.022186254151165485,acc.:100.0%][Gloss:1.0690982341766357]4/4[==============================]-0s19ms/step5[Dloss:0.02205378096550703,acc.:99.609375%][Gloss:1.0041770935058594]4/4[==============================]-0s19ms/step6[Dloss:0.013312488794326782,acc.:100.0%][Gloss:0.8775625228881836]4/4[==============================]-0s19ms/step7[Dloss:0.012264562770724297,acc.:100.0%][Gloss:0.8916089534759521]控制台部分输出如下所示:8.2.2案例:基于DCGAN的手写数字数据生成32训练模型54/4[==============================]-0s19ms/step8[Dloss:0.013370208442211151,acc.:100.0%][Gloss:0.8573434948921204]4/4[==============================]-0s19ms/step9[Dloss:0.014323177747428417,acc.:100.0%][Gloss:0.7378480434417725]4/4[==============================]-0s19ms/step10[Dloss:0.013277935329824686,acc.:100.0%][Gloss:0.5991072654724121]……9990[Dloss:0.5957333147525787,acc.:66.796875%][Gloss:1.0793101787567139]4/4[==============================]-0s35ms/step9991[Dloss:0.6562187075614929,acc.:58.59375%][Gloss:1.0323991775512695]4/4[==============================]-0s37ms/step9992[Dloss:0.6218872964382172,acc.:63.671875%][Gloss:1.003770351409912]4/4[==============================]-0s36ms/step9993[Dloss:0.6109737157821655,acc.:66.40625%][Gloss:1.0452868938446045]4/4[==============================]-0s36ms/step9994[Dloss:0.6423585116863251,acc.:60.9375%][Gloss:1.1632988452911377]4/4[==============================]-0s38ms/step9995[Dloss:0.6572837233543396,acc.:60.9375%][Gloss:1.1135332584381104]4/4[==============================]-0s36ms/step9996[Dloss:0.624827116727829,acc.:66.015625%][Gloss:1.039461612701416]4/4[==============================]-0s36ms/step9997[Dloss:0.6377258896827698,acc.:62.109375%][Gloss:0.9489682912826538]4/4[==============================]-0s35ms/step9998[Dloss:0.6573163568973541,acc.:60.546875%][Gloss:1.0149815082550049]4/4[==============================]-0s37ms/step9999[Dloss:0.600916177034378,acc.:66.796875%][Gloss:1.0203871726989746]控制台部分输出如下所示:8.2.2案例:基于DCGAN的手写数字数据生成33训练模型5DCGAN使用了复杂的卷积神经网络,这使得训练过程难以调参,训练过程很难达到稳定的状态。但随着迭代的持续,鉴别器的准确率也在持续下降。迭代生成的图像如图所示。8.2.2案例:基于DCGAN的手写数字数据生成34训练模型5与GAN不同的是,DCGAN在第200轮迭代就实现了白色区域的聚合,但仍与真实图像相差较远,随着后续的不断迭代,生成的图像逐渐逼真。相比于GAN,DCGAN在边缘的处理上更加真实,图像中的噪声点也更少。CGAN网络CGANnetwork3503

8.3.1CGAN网络模型36原始的GAN的生成器只能根据随机噪声进行生成图像,图像生成的结果完全取决于噪声,判别器也只能接收图像输入进行判别是否图像来使生成器。相比之下,CGAN允许用户指定要生成的数据的特定条件。这些条件可以是任何形式的附加信息,例如类别标签、文本描述、图像等。通过将条件信息输入生成器和判别器,CGAN可以学习在给定条件下生成更具结构性和多样性的数据。CGAN的结构通常包括两部分:一个生成器网络和一个判别器网络。生成器网络接收随机噪声和条件信息作为输入,并生成与条件匹配的合成数据。判别器网络接收真实数据和条件信息,或者生成器生成的数据和条件信息,然后尝试区分哪些数据是真实的,哪些是生成的。CGAN中的鉴别器与生成器均有两个输入,如图所示,下方为CGAN的生成器,上方为CGAN的鉴别器。生成器有两个输入,分别记作z与y,z表示随机噪声,y表示标签信息。标签信息会被转化为onehot编码,与随机噪声拼接后输入至生成器中,此时图像生成的信息来源既包含噪声信息,又包含标签信息;鉴别器也有两个输入,分别记作x与y,其中x表示真实或生成的图像,y表示标签信息,鉴别器不仅要判别x是否为真实图像,还要判别图像x是否属于标签y的类别。8.3.2案例:基于CGAN手写数字数据生成37本节将使用MNIST手写数字识别数据集进行CGAN模型的代码实战。98.3.2案例:基于CGAN手写数字数据生成38引入依赖库1importnumpyasnpimportmatplotlib.pyplotaspltfromkeras.datasetsimportmnistfromkeras.modelsimportSequential,Modelfromkeras.layersimportDense,LeakyReLU,BatchNormalization,Reshape,Flatten,Input,Concatenatefromkeras.optimizersimportAdamfromkeras.utilsimportto_categorical首先将CGAN模型所需要的库引入到代码中。代码如下:8.3.2案例:基于CGAN手写数字数据生成39加载MNIST数据集2#加载MNIST数据集(X_train,y_train),(_,_)=mnist.load_data()#归一化到[-1,1]之间X_train=(X_train.astype(np.float32)-127.5)/127.5X_train=np.expand_dims(X_train,axis=3)#将标签转为独热编码y_train=to_categorical(y_train,10)加载代码实战所需要的MNIST手写数字数据集,并作归一化,CGAN中不仅需要获取数据,还需要获取每个数据的标签。还要将标签转化为onehot编码。代码如下:8.3.2案例:基于CGAN手写数字数据生成40构建模型3接下来定义CGAN中的生成器与鉴别器。生成器中的输入分为噪音与标签两个部分,因此需要将两部分的输入拼接到一起,即将长度为100的噪声向量与长度为10的标签向量拼接为长度为110的输入向量。鉴别器的输入同分为两个部分,首先通过全连接层将便标签进行放大,以便能够将图像输入与标签输入拼接到一起。代码如下:defbuild_generator():

noise_shape=(100,)

label_shape=(10,)noise=Input(shape=noise_shape)label=Input(shape=label_shape)#将噪声和标签连接在一起

model_input=Concatenate()([noise,label])model=Sequential()

model.add(Dense(256,input_dim=110))#100(noise)+10(label)

model.add(LeakyReLU(alpha=0.2))

model.add(BatchNormalization(momentum=0.8))

model.add(Dense(512))

model.add(LeakyReLU(alpha=0.2))

model.add(BatchNormalization(momentum=0.8))

model.add(Dense(1024))

model.add(LeakyReLU(alpha=0.2))

model.add(BatchNormalization(momentum=0.8))

model.add(Dense(28*28*1,activation='tanh'))

model.add(Reshape((28,28,1)))

img=model(model_input)returnModel([noise,label],img)8.3.2案例:基于CGAN手写数字数据生成41构建模型3defbuild_discriminator():

img_shape=(28,28,1)

label_shape=(10,)

img=Input(shape=img_shape)label=Input(shape=label_shape)#将标签扩展为与图一样的大小

label_embedding=Dense(d(img_shape))(label)

label_embedding=Reshape(img_shape)(label_embedding)

model_input=Concatenate(axis=-1)([img,label_embedding])model=Sequential()

model.add(Flatten(input_shape=(28,28,2)))

model.add(Dense(512))

model.add(LeakyReLU(alpha=0.2))

model.add(Dense(512))

model.add(LeakyReLU(alpha=0.2))

model.add(Dense(256))

model.add(LeakyReLU(alpha=0.2))

model.add(Dense(1,activation='sigmoid'))validity=model(model_input)returnModel([img,label],validity)8.3.2案例:基于CGAN手写数字数据生成42编译模型4定义生成器与鉴别器的输入形式,并编译模型。代码如下:

#判别器discriminator=build_discriminator()pile(loss='binary_crossentropy',optimizer=Adam(0.0002,0.5),metrics=['accuracy'])#生成器generator=build_generator()#CGAN模型noise=Input(shape=(100,))label=Input(shape=(10,))img=generator([noise,label])discriminator.trainable=Falsevalid=discriminator([img,label])combined=Model([noise,label],valid)pile(loss='binary_crossentropy',optimizer=Adam(0.0002,0.5))8.3.2案例:基于CGAN手写数字数据生成43训练模型5deftrain(epochs,batch_size=128,save_interval=50):#创建标签

valid=np.ones((batch_size,1))fake=np.zeros((batch_size,1))forepochinrange(epochs):#---------------------#训练判别器

#---------------------#选择一个随机批次的图像

idx=np.random.randint(0,X_train.shape[0],batch_size)imgs=X_train[idx]labels=y_train[idx]#生成一个批次的噪声样本

noise=np.random.normal(0,1,(batch_size,100))gen_labels=np.random.randint(0,10,batch_size)gen_labels=to_categorical(gen_labels,10)gen_imgs=generator.predict([noise,gen_labels])

#训练判别器

d_loss_real=discriminator.train_on_batch([imgs,labels],valid)

d_loss_fake=discriminator.train_on_batch([gen_imgs,gen_labels],fake)

d_loss=0.5*np.add(d_loss_real,d_loss_fake)#---------------------#训练生成器

#---------------------noise=np.random.normal(0,1,(batch_size,100))

gen_labels=np.random.randint(0,10,batch_size)

gen_labels=to_categorical(gen_labels,10)#训练生成器

g_loss=combined.train_on_batch([noise,gen_labels],valid)#打印进度

print(f"{epoch}[Dloss:{d_loss[0]},acc.:{100*d_loss[1]}%][Gloss:{g_loss}]")#保存生成的图像

ifepoch%save_interval==0:

save_imgs(epoch)defsave_imgs(epoch):r,c=5,5noise=np.random.normal(0,1,(r*c,100))最后对模型进行训练。代码如下:8.3.2案例:基于CGAN手写数字数据生成44训练模型5sampled_labels=np.array([np.random.randint(0,10)for_inrange(r)fornuminrange(c)])sampled_labels_onehot=to_categorical(sampled_labels,10)gen_imgs=generator.predict([noise,sampled_labels_onehot])#归一化到[0,1]之间

gen_imgs=0.5*gen_imgs+0.5fig,axs=plt.subplots(r,c)fig.suptitle(f'Epoch{epoch}')cnt=0foriinrange(r):forjinrange(c):axs[i,j].imshow(gen_imgs[cnt,:,:,0],cmap='gray')axs[i,j].set_title(f"Label:{sampled_labels[cnt]}")axs[i,j].axis('off')cnt+=1plt.tight_layout()fig.savefig(f"images/mnist_{epoch}.png")plt.close()#开始训练train(epochs=10000,batch_size=128,save_interval=200)4/4[==============================]-0s3ms/step0[Dloss:0.6773707866668701,acc.:49.21875%][Gloss:0.44435515999794006]4/4[==============================]-0s2ms/step1[Dloss:0.416949151083827,acc.:50.0%][Gloss:0.4563189148902893]4/4[==============================]-0s2ms/step2[Dloss:0.3835413046181202,acc.

温馨提示

  • 1. 本站所有资源如无特殊说明,都需要本地电脑安装OFFICE2007和PDF阅读器。图纸软件为CAD,CAXA,PROE,UG,SolidWorks等.压缩文件请下载最新的WinRAR软件解压。
  • 2. 本站的文档不包含任何第三方提供的附件图纸等,如果需要附件,请联系上传者。文件的所有权益归上传用户所有。
  • 3. 本站RAR压缩包中若带图纸,网页内容里面会有图纸预览,若没有图纸预览就没有图纸。
  • 4. 未经权益所有人同意不得将文件中的内容挪作商业或盈利用途。
  • 5. 人人文库网仅提供信息存储空间,仅对用户上传内容的表现方式做保护处理,对用户上传分享的文档内容本身不做任何修改或编辑,并不能对任何下载内容负责。
  • 6. 下载文件中如有侵权或不适当内容,请与我们联系,我们立即纠正。
  • 7. 本站不保证下载资源的准确性、安全性和完整性, 同时也不承担用户因使用这些下载资源对自己和他人造成任何形式的伤害或损失。

评论

0/150

提交评论