( A. `) C; z1 U: Q model.add(layers.Flatten())# Z. k4 ?+ s, T/ M# O' f
model.add(layers.Dense(1)) 1 ?9 c9 ]3 a+ T, l4 a! g7 W- R ' ^2 V# B/ \8 x; N [3 b9 D! p
return model6 c4 A# q% S9 V, I' t9 W% U" L
用tf.keras.utils.plot_model( ),看一下模型结构( ~2 m# ?/ U$ N* n- Y6 m- u' z
8 i1 s. c) i4 D0 Z, n ) {. g* b' ?5 J3 w% x4 O; x5 C- _2 V4 e
- Y& c# y, F. `8 ^4 I8 J# p/ s% _
$ K3 E9 Y1 w2 O; L Z" M1 p% j
* h* V' f% R8 y5 k) m/ o用summary(),看一下模型结构和参数 % l6 E1 q& w! R" N) U : q m" k9 \8 o: m' ]" M4 T# r' E 7 q2 B- L9 h; V4 D' d8 k : \) \: B3 H: b6 y8 p9 N) k: y6 ^, W, Q# H: ~$ l" X7 {8 q W8 a
& K) R7 w @( [6 y
* c' j/ S4 {. L, P. g* m
四、定义损失函数和优化器6 D' v' X0 S/ T" x: o1 ]+ N
由于有两个模型,一个是生成器,另一个是判别器;所以要分别为两个模型定义损失函数和优化器。 2 G @) {0 k8 f) F& O7 ?2 w- |0 @" W' K P3 }& b- `
" v/ y. z/ p' y, T首先定义一个辅助函数,用于计算交叉熵损失的,这个两个模型通用。9 o0 k0 y3 i N) \7 e ?
* n1 `- e" P" z7 P0 O4 H6 I
: K) b/ M# ?" Y4 c: l W* Y
# 该方法返回计算交叉熵损失的辅助函数: a9 Z/ r+ x) t) C
cross_entropy = tf.keras.losses.BinaryCrossentropy(from_logits=True)" Z; d% t# I$ [% L
4.1 生成器的损失和优化器 % L; u5 h+ p6 g5 `4 n1)生成器损失3 n6 e" I; h$ C1 Q6 O0 b
7 O# P7 F7 A$ }/ A6 y- F , c3 G5 t, _$ ?, g( |/ |生成器损失,是量化其欺骗判别器的能力;如果生成器表现良好,判别器将会把伪造图片判断为真实图片(或1)。 0 { j4 H$ w% x: V , L8 X' B+ @8 B $ y( x& Q" a1 u5 W这里我们将把判别器在生成图片上的判断结果,与一个值全为1的数组进行对比。 & \3 w) Q, A) V( s Q' [/ ^$ ]% G- I Y8 X& K" ?
# ~* Y; h* {$ G) y3 @1 A
def generator_loss(fake_output):4 v3 M7 y$ h5 p% W- m/ _
return cross_entropy(tf.ones_like(fake_output), fake_output)* T2 g7 b& j l* T0 |
2)生成器优化器7 g' ]1 N/ D |
( w3 K" [+ ?$ D0 D/ C. V 4 x7 t" c; r6 a1 Sgenerator_optimizer = tf.keras.optimizers.Adam(1e-4) 6 U* w7 \8 \8 V7 G# q6 `4.2 判别器的损失和优化器 0 H5 X% c! i1 o4 b1)判别器损失 6 u' u" ]& X8 f: C$ ~- o: {/ n% |$ z$ F$ y5 `' e4 L+ f
% y% ^, \# \/ ]; c: Q' ^0 I
判别器损失,是量化判断真伪图片的能力。它将判别器对真实图片的预测值,与全值为1的数组进行对比;将判别器对伪造(生成的)图片的预测值,与全值为0的数组进行对比。 & x( Y! c7 `% g+ L( G7 @0 r; a# ~+ [# z" `) N2 a) i
$ X1 R6 Q# `. P" m
def discriminator_loss(real_output, fake_output): " B {% @, n5 v; A2 s a/ u9 ` real_loss = cross_entropy(tf.ones_like(real_output), real_output)5 V/ L* ?# ~ `& r
fake_loss = cross_entropy(tf.zeros_like(fake_output), fake_output). m) ]1 ^ M* l4 Z" u' y
total_loss = real_loss + fake_loss# Y: U( ~$ e; v3 s2 ^
return total_loss0 ^, [3 U* }6 b% N3 l1 o$ ?* i- t
2)判别器优化器 / \1 p7 h6 U3 u6 a+ G" H' }; w$ }) G4 d# I
! x3 g/ [" O- i( g1 q5 b
discriminator_optimizer = tf.keras.optimizers.Adam(1e-4) " l0 E0 M O& [# [五、训练模型1 J+ e1 |" j) G \3 ]9 M
5.1 保存检查点 4 ^$ E, E& p3 @# ?; b保存检查点,能帮助保存和恢复模型,在长时间训练任务被中断的情况下比较有帮助。 % }2 u, W+ m. z1 D, S. U. d. a 3 e0 c% {* b' c : c7 L w2 q0 @* ?checkpoint_dir = './training_checkpoints'( w) w3 F) i: M: o; {1 F
checkpoint_prefix = os.path.join(checkpoint_dir, "ckpt") 0 P0 ^0 N8 N1 P# Kcheckpoint = tf.train.Checkpoint(generator_optimizer=generator_optimizer,3 a) P9 q+ g# @' M6 i8 o
discriminator_optimizer=discriminator_optimizer, 5 U4 F0 q# r, H: h) |; l1 M3 K0 i# L& q generator=generator,+ t' X" Q5 V: ^ B
discriminator=discriminator) 8 R& A0 M% Q& t0 V' F- B+ c& `) ^% p5.2 定义训练过程 7 c5 Y2 z8 f2 e* G' Q, l7 f$ |3 O3 eEPOCHS = 50; b0 k) p, G- _6 ^* R8 @
noise_dim = 100% r- I- p, R3 s, n* E- c1 @
num_examples_to_generate = 16$ N4 j4 {8 s/ A6 V
) D/ h% M8 |6 e j % j F; W) f. p6 i# B/ n
# 我们将重复使用该种子(因此在动画 GIF 中更容易可视化进度) 4 R- W/ n/ N1 |seed = tf.random.normal([num_examples_to_generate, noise_dim])! c) }( p: C* w" w! \/ }
训练过程中,在生成器接收到一个“随机噪声中产生的图片”作为输入开始。5 j+ e4 q! S9 B
]$ i& Y0 l/ _& \
* o8 G/ ^$ N* N$ A6 [# r7 t判别器随后被用于区分真实图片(训练集的)和伪造图片(生成器生成的)。 3 c2 x" P8 c6 K G+ M, Y, h) h
7 ]7 s! ~! {' O两个模型都计算损失函数,并且分别计算梯度用于更新生成器与判别器。 7 k; c/ |' s+ K Y& e 4 ~( J; c- X2 @' m/ n) Q# A( A- W5 H4 E4 C: H* C h) `
# 注意 `tf.function` 的使用1 L% _* e8 e7 a- P ]2 n5 r
# 该注解使函数被“编译”# x9 U# o+ `& u2 f1 [
@tf.function4 H1 Q: s! }. O; i3 E% a* R
def train_step(images):9 ]! v: H- j- H# h# ]
noise = tf.random.normal([BATCH_SIZE, noise_dim])9 x( b6 |' `+ p6 V; L
- ~; ~. G+ N) f% v
with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape:9 y" e# ~3 `, z) F* f- \
generated_images = generator(noise, training=True) - X d- X: \; q) A ) H, E/ j3 N2 o2 T3 X real_output = discriminator(images, training=True) ! ~* v) e+ ?, r; Z6 ?1 e: g7 U" P fake_output = discriminator(generated_images, training=True) a2 x: G3 L a, @ - _! X2 G) ?% ]% E# _1 L5 Y gen_loss = generator_loss(fake_output) 1 M2 b+ @/ v/ g* Q) K( q disc_loss = discriminator_loss(real_output, fake_output) , ?6 v6 B7 Z4 ^: ?$ }$ t0 J+ } [- v9 ` V1 r( s
gradients_of_generator = gen_tape.gradient(gen_loss, generator.trainable_variables)7 T' M5 I) r) B' |/ ?: J; f. }1 p
gradients_of_discriminator = disc_tape.gradient(disc_loss, discriminator.trainable_variables) 5 X& p$ ~( U$ B) d, W8 `; b* Y ( Q4 b7 P3 f3 o
generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables))6 F4 s( a: k' E( j
discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables)), z4 _! R, A3 Y. d
9 A' G6 z6 r' C) O# E: K x1 idef train(dataset, epochs):6 i. g ^1 }% E
for epoch in range(epochs): & w, ?2 i& T# U3 B! d: I start = time.time()1 ~ z0 x# F: f/ g/ q" k, z
J* W. K( q$ t
for image_batch in dataset:4 j+ B+ z, E4 k! B1 \9 s$ E
train_step(image_batch) 5 v5 B4 }- G2 L2 j5 U9 j6 ]9 w8 W 9 F) Z7 Z+ l' J# C) W: D # 继续进行时为 GIF 生成图像4 s: l. A' }, s9 W8 ~
display.clear_output(wait=True) " V8 B/ c) }0 D# ]; W6 D, N6 q! S generate_and_save_images(generator, N3 K# R1 l% f. }; D( I
epoch + 1,; z, i3 H9 ?# z4 E3 z& ~! g
seed): a3 G6 }* \& t4 j
9 u: Q6 c1 R9 v H! b4 R0 f; w0 F
# 每 15 个 epoch 保存一次模型 9 F" k. S$ x8 h if (epoch + 1) % 15 == 0:/ W g" ? y8 D$ m# V- T
checkpoint.save(file_prefix = checkpoint_prefix) : A/ q3 e2 f' Y6 e" g# B9 E: p " m/ Q" K+ y' }. r( d
print ('Time for epoch {} is {} sec'.format(epoch + 1, time.time()-start)) : f4 @: v5 F2 V2 L7 l ' w+ @* N. `0 E7 ]: s, G1 G
# 最后一个 epoch 结束后生成图片 7 S. X8 } l, ? T: G display.clear_output(wait=True) , Q1 M2 G! G& }! G) o3 w, z generate_and_save_images(generator, * ^3 g- x) `$ u, [9 P epochs,; o* C1 N* t! l) X6 Y' o
seed) r( ^ ~/ `* T# j" y 6 R( c' T4 b7 g# 生成与保存图片 5 S& h5 W$ X# _6 Mdef generate_and_save_images(model, epoch, test_input):' ?( u P m% X! Q7 m6 X" @
# 注意 training` 设定为 False * H4 q& N) r# W" l2 x # 因此,所有层都在推理模式下运行(batchnorm)。: d* O7 I1 D) v4 G7 e
predictions = model(test_input, training=False)& r& s) A$ h* m
, t+ R+ J5 i- }& n7 m \
fig = plt.figure(figsize=(4,4)) + I! ~! {- T/ e. F 8 ?& F" S# W3 `4 I/ U# P
for i in range(predictions.shape[0]): 7 L+ k9 X1 t+ b& J/ ^3 Q4 m plt.subplot(4, 4, i+1) 9 i+ c% d0 }: _2 v+ ?( ^; ? plt.imshow(predictions[i, :, :, 0] * 127.5 + 127.5, cmap='gray') 7 ?5 z# } `9 [$ g% m# R plt.axis('off') 4 o# k# i- l5 E' m 4 t4 h% F+ i! i) r
plt.savefig('image_at_epoch_{:04d}.png'.format(epoch)) . U) H5 T# O3 f: ] plt.show() # c: n4 q `) l8 |- v5.3 训练模型 6 G' m$ w! k# E调用上面定义的train()函数,来同时训练生成器和判别器。 , d. v1 C. O, c" f$ Y3 i3 k6 J ' Q3 A* P$ ^4 O- ?5 t9 I) x6 \0 d6 o) y# W" N
注意,训练GAN可能比较难的;生成器和判别器不能互相压制对方,需要两种达到平衡,它们用相似的学习率训练。 % {; v; N: O& T % |; m4 G% R; E% H, G7 P4 I# s ! i/ k9 V# F! j9 w/ {, Q%%time & N0 ^! V1 ]$ v3 z; ntrain(train_dataset, EPOCHS) : u( N# I' J' x d9 m在刚开始训练时,生成的图片看起来很像随机噪声,随着训练过程的进行,生成的数字越来越真实。训练大约50轮后,生成器生成的图片看起来很像MNIST数字了。 , g6 H& ?! C+ R8 n+ b4 t 4 g, O4 u( s/ R" E" [ 7 R/ S4 y" } K9 r* |" `5 C+ W. g训练了15轮的效果:% f8 \) e( s# a& X* U2 C5 G
* |5 y+ X) [+ D" Q, a# x" B4 B. H o E9 t# p
7 y; r# {: q/ {- X3 H `* o. M( s3 ]) e2 r; h
3 @. E) u1 G9 N7 r6 j
0 t7 H8 G+ A$ s% Y( E& W9 M训练了30轮的效果:9 i0 L) N+ C* I4 o6 ]
5 F( J. z. u4 x
5 I9 x6 t u: g2 W6 C& A; e) j9 l
+ o9 T" }- B6 g8 n$ O: C) P2 L
& [8 I8 H# c. X0 @! M ( _; x& ~8 K* U3 ~' R! ` $ ?" ]6 n1 X' I6 N& W. U) r3 y训练过程: 4 x1 p& j5 R }: c& j# W& [' A3 w. }/ J( r
! |$ t, k7 x! y
; A. e q5 R2 j1 ~8 O, M 2 T" V' \+ ^/ u8 ?' k( k+ E- ~4 r2 j0 d( R
8 t# s, E8 M" z/ ]% L恢复最新的检查点 ; b, S/ b- C# }8 @4 ~- U2 D5 U; g `! a) b$ u
1 N8 g+ x: h' \; l" V. Mcheckpoint.restore(tf.train.latest_checkpoint(checkpoint_dir))# A/ b& X0 r3 i& k* c
六、评估模型6 P3 s( s# _8 [8 j* K: W
这里通过直接查看生成的图片,来看模型的效果。使用训练过程中生成的图片,通过imageio生成动态gif。% V+ {8 g# Q" y3 ]7 Z
5 l) d1 {$ ]- o# _( n
, ?0 F0 a }$ E8 v6 F: o# 使用 epoch 数生成单张图片 ( \, ^1 G) P Y, Cdef display_image(epoch_no): * E$ ~: x8 M) S) {+ ] return PIL.Image.open('image_at_epoch_{:04d}.png'.format(epoch_no)) ! H {' [7 D0 E ^9 n, w& ?+ d/ F 4 B' N' U$ f" F( A% U5 P% `
display_image(EPOCHS) 3 T* f. y9 V3 O, \7 @0 z% uanim_file = 'dcgan.gif' - Y6 Y; F/ M! k0 i$ C8 C5 s# o ! w2 d) ~2 Q) z4 `$ B% R; Z' {with imageio.get_writer(anim_file, mode='I') as writer:* o' y. p4 \8 I
filenames = glob.glob('image*.png')) ^9 }( l& P# Y$ l7 a6 v. Z, }
filenames = sorted(filenames); E! A1 w4 \, p9 C/ k! k$ ~, K
last = -1- O: c0 e; M! K' ^% {5 W
for i,filename in enumerate(filenames): 9 j6 K! T1 i9 _! K frame = 2*(i**0.5)+ n5 b! t7 ~& G" c) Y( I. }+ R6 X( B
if round(frame) > round(last):" B; {5 D! z. k0 V
last = frame! J! n& Z) ?) I4 g
else:9 |* e6 q0 I Q! o0 G
continue 7 T W# @" C8 A2 R/ l image = imageio.imread(filename) ( p8 ]7 L9 |. ?7 P writer.append_data(image) ) o; j. h9 R' A: c0 U: p+ h image = imageio.imread(filename) 4 ]1 n* V$ M V( n writer.append_data(image)! Y6 H3 p: U( \
9 P/ ]& V" i$ O7 o1 Y2 Himport IPython * F; J6 |2 |$ `: Z+ E, Iif IPython.version_info > (6,2,0,''):+ o b9 W9 h, P- i, m) G+ H/ g
display.Image(filename=anim_file)4 J& S. \8 o, j# N, t
! @- z% W3 \/ @; a9 a
, z! P5 S* E: H6 p: q1 L2 d9 ?- P: T! |
: Z& `. @# ^( @1 t8 e
完整代码:* e- \& Y; X$ J* p2 P
3 r6 C; R5 g: I/ I# L" C) B
8 s! N+ H- I3 r4 y! h w& h
import tensorflow as tf$ c- k# w& `+ i7 q1 N% l$ d
import glob' [6 F3 P6 V9 d1 S5 m& m* W
import imageio2 u4 o7 S8 X' D# x) p3 T" }
import matplotlib.pyplot as plt7 z: V A, {$ M6 P, a# T9 P% f
import numpy as np 4 @" H; f! H$ _5 P7 T! gimport os 9 l$ \& v$ S& M9 P8 z; E% Pimport PIL + j0 S" _$ P% J/ @- T8 P+ yfrom tensorflow.keras import layers - ]- L: I/ S# T: }import time5 l! r2 p: o+ J$ B: S2 H/ {
/ \+ r" k: v) t/ y
from IPython import display8 w+ f0 X: a0 M