& H2 Z- p' `3 v+ m+ m7 Ggenerator_optimizer = tf.keras.optimizers.Adam(1e-4) ( X" T9 J" K* o. A. `4.2 判别器的损失和优化器 8 k# B! f) k+ m& D1)判别器损失% @) k3 }2 M6 S* x
: K% b. C) V0 [$ G; y- y" q; | ; V2 L5 f0 h, Y8 y* h' e" _1 T判别器损失,是量化判断真伪图片的能力。它将判别器对真实图片的预测值,与全值为1的数组进行对比;将判别器对伪造(生成的)图片的预测值,与全值为0的数组进行对比。3 M. e! z9 H# ?' R1 t
3 |) }* n0 V" ^0 E
, g. S5 z* a0 L7 s5 _def discriminator_loss(real_output, fake_output): 0 i( \ p4 t+ [3 f# H# \ real_loss = cross_entropy(tf.ones_like(real_output), real_output)5 j$ h& m. D7 X7 C5 i
fake_loss = cross_entropy(tf.zeros_like(fake_output), fake_output)& n8 B( {" K" g" {$ J% O
total_loss = real_loss + fake_loss ( y% {5 S! A/ e. w1 s return total_loss 8 o* k3 u3 q, N; z( u4 e2)判别器优化器 6 r. p4 Y& s% e" s. D3 W$ G D$ d& v: w
% I: W# y- y2 i/ d( F( t
discriminator_optimizer = tf.keras.optimizers.Adam(1e-4)6 L! V" X) S- r( K
五、训练模型- d1 j' r6 j- w, Q' ^
5.1 保存检查点% r. T1 |+ _5 `
保存检查点,能帮助保存和恢复模型,在长时间训练任务被中断的情况下比较有帮助。; ^- ~: H4 H8 m3 Y. B
+ B3 q% Z5 Q+ Q3 s X' V% ^9 c }8 R# i! Y8 t' O3 M
checkpoint_dir = './training_checkpoints' 6 i( i- i( u; ]+ m' }checkpoint_prefix = os.path.join(checkpoint_dir, "ckpt")* v% y5 ~0 f* ]" A' V
checkpoint = tf.train.Checkpoint(generator_optimizer=generator_optimizer, * U" I+ Y! I+ M) i- ~6 @ discriminator_optimizer=discriminator_optimizer,* g& A' I- \3 \6 u+ }3 D
generator=generator, " ?$ f3 t9 T; R# W m! n discriminator=discriminator)1 t; Q' E# N/ [$ a5 m
5.2 定义训练过程 3 P% v& c2 ?" ^9 hEPOCHS = 50' g g7 B3 d& h) K u
noise_dim = 1000 r/ y- u+ _; b# B7 X
num_examples_to_generate = 16 7 W3 Z% S* c6 M- B+ P3 T! A 5 Q, Y; W9 i" y5 g6 y/ ^
, n& s* N. y6 U3 f& t
# 我们将重复使用该种子(因此在动画 GIF 中更容易可视化进度) & S+ X5 }: |) i1 Y5 fseed = tf.random.normal([num_examples_to_generate, noise_dim]) ) J3 L" e5 f. J7 ?训练过程中,在生成器接收到一个“随机噪声中产生的图片”作为输入开始。 4 k' y, e6 \" ?+ p 2 h6 f& M9 [* B- x* ~3 s/ ~6 ^% E, e7 [
判别器随后被用于区分真实图片(训练集的)和伪造图片(生成器生成的)。 + B# A) d) g( U+ K# `$ z! @! v) ` , A; m9 d1 C0 w5 f, Y- V; z- }! G2 p* P$ s/ k, U
两个模型都计算损失函数,并且分别计算梯度用于更新生成器与判别器。6 R1 H" G5 z2 @* L8 I8 k) W
3 @& j f5 n+ f0 v
* d) ]/ J7 k Q/ a" \) t7 q# 注意 `tf.function` 的使用 " z: H7 q, |/ y# 该注解使函数被“编译” s4 S$ o8 g! l# M" d@tf.function1 X2 S! v& X. ^9 u M' ~; {: T
def train_step(images): ! T% |1 D9 ^1 \6 F; g noise = tf.random.normal([BATCH_SIZE, noise_dim]). b$ k3 `, u: u
, P$ Y9 s( {7 k0 f' {4 U with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape: ' G6 l- y }8 a6 @" y1 L generated_images = generator(noise, training=True)5 l2 d+ U# K. g3 O5 b3 V9 \
* j4 d# ^2 l& [6 i/ r real_output = discriminator(images, training=True) ; ~7 M2 q# k- v7 i fake_output = discriminator(generated_images, training=True) 8 q# R) F0 a* R# A4 [. J8 P+ U/ J' K" m# O ' y" V4 v( s. t gen_loss = generator_loss(fake_output) # z. a4 M% }. ^ ^( \1 a disc_loss = discriminator_loss(real_output, fake_output) r4 m8 r4 s# f
' o1 m* ], a# u) F/ q( A
gradients_of_generator = gen_tape.gradient(gen_loss, generator.trainable_variables) - p2 \% t$ \' T2 a: K: l5 k gradients_of_discriminator = disc_tape.gradient(disc_loss, discriminator.trainable_variables)) F; h$ y0 v0 L- e" f* b& k" f6 d K
( F7 x, L" o; b& T0 E+ G generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables))$ z+ m* _7 E' @* _- F9 K) i
discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables)); E4 ^& f u/ a4 B2 ]) p" r1 `
0 l& ?7 X) V$ f& m- y
def train(dataset, epochs): 2 p1 a/ M1 T, t6 r; s w for epoch in range(epochs):% `2 g0 S2 M) E6 s: t
start = time.time()3 W$ P5 w$ ?. r' _, h, [" m( V0 {* _5 a
8 s! P; h9 M6 G) i% n' p
for image_batch in dataset:9 W- w# N$ D! i# `" i8 G1 b! b
train_step(image_batch) " Z- F: t& y! F3 r, q9 E ' p5 I0 l; N8 i3 j# e# q # 继续进行时为 GIF 生成图像 u5 Q& e. Q: z7 z display.clear_output(wait=True)/ M3 X6 k4 K) b) F" f% f
generate_and_save_images(generator, / z8 y5 d2 |( o$ N# y9 f: | epoch + 1, 6 I- m# [( |$ {' I; l; v) G1 O0 I: ~ seed) - i+ e* T* [$ g" c! m2 G ' `+ z( y, @6 l
# 每 15 个 epoch 保存一次模型6 i: ]# t1 l; z
if (epoch + 1) % 15 == 0: ; H% Q. Y& s; @& ^# W checkpoint.save(file_prefix = checkpoint_prefix) $ H. n' _. z( ^! V# I4 f1 J3 ] ) \6 J& M# T' o+ N" S" A: f* g
print ('Time for epoch {} is {} sec'.format(epoch + 1, time.time()-start))9 f6 M( J: }! o6 e: ~+ m$ Y6 G0 k
; t4 j) ~$ R8 i$ J8 a6 ?( y # 最后一个 epoch 结束后生成图片 + W# H( ~/ ]* T8 @ y! t7 { display.clear_output(wait=True) ) d3 B8 X8 l" q' H; e) S generate_and_save_images(generator,7 V3 v; @; N# Z$ i, _+ v
epochs,+ J6 X( Z% u+ ^. O2 S/ B
seed) 7 ]* s. V0 l! I) Q. o. w 1 w( v9 z8 A6 O/ Q' j: ^
# 生成与保存图片 ' d& J( {/ G) s( udef generate_and_save_images(model, epoch, test_input): & A* U% D+ V: t n" ?- | # 注意 training` 设定为 False / x+ l6 D5 J8 w' P: c3 s N # 因此,所有层都在推理模式下运行(batchnorm)。 1 X4 j+ G5 s7 K- p predictions = model(test_input, training=False) 3 G6 h" l5 G5 h$ C 2 s' E5 S: V5 [9 y7 p, m: {
fig = plt.figure(figsize=(4,4)) / q& D ^8 Z c% _9 s6 a 2 h3 A4 w; g# a; C4 d& }
for i in range(predictions.shape[0]):7 ^* @ j+ X; u& w
plt.subplot(4, 4, i+1)9 z. k/ z1 f& e e, P% q* w7 _& W
plt.imshow(predictions[i, :, :, 0] * 127.5 + 127.5, cmap='gray')' i7 q# U) L/ q* f+ H- \2 f9 l$ k7 c
plt.axis('off')! |+ D5 M2 l* h3 I6 @