数学建模社区-数学中国

标题: 深度卷积生成对抗网络DCGAN——生成手写数字图片 [打印本页]

作者: 杨利霞    时间: 2021-6-28 11:54
标题: 深度卷积生成对抗网络DCGAN——生成手写数字图片
+ {+ t9 m% e# n& B
深度卷积生成对抗网络DCGAN——生成手写数字图片
: B' V9 ]: R8 N* J* i9 k前言
$ \' m5 Y, A2 l: V+ h, t/ ~  f* h本文使用深度卷积生成对抗网络(DCGAN)生成手写数字图片,代码使用Keras API与tf.GradientTape 编写的,其中tf.GradientTrape是训练模型时用到的。
+ |$ p7 ~; c4 J/ z# ]2 C# }: E2 F' D9 w2 @; L( \

/ t2 Q! ^9 A2 e: @: E5 c 本文用到imageio 库来生成gif图片,如果没有安装的,需要安装下:( n8 u6 z0 ]. k9 f) K
  F' a% I' b2 H: o
% X, G# M1 M) U3 G
# 用于生成 GIF 图片7 h* o9 S3 S2 k- ?2 O4 U! K9 J/ X
pip install -q imageio$ f+ T0 p$ G1 Z7 [
目录
( A2 V6 V. k5 n9 j! a
6 u3 w4 g$ T" n/ F# A- z% I

1 L+ N# _  Q( n  P, V0 _) U前言. K$ y" u4 ]' Z/ l
6 l, Z1 f; S( p; P& S* Q/ \

) X; e) b! Z: H$ Y1 N一、什么是生成对抗网络?
+ h% E6 a3 G& s7 \" A0 p; _7 x' C2 q0 l- y+ y

$ {, F+ j3 \' |# c二、加载数据集
' v  d; C7 Y7 O$ T3 q, ~$ O- _  ^0 d; |% c" [
& N5 S( M7 Q9 J' e' q
三、创建模型
$ B, {7 C# g: P: k1 e2 y- I- n+ j( n) H3 \
/ L' a3 ^' Q8 p3 b: s+ t! `8 B$ Z
3.1 生成器* \6 G6 ^, e; G9 e

: I! ^3 e' h0 [/ n
; c+ s- r  N( q, e0 p; D* p" D
3.1 判别器
0 P4 ?8 s0 C9 y+ z# Q" [
' H  ~, [+ A: _0 V8 S; d
7 Q/ U( A% E, l: B8 w0 Y
四、定义损失函数和优化器
+ M/ c( g# n* _: ~, |7 T0 a. G" g4 F! |8 Q  a

" f( R. k, S. a  Z- W9 I$ v4.1 生成器的损失和优化器; A6 @, H1 y8 e, ~4 b

. a3 e, c, _: p5 l

4 G0 U9 X+ Z7 N! |4.2 判别器的损失和优化器
# F2 c) J0 |* q2 x
+ I4 F7 O1 V1 \# q( u! I
, Z3 y+ e8 T4 g" Y
五、训练模型  M, J( B* N6 o' Y; b! C# d
( r: u& B/ u9 t9 b+ w3 h3 S

9 H/ p& Q) R* E& r1 ^5.1 保存检查点; o3 V- Q" L2 U+ ]

  {  t, f; n! o" r# L& ^: b8 V

0 M3 o7 ^" n+ P( {5.2 定义训练过程
; W  A9 p5 i: F1 g: ]9 m
8 C5 s1 q1 N: Y+ Q. \

7 V+ X& @$ A) ?$ V) u5.3 训练模型9 |$ L1 Z# c, O
; e7 _4 Z2 L% ^8 ^, w+ x
1 h5 b8 |- T  M1 H# o
六、评估模型
+ x# p" ?6 g! j( j+ D* t" C" c- @+ G5 [

. J8 w9 z: X( ]" J" a, l) v3 g一、什么是生成对抗网络?- b3 [) g+ C- i
生成对抗网络(GAN),包含生成器和判别器,两个模型通过对抗过程同时训练。+ q* B, N+ n- l9 [6 \9 b' m

0 m, J2 G, ~5 |: H. w
9 V" n2 r5 P7 _- \+ A0 _
生成器,可以理解为“艺术家、创造者”,它学习创造看起来真实的图像。
$ p, I0 B1 U* j" q2 O. Q
: s, e. b1 y# o" I
7 p# |7 H: |% Z5 L
判别器,可以理解为“艺术评论家、审核者”,它学习区分真假图像。
  i  A" o( l0 J3 q4 R( H. C4 r( `
8 ~4 h# x/ |/ T: B0 i; U' y

; V& M7 h! S5 R. N# Z训练过程中,生成器在生成逼真图像方便逐渐变强,而判别器在辨别这些图像的能力上逐渐变强。
: T, l: j. N5 u9 r1 ~! \! P" i: X# G  p4 ?! N& h6 ?0 U
2 \* c* _' S, b* H. E8 U+ J8 N7 h
当判别器不能再区分真实图片和伪造图片时,训练过程达到平衡。
7 J" w+ H  I% ^  ~" J! i/ A1 G+ O' D4 U

! |: z4 u$ L9 [4 K* u本文,在MNIST数据集上演示了该过程。随着训练的进行,生成器所生成的一系列图片,越来越像真实的手写数字。1 V* }$ m% A, X, R3 p* ?( B
0 k5 O  B" x9 o
  Z/ Z' t- i) ~; x. J/ {
二、加载数据集
  L/ \! d- T* C( u5 Y" k使用MNIST数据,来训练生成器和判别器。生成器将生成类似于MNIST数据集的手写数字。, L* ~' b, \: t, q' x

. u' h* A# ]# n+ L- W

1 P/ x+ L; P4 X+ ?2 _1 x8 N(train_images, train_labels), (_, _) = tf.keras.datasets.mnist.load_data()6 Z) W) c* i$ \$ Q

( Y: g* o& t- |( V" o5 Atrain_images = train_images.reshape(train_images.shape[0], 28, 28, 1).astype('float32')$ l) O8 L8 T: d+ ]: P
train_images = (train_images - 127.5) / 127.5 # 将图片标准化到 [-1, 1] 区间内
" _0 _$ g( y: Z# ?6 I ! }5 |, ?& ]- C
BUFFER_SIZE = 60000
" C# \- `) c9 z( U+ ~& B. @BATCH_SIZE = 2562 U, w7 Y6 P" T! F8 r5 W# A9 X) s

, _" p1 I+ k  x# 批量化和打乱数据
* A' y: X2 s0 j3 qtrain_dataset = tf.data.Dataset.from_tensor_slices(train_images).shuffle(BUFFER_SIZE).batch(BATCH_SIZE)
& I/ N( h- Q& F% u, l三、创建模型5 _5 k9 X6 F' l! Q" r2 A: s4 }
主要创建两个模型,一个是生成器,另一个是判别器。
( V/ I: B5 g/ q8 D" M8 \0 x6 r# B7 Y
5 ]) C- Q1 {2 j$ s

6 P0 @, ], L1 Z3 Z3.1 生成器! V. A4 r! @& }
生成器使用 tf.keras.layers.Conv2DTranspose 层,来从随机噪声中产生图片。3 P5 v4 Q* h0 H# U  \2 v1 T  U4 O

! {4 \8 _  L; R

" }) q1 U. A, v8 f/ ]% `* k然后把从随机噪声中产生图片,作为输入数据,输入到Dense层,开始。
: z  u' [4 l. Y% E3 ?9 M4 p, B, P' O3 P4 O, X9 p. w1 [/ ^! _  M

0 i& g, F3 t, |8 c4 f后面,经过多次上采样,达到所预期 28x28x1 的图片尺寸。
( }# Z% I. C' N/ f# n, v0 c2 V( o4 b' N' ]  ?7 ~+ c" @
" a! j1 U; W" t) T1 L: A  J) i' v
def make_generator_model():. I' _/ U& Q5 u0 G* k* o
    model = tf.keras.Sequential()
# J( r& e! g. b4 n    model.add(layers.Dense(7*7*256, use_bias=False, input_shape=(100,)))0 u$ ^$ N8 w  L) h
    model.add(layers.BatchNormalization())
4 Y+ h2 o9 o0 P* L    model.add(layers.LeakyReLU())' _8 R0 D6 s# x! T! @
2 F$ L2 K- w6 b9 y
    model.add(layers.Reshape((7, 7, 256)))
- B+ X5 ]6 i* `    assert model.output_shape == (None, 7, 7, 256) # 注意:batch size 没有限制2 {1 w4 u0 d  \/ [" y, k1 @
: h4 j1 Y9 \9 u4 n( R: y
    model.add(layers.Conv2DTranspose(128, (5, 5), strides=(1, 1), padding='same', use_bias=False))
0 d. M8 q5 @9 ]4 [: F    assert model.output_shape == (None, 7, 7, 128)
0 p9 m% o; L' s, K5 A+ R    model.add(layers.BatchNormalization())
2 m( J4 j1 P6 b8 E; w. K- B% O% C    model.add(layers.LeakyReLU())
( m- v  R. g9 F+ L2 L3 ~" ] 7 ~) Z/ X- u5 u. x. S
    model.add(layers.Conv2DTranspose(64, (5, 5), strides=(2, 2), padding='same', use_bias=False))
1 r+ _" g, L- ?2 t8 q    assert model.output_shape == (None, 14, 14, 64)
8 _, Y3 E+ K% X0 e    model.add(layers.BatchNormalization())
2 C7 {' |4 n) X1 R    model.add(layers.LeakyReLU())
# A' T, r8 U# }  p2 f2 K* C * j$ h4 Q+ N# ]8 ]1 a0 t+ U
    model.add(layers.Conv2DTranspose(1, (5, 5), strides=(2, 2), padding='same', use_bias=False, activation='tanh'))
/ c9 a7 V" t& v5 O    assert model.output_shape == (None, 28, 28, 1)
& R  m9 f6 G3 h7 m2 c5 a+ [- s 5 b& j% S5 }! t. O/ J& ?  ~
    return model
  \! _& w6 Z" M# }( h# s用tf.keras.utils.plot_model( ),看一下模型结构% V; }, W$ r; n8 s
- Q- Q* Q0 B. q! Q  m

9 v8 @( A$ d0 P8 E  V) O% n
9 ~! M/ E! z# V7 J7 _% d" G3 X1 h. ]7 @* K/ C$ U
, w4 E7 V/ [& \' D0 }
用summary(),看一下模型结构和参数
% l/ e3 G- F' \' J5 G, m" y' Q- Y- |1 W9 X& x

& @; {" b! r7 A. _/ k5 ?" [7 s! `5 I, s: Q5 Q0 B( C9 I! z/ E+ J

. X) Z( u3 ?, d# H: D! a9 o+ D  ^) c- A$ M0 L4 }" O5 P( s
+ C, f. ~4 [# E7 a' s( W5 s
使用尚未训练的生成器,创建一张图片,这时的图片是随机噪声中产生。
- L9 P) `! g  q! O4 r/ y0 l$ B: L6 P( ]1 l& f1 _: C  u! m
* a- u& ?4 U% _! S, _0 E
generator = make_generator_model()
% c; W8 q0 p. e. r. t8 F 4 `( G; r( J3 g; L( k5 `8 l0 I
noise = tf.random.normal([1, 100])
3 p" r; P4 n. Pgenerated_image = generator(noise, training=False). A3 Y) P5 [7 x" D( z! r4 A
! {+ v: T6 _+ j3 m( F' U
plt.imshow(generated_image[0, :, :, 0], cmap='gray')+ k: R, h5 j/ _3 @
3 `+ `/ f, u! b6 \4 b
: y/ w9 z( R. e. U1 K4 h: w$ c
8 `4 G. h- c  N! D
: S0 c8 T7 s/ g3 |% D4 e# q
3.1 判别器" j( l) q7 v2 [/ h7 O$ n
判别器是基于 CNN卷积神经网络 的图片分类器。
" G- i# K/ }2 f3 P  E
% e8 p  I3 v, V& A0 l: I+ [( F

' i4 Z: e3 A& Edef make_discriminator_model():
( ]/ f& z' Y& I' O4 v    model = tf.keras.Sequential()0 H8 Q7 l( X! ?4 J6 E& s6 e- h
    model.add(layers.Conv2D(64, (5, 5), strides=(2, 2), padding='same',
) q2 T9 u1 J6 d  J7 ]' Q9 F; O                                     input_shape=[28, 28, 1]))
( B7 M% p1 @/ _) U6 |    model.add(layers.LeakyReLU())- U2 Q* o% S8 `! \0 U8 _
    model.add(layers.Dropout(0.3)); ^# A4 o- f: x1 S7 f3 }
5 [: w1 \) C: }
    model.add(layers.Conv2D(128, (5, 5), strides=(2, 2), padding='same'))
8 Y9 n% E4 |( J  K    model.add(layers.LeakyReLU())
, _% V0 a0 C  ^8 \4 H    model.add(layers.Dropout(0.3)); x9 k; k9 D+ v
, y5 F! w8 z2 ]0 a2 C" o  |
    model.add(layers.Flatten())/ ]! J; \$ S* w8 P
    model.add(layers.Dense(1))7 {& x8 F. M, v# v3 ]

$ [, m- d) n& o; \0 u    return model
3 r+ m/ Y+ w* @' y: _用tf.keras.utils.plot_model( ),看一下模型结构# f8 d' e: w" g: z  C

' Z, N; x- X* _6 s9 E2 c: K3 f3 @
) L0 t8 s$ ]& A8 O8 Z

( n4 v6 t) K: B$ y& a

# g9 H. d) P6 ?' W( {
* g  R, ~% N4 ?! P$ D" ~  h

4 z3 Z0 h* `9 d+ L: G% ~用summary(),看一下模型结构和参数
" Y7 E6 E3 o9 P+ J$ M5 C/ j: ^( L2 y# c! I
& k: n. j9 e# K  p! y% y
# w( R8 g* G; D. n- z

7 F. C. G7 I; Q% a9 \% I' v' Q* z+ m1 \: F4 v3 O1 F

* S0 U# Y: s! h, }) S% n四、定义损失函数和优化器
. h' z7 l+ |, |1 ]' F7 r/ U7 U/ s# x# u' _由于有两个模型,一个是生成器,另一个是判别器;所以要分别为两个模型定义损失函数和优化器。
; t- b! [4 X* p8 b0 A2 q- T
: \  j% @8 c1 K% r' v* n  l
6 j& ^5 Q+ |5 ]9 J" L! B
首先定义一个辅助函数,用于计算交叉熵损失的,这个两个模型通用。$ x3 }3 |  `% S

2 L0 D8 \/ U1 Z& E$ \' h* }( [
- A/ V% G: H& m* E" O! g) E  o) ^
# 该方法返回计算交叉熵损失的辅助函数" ~$ o) A3 |0 B
cross_entropy = tf.keras.losses.BinaryCrossentropy(from_logits=True)
+ O2 \) l* H* p0 Q( I: I( I4.1 生成器的损失和优化器3 Y0 D$ R$ U! j  b! f- Z" s; c& x
1)生成器损失. n9 P, p" A: j5 I

9 Q' U2 g: ^0 }+ D& ~6 F2 p

, T7 n& c1 i( q: m* K生成器损失,是量化其欺骗判别器的能力;如果生成器表现良好,判别器将会把伪造图片判断为真实图片(或1)。; Y4 @+ `  [  z

/ P; R% P5 c1 d. z: d' i: s
/ L& W, s2 v. n7 c7 {% }
这里我们将把判别器在生成图片上的判断结果,与一个值全为1的数组进行对比。$ h- B% o8 |. w/ }/ |( D( o1 a

! o" D" j7 v1 S

% b" }" s+ m( @6 B8 udef generator_loss(fake_output):2 y0 k9 r$ c3 `# g2 I7 q" Y
    return cross_entropy(tf.ones_like(fake_output), fake_output)+ o' H9 E9 S: ^2 E5 I: G
2)生成器优化器' s; W+ @: x2 k

0 Z) J  F! g- M3 b# H6 _
9 [+ r8 X2 f1 p9 L+ [
generator_optimizer = tf.keras.optimizers.Adam(1e-4)
# K  v5 B# L. r2 v: z3 _4.2 判别器的损失和优化器& I6 ]$ K0 C; h; U: H+ ^5 o
1)判别器损失
( |. N5 b3 }$ h: X, D0 Q
: U$ ?, w& V1 i. c" d( U

/ `& p; a; p! _8 c9 U9 y( H判别器损失,是量化判断真伪图片的能力。它将判别器对真实图片的预测值,与全值为1的数组进行对比;将判别器对伪造(生成的)图片的预测值,与全值为0的数组进行对比。3 x8 s+ I5 A7 l( u3 k

) c- F2 v0 e, }# ^2 ~
6 C; V' m  G# M& _/ h0 }4 Z
def discriminator_loss(real_output, fake_output):' L+ \& e& w2 ~  ~7 z
    real_loss = cross_entropy(tf.ones_like(real_output), real_output)
2 x  f3 }' C) X7 k! m; y    fake_loss = cross_entropy(tf.zeros_like(fake_output), fake_output)
& _' k0 |2 j+ A9 T, y! p4 z) r! @; X. P5 n    total_loss = real_loss + fake_loss
3 U! y8 ?$ Z8 m+ t+ M/ R    return total_loss+ O) u& g2 G6 l$ y$ J& B
2)判别器优化器
3 w2 \/ _, Z5 _) |; k6 L- q; m% A+ B0 t
7 I* o% J+ w2 L# b  N/ ^
discriminator_optimizer = tf.keras.optimizers.Adam(1e-4)
( R% b; A0 M/ Q9 b2 D! Z五、训练模型* |2 `9 @) X/ O6 @
5.1 保存检查点
6 k" h" N! S1 x; q5 T7 p8 L5 x' o保存检查点,能帮助保存和恢复模型,在长时间训练任务被中断的情况下比较有帮助。
) b6 y; y! ]; V+ N# f/ a
8 X% g2 O0 c4 K! Q, G. `: S$ V
8 d- K$ \! C* P9 {* E& u
checkpoint_dir = './training_checkpoints'7 N# _+ B9 E+ L6 y, B: W
checkpoint_prefix = os.path.join(checkpoint_dir, "ckpt")
' b5 f" v" _( H  D% ocheckpoint = tf.train.Checkpoint(generator_optimizer=generator_optimizer,
1 u$ Y5 `7 t- i1 z8 R* S9 d2 [4 n                                 discriminator_optimizer=discriminator_optimizer,) ~# e: S1 H) s7 Y
                                 generator=generator,
: b& A0 w7 F& `- N                                 discriminator=discriminator)
4 _7 {% S. f/ |% N2 L" J5.2 定义训练过程
8 S( j$ x  {) VEPOCHS = 50
- m' I" c9 [; }0 lnoise_dim = 1006 T$ r1 R, s3 g9 u* G
num_examples_to_generate = 16
9 ^, C) t/ ?" {+ q7 n5 D  \" f # I, q. y/ e/ |

, ~5 E$ r  K0 @2 `2 Z# 我们将重复使用该种子(因此在动画 GIF 中更容易可视化进度)
/ @  h2 h# V# r6 vseed = tf.random.normal([num_examples_to_generate, noise_dim])
% P7 Z( I, \4 L7 d- A" I) D训练过程中,在生成器接收到一个“随机噪声中产生的图片”作为输入开始。
& n1 p' h8 Q! j, v6 t7 r7 b8 l/ P) Z# y3 N& ]6 p2 W7 Z0 H7 o

( T( m% n; ^; j; k1 D判别器随后被用于区分真实图片(训练集的)和伪造图片(生成器生成的)。
# _: O7 u. Z* _4 t' A- r0 d% v# P( A4 y
7 u% Q% E% `! b. v  w2 P5 g
两个模型都计算损失函数,并且分别计算梯度用于更新生成器与判别器。" r' s: i* _( B8 u; q

1 d4 ]/ y, w: }% C5 q1 b5 h

* [# I; P2 i$ L2 {" o8 t2 C7 b' G# 注意 `tf.function` 的使用9 u: _/ z" Q, Z6 o1 j
# 该注解使函数被“编译”+ h1 a+ ]! e& t
@tf.function
/ |; K- C% N0 Q" T. ?7 Ddef train_step(images):
" R9 C! {( @8 J0 f3 r6 ]7 M; Z    noise = tf.random.normal([BATCH_SIZE, noise_dim])9 ?7 L' b; ]5 f  H
$ y. |" Z0 m- c! f, z- B5 m( `
    with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape:/ h: M6 K  Q, n! _8 s
      generated_images = generator(noise, training=True)4 @4 I. B4 l6 A) b) D7 c4 }
2 l; b4 c1 A3 |( m9 ~
      real_output = discriminator(images, training=True)( n) R4 Z2 N/ ?
      fake_output = discriminator(generated_images, training=True)
8 B+ N+ \: x" t- U: p' g
. e3 H' R6 F/ r- K) d. ^  Q* l      gen_loss = generator_loss(fake_output)  s' J5 z) b( t* g
      disc_loss = discriminator_loss(real_output, fake_output)8 m9 S) B; P& n
5 z8 L* y# u5 K$ `; o
    gradients_of_generator = gen_tape.gradient(gen_loss, generator.trainable_variables)
, R( B' t8 V5 q$ b- s- ]    gradients_of_discriminator = disc_tape.gradient(disc_loss, discriminator.trainable_variables)
, p7 A; r2 V: D( V+ ] ' s& n( A! [4 b- }
    generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables))! u- a4 M2 [$ h, @5 H( ^" r
    discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables))
. C1 i7 A* T* ?# \ 2 j7 ]4 w% G/ i4 @* P$ [
def train(dataset, epochs):$ X1 K0 Y4 Z5 B" M" Y1 M
  for epoch in range(epochs):) R7 R. V; p2 c+ X5 m: Q" @; j
    start = time.time()3 [9 z2 ?! ], u: c& O5 q! ?
0 v  I3 _/ l$ R$ L* J" D, r
    for image_batch in dataset:+ f0 k; b( b- p7 o1 R+ h) H
      train_step(image_batch)
  n# l% p3 c& I+ p + v7 E0 w/ X3 s& q! C
    # 继续进行时为 GIF 生成图像
& R$ `/ u1 {% _3 a    display.clear_output(wait=True)) ^+ |' Z3 ~5 l/ }- o0 k
    generate_and_save_images(generator,  k9 j. H: E  m2 W# R$ K1 O( `2 I
                             epoch + 1,
1 T% U; t0 t. L! Q) U& `/ f' k1 H                             seed)" p0 j) l% q* x  Q

/ S0 X& \% p$ C1 B( t0 |! W# k    # 每 15 个 epoch 保存一次模型
& p7 o# @5 `4 |    if (epoch + 1) % 15 == 0:
6 @' M: U3 d5 a; S      checkpoint.save(file_prefix = checkpoint_prefix)
/ V  l2 b# V/ J! ]
  ?& q0 A. W7 W9 P5 A& K    print ('Time for epoch {} is {} sec'.format(epoch + 1, time.time()-start))8 ?) G* T# w, k1 b9 a1 m) Y
5 P. k  ?  T1 u; l. c
  # 最后一个 epoch 结束后生成图片
* h" a" S, A& H: F* R  display.clear_output(wait=True)7 [0 s& u8 g5 o6 u1 |8 E% ~
  generate_and_save_images(generator,
9 h2 B% ~* Z# L                           epochs,6 D% q# ]- [- i* X6 K
                           seed)
4 r1 e1 ~7 K+ ] : u+ |3 e) `, V$ x
# 生成与保存图片( O, Y$ Y5 d$ @4 g
def generate_and_save_images(model, epoch, test_input):
) u% t5 t2 I* X  u1 ]  # 注意 training` 设定为 False& d" \* s- f' E" z- u. x9 p; n
  # 因此,所有层都在推理模式下运行(batchnorm)。& c+ w" a2 h( w- ~1 D' X0 x
  predictions = model(test_input, training=False)
1 U3 N. ^# d& ^$ z; I/ i5 F
. d- q! q/ a1 l# k1 b" }* ?4 Z  fig = plt.figure(figsize=(4,4))  s; `1 k& Q, S; L& k5 y4 j: o
1 _4 M7 W% h( ^0 ?; p6 e. H  d8 X, s
  for i in range(predictions.shape[0]):
# W% a$ L" |: y' \      plt.subplot(4, 4, i+1)
( N+ p& l) g# }      plt.imshow(predictions[i, :, :, 0] * 127.5 + 127.5, cmap='gray')4 M, d8 u2 O. Y8 {3 Z( |% |3 b% k5 Y
      plt.axis('off')! b5 x( x' s; o' S  N. M

* j8 y. i: e0 X5 S) x  {' ^5 e  plt.savefig('image_at_epoch_{:04d}.png'.format(epoch))
) ^4 f. E9 r+ s" l1 K; |. ^  plt.show()! O- t$ y  t" y! K2 v; h; r
5.3 训练模型2 v! F# s0 b; ?/ G
调用上面定义的train()函数,来同时训练生成器和判别器。2 _& u0 m1 o% L# a# z* p
8 ]) e8 R7 i2 r2 E" z( L
2 o/ v, x7 ]7 }! a  z
注意,训练GAN可能比较难的;生成器和判别器不能互相压制对方,需要两种达到平衡,它们用相似的学习率训练。
3 d( g: p/ v, j$ z+ H8 ^
+ l5 g( t8 b1 w% F2 U" \& U# r, k

0 o+ W; [' i, x3 D%%time
& Z/ e5 }/ b) l+ w8 O9 gtrain(train_dataset, EPOCHS)3 |& p) x2 ?8 |6 X* C
在刚开始训练时,生成的图片看起来很像随机噪声,随着训练过程的进行,生成的数字越来越真实。训练大约50轮后,生成器生成的图片看起来很像MNIST数字了。( m2 A7 a, B! o

2 g3 Z5 H4 ~% D' z' T

" F: a( X/ @5 n训练了15轮的效果:6 c1 X* h5 B# ^; @
- @1 R2 K6 L- }' b) g4 P9 [

5 E4 ?$ }) `6 I& A* i3 u
* Y( o( V7 |# u9 _4 {$ U5 T

6 _2 a! z' a9 ?" a
, L# P7 X1 z. u1 T; x: b# m1 c
$ Z# z( G; s! z  D3 b. \
训练了30轮的效果:
1 u; p. E- {  X- L% y
$ @, x- Q7 d/ ~; u! `- j

! g2 {" z+ }; V3 Q8 I
- z% i; C" T% X/ b

5 ~6 U; x! u8 H. a0 t, v' R5 o+ h9 f, X. A) r# X; _# c
* q- o- e3 r6 ~5 \, \
训练过程:% }+ H/ ^1 D" |/ G- N0 [

8 p2 T/ X$ R/ s5 j6 i% R- p8 U
- e/ H# T5 O1 ~$ t
( K- q* A1 n8 C/ c4 U9 \

% u, \7 N/ C' R# M# s+ C1 B7 \* C- g2 S; }: |

' a: R$ M, L2 k5 \恢复最新的检查点3 D6 m) n  z9 X, }4 b

% {0 Y3 ?  V+ M; j7 U2 L  M

1 x( Z- X4 v7 X9 M' d5 dcheckpoint.restore(tf.train.latest_checkpoint(checkpoint_dir))
/ ]. M0 J- f# Z6 w六、评估模型0 V" e3 x" l9 s$ Q7 T% x; I
这里通过直接查看生成的图片,来看模型的效果。使用训练过程中生成的图片,通过imageio生成动态gif。4 I, u" o, H" t6 r7 z% Y

$ y3 ^$ t' i, X

+ G$ j7 _: ?( x# 使用 epoch 数生成单张图片" o* J' `1 V5 w1 y/ a' F. g
def display_image(epoch_no):
+ W% T) @' c4 \6 @& h  return PIL.Image.open('image_at_epoch_{:04d}.png'.format(epoch_no))
$ Z( \2 |3 Z' p9 V; m. j 7 W( @) Q& r6 @6 B7 k; ?1 F
display_image(EPOCHS)
* D$ x' u- }9 Z# Z% B1 H8 [anim_file = 'dcgan.gif'4 J9 E! i. z0 o7 s; [9 Q' D

$ T7 Y' r/ R4 h9 \; y  }3 wwith imageio.get_writer(anim_file, mode='I') as writer:
% F0 x) W0 i( Z. R, g  filenames = glob.glob('image*.png')
2 M; I9 V- N+ G2 W/ L  filenames = sorted(filenames)
) `  E7 O+ W/ t7 B0 |& d0 d  u( q  last = -18 ~5 g% i% o3 e& X6 u& a8 m  B) e
  for i,filename in enumerate(filenames):: C8 k  {  h4 }; Y
    frame = 2*(i**0.5)
& X/ Z) j- @  K9 {6 ?) k4 c    if round(frame) > round(last):( y3 |# a4 Y8 N( z- y! b. M
      last = frame
& J$ u$ k' Y, d( e    else:
  R4 k/ A2 P9 U, d3 W5 `. N      continue* }+ e* X4 ]( g" A
    image = imageio.imread(filename)
4 h( Z9 [1 v* p& d3 T    writer.append_data(image)
7 Q  i8 u- [' m* G. a5 U. Y6 p  image = imageio.imread(filename)
) C9 l; L# X9 F! r& m" _, d) ^8 Y  writer.append_data(image)+ k' z5 p; |' N' Z
: B% a& W( c' c  Q6 W
import IPython
7 Q2 ]3 R6 O+ K& K+ V! o6 N. Eif IPython.version_info > (6,2,0,''):
! P0 i2 W, s# M0 D2 |2 N  display.Image(filename=anim_file)* f/ u; F1 E& M+ `7 l
: w$ M3 I* o5 E2 n! p, ], g% ^

0 F7 Y/ {) p3 v2 U! z( ~5 F, a1 n% w) d- z) k7 h' o4 k* U

# b/ Z# e9 {+ z( y+ T' W完整代码:
( K% @* B  F5 D, }* A  q2 }0 a
8 E  f: P' d; |. @

0 g$ x! g* B0 f1 c  v8 qimport tensorflow as tf
. I; x0 m( x" R2 ?7 L/ Vimport glob
: z! O$ l) D; {& p  r& _9 M9 Wimport imageio  p) ~9 l2 F. [# ^1 b
import matplotlib.pyplot as plt6 _" n* H2 l3 H; u) \
import numpy as np" o$ U9 t) O% Z; ^2 H: g' m6 n
import os
1 b6 [) c9 p( c1 ^3 zimport PIL
/ V1 m% z1 z* B1 hfrom tensorflow.keras import layers% P8 h, _! L5 j0 b, I/ Q
import time
/ Q  |: Z/ V6 L4 G- Z ! F8 H: S% R: B- H, \" n9 o! w1 R
from IPython import display4 b* t0 l1 [0 L8 ]% ~% {

- R" {. N# `% ^2 i1 [, M2 S  J2 T(train_images, train_labels), (_, _) = tf.keras.datasets.mnist.load_data()
2 {7 l5 Z0 v& K1 D4 J+ I
9 x8 {. I, k% W/ etrain_images = train_images.reshape(train_images.shape[0], 28, 28, 1).astype('float32')
9 M* j2 ^. u' z  Ttrain_images = (train_images - 127.5) / 127.5 # 将图片标准化到 [-1, 1] 区间内
" B, \" I2 v! k5 R2 I  `" l+ d( G
) k$ n& f& A0 b7 c' s7 fBUFFER_SIZE = 60000* o: T3 v7 S$ Z' s3 G7 Q0 E5 V
BATCH_SIZE = 256# Q9 F' U( b% ?$ Y0 o$ z
1 ~" L% u; A+ O5 W& h
# 批量化和打乱数据
% x& Z! `2 t  w7 ]% G4 wtrain_dataset = tf.data.Dataset.from_tensor_slices(train_images).shuffle(BUFFER_SIZE).batch(BATCH_SIZE)3 W1 H8 n1 x: Q
' E7 e, R3 B4 a  R% N7 ]( j$ N0 I
# 创建模型--生成器* a, p, T! n. a+ u
def make_generator_model():9 g+ ]  ^4 j* o) A8 J7 H9 q
    model = tf.keras.Sequential()
: ]8 O& g8 f: I8 b0 \    model.add(layers.Dense(7*7*256, use_bias=False, input_shape=(100,)))
8 a- }# q8 I/ s2 g0 X3 Y    model.add(layers.BatchNormalization())8 Z2 C1 N+ y* f& @7 u
    model.add(layers.LeakyReLU())4 Y2 U$ B+ g. \8 ~% K4 B

. d; Z  N+ ^& |1 d3 `7 c) f6 a    model.add(layers.Reshape((7, 7, 256)))
5 h$ m- C! r0 A. c) W/ C    assert model.output_shape == (None, 7, 7, 256) # 注意:batch size 没有限制  @- x7 _/ k2 _! M6 ?) X
0 i+ v! c6 \' V" d* L# ^
    model.add(layers.Conv2DTranspose(128, (5, 5), strides=(1, 1), padding='same', use_bias=False))
8 N: M0 u, F0 T, }3 k' y9 [    assert model.output_shape == (None, 7, 7, 128)6 j" k% x' J3 K
    model.add(layers.BatchNormalization())5 H5 X/ g0 S" z/ l
    model.add(layers.LeakyReLU())8 o# ^; Z) {& q4 Z2 G5 G, q: `

5 g3 r8 j- H5 {1 F    model.add(layers.Conv2DTranspose(64, (5, 5), strides=(2, 2), padding='same', use_bias=False))) M( @8 Z4 x5 r  L5 q% W* i6 i' e
    assert model.output_shape == (None, 14, 14, 64)
! h2 C7 e# k' r8 j    model.add(layers.BatchNormalization())
* p% \# d! `# s3 K: m    model.add(layers.LeakyReLU()), W% R0 k( {  k9 K

4 h3 f1 h6 U1 O0 k: t; V    model.add(layers.Conv2DTranspose(1, (5, 5), strides=(2, 2), padding='same', use_bias=False, activation='tanh'))6 m+ G- F4 \6 o3 {) b) h
    assert model.output_shape == (None, 28, 28, 1), e# X; m$ L) N) m
3 W/ d5 ~0 |: M( A+ [. R
    return model
, \! x1 H0 R( s5 U0 G
/ c$ q! t) A  b! X& @; |) U7 j# 使用尚未训练的生成器,创建一张图片,这时的图片是随机噪声中产生。
/ A( `3 F% [! K5 ]6 G" [2 Wgenerator = make_generator_model()
: d* c; ~* G  Z, ^1 x
# K0 h) j1 C, I3 G/ @. ~' unoise = tf.random.normal([1, 100])9 G# U- ~& V1 u
generated_image = generator(noise, training=False)& @( n; ^- t4 H8 Q2 n  Y, h, }: w; M
" X" m* T  Z1 D3 g% P
plt.imshow(generated_image[0, :, :, 0], cmap='gray')' v* C* f4 L, P
tf.keras.utils.plot_model(generator)& A' W3 X7 d" N
- n$ I5 {* _+ Y% J
# 判别器1 j* K* X3 o) A% a; j/ l1 E
def make_discriminator_model():6 O  {# r. @4 p9 m$ ^' v, x' N- v
    model = tf.keras.Sequential(); t% L0 v' q# {" C4 i& v) ~! Y  \
    model.add(layers.Conv2D(64, (5, 5), strides=(2, 2), padding='same',
2 q/ r; k0 ?3 J0 k                                     input_shape=[28, 28, 1]))  M; R6 g" Y! s3 W; Q+ r
    model.add(layers.LeakyReLU())
; D" G# e' U8 q! w( a" J+ Y    model.add(layers.Dropout(0.3))) X  m$ e1 @; X2 h& U- x. m
# {1 a) [7 {8 ~1 Y9 U3 \
    model.add(layers.Conv2D(128, (5, 5), strides=(2, 2), padding='same'))* N8 K  J4 S) W9 R
    model.add(layers.LeakyReLU())5 I$ K) r. T' G; ~( S9 ?- u
    model.add(layers.Dropout(0.3))
+ i! G+ x1 E5 J" |/ t 1 Z/ z7 S7 D& f( V  p4 }7 r- q
    model.add(layers.Flatten())6 F$ E9 K( E3 W: Q( S) i
    model.add(layers.Dense(1))
! v; _) c6 {( Y( h% B  u - @- J4 V3 G% R) V
    return model: x  C- v4 v. B& _* W0 V1 T

/ x" \4 T) z" e( [+ w# 使用(尚未训练的)判别器来对图片的真伪进行判断。模型将被训练为为真实图片输出正值,为伪造图片输出负值。
( j' `. B6 I7 z# {& c( Udiscriminator = make_discriminator_model()
$ y9 n8 E3 w4 v- I; m8 edecision = discriminator(generated_image)
8 j! l! M1 B, r1 }" T& uprint (decision)
& k% P/ t3 g! h. A3 s- \
" N3 e# z( v% m! {! |" K7 E, [4 T# 首先定义一个辅助函数,用于计算交叉熵损失的,这个两个模型通用。; d5 M  V; S( L: V* [1 X$ h
cross_entropy = tf.keras.losses.BinaryCrossentropy(from_logits=True)
" z2 z6 k2 {  @; g) a
6 `. M% x- C. w, B8 o. A# 生成器的损失和优化器
2 h' W8 F5 b1 Tdef generator_loss(fake_output):4 I* l2 D, r/ ^" O8 y7 q# r
    return cross_entropy(tf.ones_like(fake_output), fake_output)+ N6 _& k- T% _  E
generator_optimizer = tf.keras.optimizers.Adam(1e-4)
  i  E# U2 j5 k2 O
+ q9 A  K3 a  @# P  L% L8 _# w  j# 判别器的损失和优化器
) a% Q! l: ^; ^  q: |! N% B  adef discriminator_loss(real_output, fake_output):* u. H+ j' E7 a: b
    real_loss = cross_entropy(tf.ones_like(real_output), real_output)
- D7 w! p8 r( O3 P" f6 m% q. Q    fake_loss = cross_entropy(tf.zeros_like(fake_output), fake_output)
) V5 y) i  |* r6 A. S% l    total_loss = real_loss + fake_loss1 I, r5 p  X% o  f3 W+ d  O" R
    return total_loss9 w8 v( }- A/ l3 T  g, T! J. ]
discriminator_optimizer = tf.keras.optimizers.Adam(1e-4)( o9 E& g8 x' x2 {

# n0 D; {% t- u; o# R, c# 保存检查点$ a7 Y/ D. Z$ C
checkpoint_dir = './training_checkpoints'
: b) H* ~6 [! Y9 z3 {# z' ?checkpoint_prefix = os.path.join(checkpoint_dir, "ckpt")
$ k# T  O! c" u+ _) c6 O% F$ X' Jcheckpoint = tf.train.Checkpoint(generator_optimizer=generator_optimizer,
( V& O9 t( V/ Q5 o0 n                                 discriminator_optimizer=discriminator_optimizer,
% n' j3 w0 i0 s, n                                 generator=generator,- f! t* F6 b, T, X2 |! u% P5 i' F
                                 discriminator=discriminator)* s" H+ O$ T1 L3 X) J! ~

# C0 n* Q. w. l; G5 ~# 定义训练过程
0 T1 @$ m* N+ ^/ f- CEPOCHS = 507 m8 i8 J: y+ q! W7 U
noise_dim = 100
( J& f1 R! ?. e. D8 X: i: q6 P7 _) lnum_examples_to_generate = 16
  n( _- V2 f9 {$ o8 G  ` & x: u; `; Z) W/ P
# 我们将重复使用该种子(因此在动画 GIF 中更容易可视化进度)+ T+ S0 s! f0 @/ o, j/ H- T' b
seed = tf.random.normal([num_examples_to_generate, noise_dim])0 K4 @, Q1 n, `5 y& x7 D' L

5 H1 K' T3 [# n1 }2 E. h) O0 N# 注意 `tf.function` 的使用
) t7 l" |+ J8 u, H/ ]# 该注解使函数被“编译”
0 q! k9 ^# ?, K4 w@tf.function( D9 S' x* \' c' r- ^: w" s
def train_step(images):2 {; k/ f2 Q$ e
    noise = tf.random.normal([BATCH_SIZE, noise_dim])
4 P! K$ V: J9 z* E9 |
+ ^. s; w& n' A9 j1 H4 I3 h    with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape:2 M" M2 I2 R' |# o/ ?; O$ ~- x+ a
      generated_images = generator(noise, training=True)7 [& O9 A, v* i
9 V5 ?2 [4 f' h! j. j
      real_output = discriminator(images, training=True)
# ~5 A1 c. m- V8 i5 C2 q: t' I      fake_output = discriminator(generated_images, training=True)
5 \5 S+ Z% U5 ~1 c3 k   e/ b; `, i; Z) P$ z
      gen_loss = generator_loss(fake_output)5 ?- y' ]# I5 R, k  I
      disc_loss = discriminator_loss(real_output, fake_output)
7 y/ C% x. ?2 ~# ^3 W - g8 f3 x/ |- B( X* l0 n7 w
    gradients_of_generator = gen_tape.gradient(gen_loss, generator.trainable_variables)
  I( H* M+ A$ e$ L/ f+ m    gradients_of_discriminator = disc_tape.gradient(disc_loss, discriminator.trainable_variables)
6 z; u+ n1 H- Y
3 N  {5 f8 t9 X    generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables))
( U+ M9 ~7 N6 Y    discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables))
; d  a- ]8 l/ E+ [. a; w ! _1 g5 c3 D! H+ C
def train(dataset, epochs):
$ `1 J+ J) t3 F' l: s  for epoch in range(epochs):
9 R9 a1 P4 U# P+ F- J& \( _( F    start = time.time()1 M3 C: f' ]2 w9 F( U) t! a6 z
) H2 L. b$ ?" w) v
    for image_batch in dataset:
$ z. Y: R& ]! z1 W0 E      train_step(image_batch)/ x' s' D- W: w6 L7 n3 ?, X! V

. `4 j( Z: O% a! Z. w    # 继续进行时为 GIF 生成图像8 O5 S  h, x2 {, u  K' J
    display.clear_output(wait=True)7 a" ~5 C7 R6 F3 H
    generate_and_save_images(generator,
- b/ H% K! u/ b8 z& G4 l                             epoch + 1,( A5 t% y* _: w' k' T2 I7 S& p
                             seed)
  I9 Y3 L" M1 T+ j4 z" J
' b! ~) ^1 {/ [, q3 W" Z, ~! [( B6 Y    # 每 15 个 epoch 保存一次模型' G$ r. v7 v+ X4 ?5 L6 E2 y
    if (epoch + 1) % 15 == 0:0 V7 X% v& Z! f4 o
      checkpoint.save(file_prefix = checkpoint_prefix)4 V  J" ]3 X' b. T5 O. f" T1 [
& c/ V1 P7 r" v5 F, a# f$ a2 q, _! t
    print ('Time for epoch {} is {} sec'.format(epoch + 1, time.time()-start))
, C3 `" \; D$ r) z$ B / o: D( y# {+ `! ?
  # 最后一个 epoch 结束后生成图片
. i: @9 C$ w: `$ u8 F8 P; M+ N1 Z  display.clear_output(wait=True)/ r$ x3 o& U8 K* i8 K
  generate_and_save_images(generator,
2 y% C# s7 I6 f, I, {3 m                           epochs,
# [; n, _. [/ b4 {+ o                           seed)
3 B1 d  b/ e; N' D7 d6 [
; ?! ~1 y, j4 W1 n2 V6 ]1 x# 生成与保存图片
7 o$ o# b* u2 |9 N! |; [& \3 jdef generate_and_save_images(model, epoch, test_input):
' x' L+ C% F( ?. N5 C  # 注意 training` 设定为 False% z' ~7 y" Y' X" K' Y
  # 因此,所有层都在推理模式下运行(batchnorm)。' V0 ]/ ?- ?$ v9 F# t# V9 q
  predictions = model(test_input, training=False)
4 S. _7 s* e8 [* L  _! z 2 t0 e- D8 ]4 I
  fig = plt.figure(figsize=(4,4))$ U) z, u, I5 u0 z) L  x# M

$ Y& B* U. b' _' P  for i in range(predictions.shape[0]):
! D9 j. V+ Z" s* L1 x      plt.subplot(4, 4, i+1)
, c2 v  ~, C7 e1 W- P: Q1 k* f- C      plt.imshow(predictions[i, :, :, 0] * 127.5 + 127.5, cmap='gray')
) Q" |) Z+ }  k( Y, A% l, }$ m      plt.axis('off')
4 _: w, Q( D4 d+ c  I  ~5 O ; t3 t; Z9 q# l4 D
  plt.savefig('image_at_epoch_{:04d}.png'.format(epoch))! u* l' }& y, E  ~( R: h* `
  plt.show()
. \0 n  y# U8 C+ g
; W7 \+ |7 m% Y% ?6 g1 B$ G9 R# 训练模型
0 L- F- d4 E% n8 ]train(train_dataset, EPOCHS), \% z" S" z) m$ [2 J$ @

& I% j1 m% C; [$ @# 恢复最新的检查点2 I9 f0 m/ [  g( Y  y
checkpoint.restore(tf.train.latest_checkpoint(checkpoint_dir))
( a8 h! N& |+ e( p2 Q# m* s
1 m  ^, D" m9 G) z( }# 评估模型
, M# T: [3 t- O4 X# 使用 epoch 数生成单张图片
4 C/ _, \+ W, k3 M# xdef display_image(epoch_no):# h; T% B4 S4 Y8 ]! @+ T6 [
  return PIL.Image.open('image_at_epoch_{:04d}.png'.format(epoch_no))
8 t) C- r8 s% D # ?% M) L) G1 o3 M# s0 C
display_image(EPOCHS)
6 P5 _, ?/ {  l 3 K  R1 ?4 L3 V0 m' q$ f
anim_file = 'dcgan.gif'
% B& c" N! T0 o3 J, g : O/ R- l0 K6 r" }
with imageio.get_writer(anim_file, mode='I') as writer:! {2 c9 x& N) u! B6 Q9 L
  filenames = glob.glob('image*.png')1 e' k$ g  A: S3 V: {; u/ \
  filenames = sorted(filenames)! y) Y1 f4 j- B! ~& C
  last = -1
0 }* w- p4 x! ~6 ^+ b  for i,filename in enumerate(filenames):
; w& I; P" p9 l- p4 `- D+ o    frame = 2*(i**0.5)
0 Q; I. A1 s( M& m4 y' H$ V    if round(frame) > round(last):
$ Z: i0 C% |8 \- W# E( P4 P1 f      last = frame" Y5 `- D" y- X1 e* ]% p
    else:
( U) C0 f. L, ~9 q1 S8 ~3 j- t$ @      continue2 Q# j  e) o+ E5 q3 I
    image = imageio.imread(filename)/ H% g) f; L9 P3 ?3 v5 u5 D% R
    writer.append_data(image)  `9 A& q0 L  w* r7 B' M$ @
  image = imageio.imread(filename)
4 ^6 s, O' R& J9 ?2 ?  ]$ p& [  writer.append_data(image)
5 i0 K; z8 g: w2 |6 @1 L , }2 a: D3 r. T$ r) o
import IPython5 H8 s' U; W- i7 `: b* ~( I
if IPython.version_info > (6,2,0,''):
; f4 D# s2 ^6 g3 h3 `  display.Image(filename=anim_file)
$ f& f8 s% W/ x* E  n  P/ g" Q参考:https://www.tensorflow.org/tutorials/generative/dcgan
! q1 ~- i1 ^* J5 {* |' R/ j————————————————
1 F+ b% u8 i' B! ]* B0 z版权声明:本文为CSDN博主「一颗小树x」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。
7 }- Y1 G( V) q! c原文链接:https://blog.csdn.net/qq_41204464/article/details/118279111
( N+ t' z: c! d1 u8 {! P( L  L/ q' y

, Z# b1 a+ j* P- h$ B$ U, K' Y$ P




欢迎光临 数学建模社区-数学中国 (http://www.madio.net/) Powered by Discuz! X2.5