数学建模社区-数学中国

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

作者: 杨利霞    时间: 2021-6-28 11:54
标题: 深度卷积生成对抗网络DCGAN——生成手写数字图片

) y9 C0 J/ j6 j! F1 w0 L2 B: p深度卷积生成对抗网络DCGAN——生成手写数字图片+ r0 n9 O4 q! X5 Z9 X' J( {- C
前言. ~( U; R7 r1 s9 W
本文使用深度卷积生成对抗网络(DCGAN)生成手写数字图片,代码使用Keras API与tf.GradientTape 编写的,其中tf.GradientTrape是训练模型时用到的。
- a- E1 }- v4 Y8 F" G. F
1 [6 R: r; b2 x& L. [/ v
8 R# w8 [; j6 x9 d/ Y/ Z6 _- m
本文用到imageio 库来生成gif图片,如果没有安装的,需要安装下:/ H' I5 U4 G/ }3 [/ ^  y9 ]
+ m+ u9 H2 k& h
+ S8 B! v/ E8 ?/ y: H8 c
# 用于生成 GIF 图片
1 e0 _3 }7 V1 a5 i. Wpip install -q imageio
, {' M/ W& H! {目录6 q- j- G# y) u- z3 a8 p
& N9 c0 {) K9 Z7 j
* m! G% q( j# D- x9 J
前言
, E: X! m0 M; f2 R. T* ^! w% Q6 R1 |( Z) m# ^

. x4 F: W. T% r+ X# Q3 e8 x7 L一、什么是生成对抗网络?
4 w. O, u" N* Y4 j8 ^9 `7 M, w) m6 r& r) \- F

% Q7 E  ?0 @8 }- Y二、加载数据集
  B& x4 a' O8 i* N! i* f
7 h6 A/ x4 N/ V; o& u2 v# G: v7 z

" g, ?3 m* h* `$ G4 i三、创建模型0 \) @4 c$ t& Z+ d5 }6 b; T/ Q

$ J( k' N7 @8 n* q6 Z
4 j" i" f2 b5 {" A) ?
3.1 生成器# L1 A; O+ X' q6 Z( G4 t) g
; y, [( m, X7 d9 k6 d. i

0 b9 U& F" a( S3 w3.1 判别器
. ?5 Y' \( j3 L" D' d6 P+ x+ q- |/ m2 r0 V% y; K. N6 c
: h1 y9 ]" b2 y
四、定义损失函数和优化器
' t5 L3 u3 o1 |2 m, K& R, a) n* ^
2 S' u4 E; J8 q& j- Q, ^& X( Q
: Z" Z7 T. Z, j( M
4.1 生成器的损失和优化器  b& ~, e0 _5 K: ]# z: X: _
& k1 w+ `# P! a& v3 \1 d
8 c" U4 S4 ~. s$ {# m- ?8 y# r. ?
4.2 判别器的损失和优化器
6 g" D0 |7 t/ p) V# [( u" o! y. P! F5 O; }6 H. ^$ @* O7 P
3 F" n: F6 ], T9 w4 q  \
五、训练模型
/ Y: H; j/ n% S! S" l( K" ^2 n: k% q6 v: u, T% z% V

+ r) H- `5 d% `; ~- \* A5.1 保存检查点" A" p) o6 v1 _. H, M9 h% ~
) i- v' O' w! {  d

! q0 G2 M" V7 ^. w( S1 j+ h$ W5.2 定义训练过程! S7 A" ]) g* a* n

! P/ N2 t6 k  p- E, H
) m6 |/ Q% d6 [& j: V# F
5.3 训练模型8 k" ]- ]: k: i/ b; |

; N- c* S: ], I! M6 {

; N7 ~7 x4 U" i$ X% Q% e4 m六、评估模型
. ~# _% y4 @# }& P' t4 |4 _# h0 b* i6 K' ~9 \1 Z, L' h# [1 y
1 T: i" j: s6 Z' I' b4 u$ w
一、什么是生成对抗网络?, Z; r: v/ k/ V3 ?6 V5 F. O
生成对抗网络(GAN),包含生成器和判别器,两个模型通过对抗过程同时训练。. ]0 q' P; X( a
) R7 o) m0 p, ?! L; L
  G/ z+ R! b" g, ~+ H3 V8 D
生成器,可以理解为“艺术家、创造者”,它学习创造看起来真实的图像。
* j/ L6 t6 ?, U# R& B  |$ \- v* }6 S% C

% w2 u7 H& W& R判别器,可以理解为“艺术评论家、审核者”,它学习区分真假图像。
& x. d) q" @3 U, @$ J- @( K' n) t) z5 c$ i3 K& V9 O& ^. \
7 t) Q9 z. X% ~/ k+ u
训练过程中,生成器在生成逼真图像方便逐渐变强,而判别器在辨别这些图像的能力上逐渐变强。9 \; S8 C( |# J) R, l7 T
2 Z9 h# Z4 Z7 `0 k

* m# B9 F6 w7 @+ k5 a( P- @2 n当判别器不能再区分真实图片和伪造图片时,训练过程达到平衡。! T  e) O5 ?  Z+ C! ~: `  M6 K
4 T, ~0 ]- e0 k

% `/ J4 M. z7 `9 \$ s本文,在MNIST数据集上演示了该过程。随着训练的进行,生成器所生成的一系列图片,越来越像真实的手写数字。
9 l& c: X+ e" y% g' Y
3 O2 Y) N+ a; T2 s" a7 r
. n2 D6 ^% o& @3 _, o4 z2 f
二、加载数据集
/ [% d( L2 n2 S2 b' P5 K5 b使用MNIST数据,来训练生成器和判别器。生成器将生成类似于MNIST数据集的手写数字。
' B  `, m/ [) ]; }3 C, C& f/ {" o4 }4 ^, z

. g+ L: i. K1 N7 m2 s, E(train_images, train_labels), (_, _) = tf.keras.datasets.mnist.load_data()' E3 k, H7 [2 P$ g* }- o0 Q$ X

' }% _0 n- _9 y% Wtrain_images = train_images.reshape(train_images.shape[0], 28, 28, 1).astype('float32')
( g& j; Y6 t& i! S6 j, m1 utrain_images = (train_images - 127.5) / 127.5 # 将图片标准化到 [-1, 1] 区间内
: }: Y7 M% t6 ^: B: {6 R- Q
/ q6 r  W& b- m* e7 G# m7 @8 IBUFFER_SIZE = 60000- t% d" `1 \" R8 t' Y) a' Z
BATCH_SIZE = 256/ j% a3 _3 |. @* f
' a5 t8 E9 _1 j8 x1 M+ I
# 批量化和打乱数据
# }& G$ P: \* F7 qtrain_dataset = tf.data.Dataset.from_tensor_slices(train_images).shuffle(BUFFER_SIZE).batch(BATCH_SIZE)
; D3 k# i% ^. w三、创建模型  _! U; l3 J# d% B9 G
主要创建两个模型,一个是生成器,另一个是判别器。
" |* f* [8 l, _' k- i8 h9 R7 O! r2 @' v7 V( j& X
1 D, ]0 Z4 ^" \3 l
3.1 生成器
: s- ^" K  W3 s  j7 O4 |0 v# W生成器使用 tf.keras.layers.Conv2DTranspose 层,来从随机噪声中产生图片。# Z1 B1 X" `) h9 [. B3 l
; U* K# n- e5 J3 O5 q" L# b

- _4 d8 {# b/ f+ o% y" S& |8 ~* a1 f然后把从随机噪声中产生图片,作为输入数据,输入到Dense层,开始。
1 E/ t% P: Y6 Z0 I0 _: Z
( g) E5 x, E0 o: J3 I6 L
$ Z3 A2 z3 K7 V4 q( H
后面,经过多次上采样,达到所预期 28x28x1 的图片尺寸。' [: f) A' ~+ c" x+ x: \  f4 T
6 v2 o; \  D' ~9 @" w
& F( a7 J+ Z8 _$ @, y
def make_generator_model():( C* @( p+ N4 M- T
    model = tf.keras.Sequential()
) x/ |/ Z; `) E7 r( s$ O    model.add(layers.Dense(7*7*256, use_bias=False, input_shape=(100,)))
$ S# L) s/ j) J- P4 ^    model.add(layers.BatchNormalization())
! ]( g2 X* b7 |: e- x    model.add(layers.LeakyReLU())
: O6 h5 H! y0 \6 {
: T; u$ b# G" k6 p- V) `0 g) I    model.add(layers.Reshape((7, 7, 256)))4 V5 i) C4 d* |7 u
    assert model.output_shape == (None, 7, 7, 256) # 注意:batch size 没有限制
# j5 E. G' c2 A) m! y $ U& Q& `% G- @! A! g
    model.add(layers.Conv2DTranspose(128, (5, 5), strides=(1, 1), padding='same', use_bias=False))& M, K7 \! m0 ^' ]) i- E
    assert model.output_shape == (None, 7, 7, 128)
9 n) S7 r4 r! p" w  t( T1 P    model.add(layers.BatchNormalization())9 L# Y/ u/ q% U$ U% [
    model.add(layers.LeakyReLU())1 F- N* z* M" P8 S0 x* K7 s
$ M/ P; ~9 _9 E& b5 E
    model.add(layers.Conv2DTranspose(64, (5, 5), strides=(2, 2), padding='same', use_bias=False))! q5 R  I4 Q3 O$ L  m
    assert model.output_shape == (None, 14, 14, 64); u# p5 U. l5 ~+ y7 R- ?
    model.add(layers.BatchNormalization())
6 t4 _* ~8 n8 z7 n( r    model.add(layers.LeakyReLU())
- p7 P5 j1 J1 B & q) z5 F; V! g3 P9 U, _  D. ?
    model.add(layers.Conv2DTranspose(1, (5, 5), strides=(2, 2), padding='same', use_bias=False, activation='tanh'))
7 v, z: O: y9 k& d/ U: Q    assert model.output_shape == (None, 28, 28, 1)
* d. `! ~  Z$ \/ E  V) ]
/ c0 r# L  c& V9 c& u$ d    return model
- r8 X2 f; G! m8 a用tf.keras.utils.plot_model( ),看一下模型结构
4 D6 V. f5 g5 Y5 x9 P, @0 K9 ~0 N" o) J# b. t
% T7 Q( t& a" w1 ?5 Q
# V# [/ ]: Z) ~

8 H0 O0 D- w" Y- ?
! @' C: S8 X/ N
用summary(),看一下模型结构和参数( y# M( }2 k" ]! u0 S
( U) j2 V! M3 v
6 o2 b. }6 \  v8 K; s
, K4 @' V) v: |* G7 x

) ?1 l! E8 \) |/ Z! i
6 X( x% |8 k( H2 q% E, B

% E# R- w. V2 Q) J使用尚未训练的生成器,创建一张图片,这时的图片是随机噪声中产生。
8 a) X; q0 i7 i5 n, {% y3 ?% v# e+ F" p0 I7 v9 O4 T

* R& X6 k2 e$ ~3 \1 |5 w, Igenerator = make_generator_model()
' s" E+ ?2 @" Y# Q1 P0 I- d4 L 9 J  n& q0 L, d: o& S/ S$ ^
noise = tf.random.normal([1, 100])+ L' e) _8 m2 G( U3 t
generated_image = generator(noise, training=False)& L$ g# N. w3 s6 S
( [# ~0 x  }8 `& R) E
plt.imshow(generated_image[0, :, :, 0], cmap='gray')
% E3 n; w8 E. k3 V! ~# C9 N' K( H; T( l4 L& {% E6 h+ i0 H

+ C( R2 G* |* J6 ]( O
) E5 l% `2 b2 o+ f$ ]- T- A
8 x# a1 `7 I  T1 R
3.1 判别器5 Q+ r) c' U+ z9 F! q- i
判别器是基于 CNN卷积神经网络 的图片分类器。
, ~  {2 l9 e* P9 R# i
/ ?" w/ \5 i/ O* p7 A/ W( {
8 s, w: O' l8 L2 y$ C, R! f0 @
def make_discriminator_model():" b; z9 S1 X# U8 u: i' `
    model = tf.keras.Sequential()
$ K! b, P; i+ U$ e    model.add(layers.Conv2D(64, (5, 5), strides=(2, 2), padding='same',9 ~- f# `$ s! f# F: J, D! E
                                     input_shape=[28, 28, 1]))
5 _  C# ?% k7 D2 m; v) S7 J    model.add(layers.LeakyReLU())
. f' n9 K% v8 S8 w    model.add(layers.Dropout(0.3))) J/ h2 C4 V6 M; A
6 u' I! O: X- C" r/ z  }
    model.add(layers.Conv2D(128, (5, 5), strides=(2, 2), padding='same'))- x+ \. Q& `9 t! U# o& R
    model.add(layers.LeakyReLU())
- s7 i1 {" Q  \. y    model.add(layers.Dropout(0.3)); w$ V  H/ }- f9 i5 M
/ i+ r8 \$ r2 p' R/ s. p
    model.add(layers.Flatten())( M2 Z0 ~+ V2 k& x. c
    model.add(layers.Dense(1))
8 ~; F, g: b# Y8 c 1 b4 L/ l& \5 P. I
    return model' X- G; M1 t1 B; X+ h
用tf.keras.utils.plot_model( ),看一下模型结构+ q1 t# [, L- ^' I

" ?* q, N0 R4 t) }$ m+ `4 O

6 F1 i& t& R% [/ U' a- L
- R2 H2 J1 o* a5 U% \9 D! L

+ ^5 F/ D% L' H0 h( }( ^! O  z# a& ^  B6 D  F% k  [4 q3 A% h$ E
8 Q: x1 \5 ^1 |  A! r1 i5 Y- \
用summary(),看一下模型结构和参数
0 T+ n2 q5 h$ p
3 V1 i! D# ]/ H% ^. M# u2 H
) L5 w  u3 ]! e2 e7 z! \
8 b, M# l' k" P: T! j

6 [, T7 e" l( |( B+ j: v, d
" _" j% Z2 }- z4 I& |/ ^# |
  H& g, W0 q) |- e
四、定义损失函数和优化器4 B8 ?6 O8 [8 n
由于有两个模型,一个是生成器,另一个是判别器;所以要分别为两个模型定义损失函数和优化器。* Q( U) R( i; y' v0 c/ ~  @8 `

9 r0 f3 d  N* j6 q$ i8 a- _7 F( z
- O, s7 S  i( A1 {# P
首先定义一个辅助函数,用于计算交叉熵损失的,这个两个模型通用。
* f, l' }. q) J9 X2 b  c3 H# p
/ w  \5 g+ d- ]9 u! V( R
9 K- _; t0 ?6 {1 ^( e0 ^( d
# 该方法返回计算交叉熵损失的辅助函数
) B2 e' Q" l& m- vcross_entropy = tf.keras.losses.BinaryCrossentropy(from_logits=True)
/ x9 V! Y" v" l* G; O7 \4.1 生成器的损失和优化器3 n0 Z+ \$ m/ A2 \% }
1)生成器损失, I, ^" ?8 h( i$ A

. Q; v4 `7 m+ n; z
5 I4 n) l: C+ b0 g) P! `! [
生成器损失,是量化其欺骗判别器的能力;如果生成器表现良好,判别器将会把伪造图片判断为真实图片(或1)。
- x. c+ Z  C( u4 k8 |9 g* R; W9 A# T

; n0 C  J. F+ \% f6 r这里我们将把判别器在生成图片上的判断结果,与一个值全为1的数组进行对比。& J& L1 w& t3 G" h6 i2 Q: B

! B& O! V, t9 |$ }' `

0 [+ X/ U9 y( M$ l) S  |) [def generator_loss(fake_output):
) |- a8 b0 e" X4 `    return cross_entropy(tf.ones_like(fake_output), fake_output)+ }+ f; t5 i2 R( [
2)生成器优化器. {/ g2 G" D9 h& a; J& q( x( [# O" y
, e0 ?7 S$ w, j$ I
. o3 s0 k5 [0 z6 ?! |' y
generator_optimizer = tf.keras.optimizers.Adam(1e-4)  V, T5 |8 b  y8 U: I4 u
4.2 判别器的损失和优化器( v# K+ g0 W9 [3 ~7 x9 D) N
1)判别器损失
: B  L1 @( S" E
; n, q- p; q* n8 O+ T

9 w& W9 C0 G1 _判别器损失,是量化判断真伪图片的能力。它将判别器对真实图片的预测值,与全值为1的数组进行对比;将判别器对伪造(生成的)图片的预测值,与全值为0的数组进行对比。$ _! B! |% y$ H; D5 K
" F+ l) b, }- ~
1 T4 `$ W( g3 b5 T7 J% @& G; L' z6 d
def discriminator_loss(real_output, fake_output):3 m4 |, q) f* }" q- K0 \" F
    real_loss = cross_entropy(tf.ones_like(real_output), real_output)
3 k- }, L. Q0 g' z9 v    fake_loss = cross_entropy(tf.zeros_like(fake_output), fake_output)
/ r4 i/ c% z! s  y7 p" v    total_loss = real_loss + fake_loss
/ b, H+ Q. d& ]$ G+ [5 R# ^    return total_loss! J3 V7 |$ v5 t$ q6 g# ?8 I
2)判别器优化器9 w; J) E2 J, h* f# W% `: f" }  X7 _
. t5 b  W; A/ v4 g$ b, ?

/ i, a) B7 b4 mdiscriminator_optimizer = tf.keras.optimizers.Adam(1e-4)
5 t; h7 ^5 }( D, Q. F( S0 T五、训练模型! s- \% v* _& _0 q& p2 P  L5 B/ @) u! v
5.1 保存检查点
" l' j3 v- X/ P5 x$ F9 \' k保存检查点,能帮助保存和恢复模型,在长时间训练任务被中断的情况下比较有帮助。
" p$ o( A7 T2 {# Z2 X
" K( s3 A- R# a0 g- ^! G6 f2 ~& o
4 l/ S3 V% I$ i
checkpoint_dir = './training_checkpoints'
. C9 k3 R0 ^6 ]. J; echeckpoint_prefix = os.path.join(checkpoint_dir, "ckpt")# I; Y' b3 C1 ]9 ?* k
checkpoint = tf.train.Checkpoint(generator_optimizer=generator_optimizer,7 P# d. z4 r0 ?1 m0 A+ \# ]7 f  G
                                 discriminator_optimizer=discriminator_optimizer,& B; f( [  @* w% U6 m5 `1 q
                                 generator=generator,
) H. Y; n/ @) ^* Y1 U0 {0 O                                 discriminator=discriminator)( K0 j" V3 Z2 k4 U/ Z- ^1 X
5.2 定义训练过程- `. c% U) h4 z) a
EPOCHS = 50
7 Y- L% m- y% f* x/ O( z% l! Wnoise_dim = 100
# w% N1 D0 w1 W/ }3 q9 tnum_examples_to_generate = 161 ]- D4 b  u" o/ ]

5 R- s& \: m6 e5 h) }; ~
- Y2 A$ t$ a- x' ^$ V# 我们将重复使用该种子(因此在动画 GIF 中更容易可视化进度)
; A! W. i0 ~( v  M. q1 a1 \seed = tf.random.normal([num_examples_to_generate, noise_dim])
1 `# h1 b+ z. j1 [2 `+ g训练过程中,在生成器接收到一个“随机噪声中产生的图片”作为输入开始。
+ F- g% i0 n- s( o7 Q& ]- z; J; K% Z2 o

, P9 y: l. m" U- t7 ]判别器随后被用于区分真实图片(训练集的)和伪造图片(生成器生成的)。6 ^& A$ z' w# i: ?0 e% I! [

1 u. i$ W8 r# }4 y% r

0 t8 R6 J7 T  u两个模型都计算损失函数,并且分别计算梯度用于更新生成器与判别器。
- {! z9 H* ]5 l9 [! P( s1 K6 E, w7 A+ E% m( B3 X: Q
" f( O& z8 y+ H! t
# 注意 `tf.function` 的使用
" F1 ^. [: V# ~# 该注解使函数被“编译”! y1 Q# _4 ^/ D3 \! h: s
@tf.function' I; e0 W: B6 D2 H
def train_step(images):* ]' S4 q/ u# ]6 U+ @3 U
    noise = tf.random.normal([BATCH_SIZE, noise_dim])
  ~* B% }$ l8 f$ p5 {9 e/ P
' `8 ^3 x: m" K# a    with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape:; b4 m- C! Q3 i4 B  s
      generated_images = generator(noise, training=True)9 L1 M) ~6 d0 ~$ F
# b  Y, m  L% T* N" B
      real_output = discriminator(images, training=True)
! E+ ]- W% v5 E9 S      fake_output = discriminator(generated_images, training=True); k: f8 ~9 T. V" s

# }4 {6 G! j8 v8 B      gen_loss = generator_loss(fake_output)
5 A6 b* ?1 e# y0 r1 H0 U      disc_loss = discriminator_loss(real_output, fake_output)0 K/ Y) _$ o- {
; g# w; l! u" z7 ?. D; L
    gradients_of_generator = gen_tape.gradient(gen_loss, generator.trainable_variables)6 Z  Q7 b0 `" p4 c8 J# K2 j$ ^
    gradients_of_discriminator = disc_tape.gradient(disc_loss, discriminator.trainable_variables)
0 a& z. q9 E2 j, M ) T: |6 Z4 I3 S, H' u7 E* k; R
    generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables))4 |( z  E6 z0 O' n( J: k
    discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables))
$ Q3 D+ q- z* M+ l, c1 q
$ x8 R# M+ P  Udef train(dataset, epochs):
& F4 n! {0 c1 s  for epoch in range(epochs):
9 x# |+ O, T! `. n2 N    start = time.time()
, `0 p6 [4 Y1 G# u0 _0 H 4 k- h+ ~) O! T2 A
    for image_batch in dataset:* n; i9 v: }3 \% v; D' j8 M+ \
      train_step(image_batch)! Y6 Q8 H( X' |+ X8 D/ O# [
; _3 l4 o9 o9 B5 X
    # 继续进行时为 GIF 生成图像
2 @: a" E. V+ F3 }    display.clear_output(wait=True)8 d6 D: O# Z3 M( x# }
    generate_and_save_images(generator,! h( }: r# E* i" a. w
                             epoch + 1,  O0 h( H7 L, z' N/ G" o- f" R
                             seed)
2 C" b# r( N( J( s , l3 C* f- t; R7 E* }$ Z. S" G0 x5 a
    # 每 15 个 epoch 保存一次模型
  J) ~" f6 h/ k8 y    if (epoch + 1) % 15 == 0:
% R% ~$ Z5 P1 R, D4 }! m      checkpoint.save(file_prefix = checkpoint_prefix)
7 M6 s  o3 [5 t) c+ S4 | % J8 P) m" l2 v, S3 ?, [) I, Z4 J. O
    print ('Time for epoch {} is {} sec'.format(epoch + 1, time.time()-start))
4 x6 N# k# h) k7 |+ f ! R) H$ I. z, V/ O- y
  # 最后一个 epoch 结束后生成图片
7 @  h  @( _8 ^, `6 H# K8 m  display.clear_output(wait=True)
3 H4 y9 X5 E; D% `/ N  generate_and_save_images(generator,
# f4 P0 B/ [7 h- j                           epochs," N! n  H" C. n5 B% }
                           seed)! m. n( e! ~* Y2 R  n. K/ t! a
4 B1 p# n( h( z2 z
# 生成与保存图片
, U# E& m+ y+ b; B  M" Fdef generate_and_save_images(model, epoch, test_input):5 B5 _$ o& g  ]8 L# ^2 m
  # 注意 training` 设定为 False
4 y1 \, X- l: E+ X2 [. ^  # 因此,所有层都在推理模式下运行(batchnorm)。
8 _3 g3 @/ X: G; C  predictions = model(test_input, training=False)
/ o, H3 \8 j) T! ]0 @ 2 Q7 f- [& L7 R. [' @8 \5 z
  fig = plt.figure(figsize=(4,4))
. S1 h/ O7 k4 E0 Q( g# Y- b" A, j , V, L: |( Z) D  b* ^$ q
  for i in range(predictions.shape[0]):
' b3 g$ h4 l) W1 }      plt.subplot(4, 4, i+1)( @3 u/ |( |" x3 c3 i( r/ D
      plt.imshow(predictions[i, :, :, 0] * 127.5 + 127.5, cmap='gray')& w+ r8 `- }- K8 k' T
      plt.axis('off')
% d% l+ L% M2 q# Q/ E# c
0 r6 w; S. o! U  plt.savefig('image_at_epoch_{:04d}.png'.format(epoch))# A) }. ~. \( H+ E2 x
  plt.show()
. f7 z2 R5 }$ y. r: c! V5.3 训练模型% l' Q0 W4 L5 h4 E1 v
调用上面定义的train()函数,来同时训练生成器和判别器。- ?8 I, j7 P2 F; K5 [( u/ {
/ y2 E8 X5 A, Y/ r0 T

! Y% ]7 n  L+ F  j* @注意,训练GAN可能比较难的;生成器和判别器不能互相压制对方,需要两种达到平衡,它们用相似的学习率训练。
. X+ S7 M$ F2 {4 H
- H: i2 g/ h! p0 n8 h& ^
+ Q. p, Q. B, _& f* ^( z
%%time
0 `' U$ M5 y; v# o0 Q& b8 L' Gtrain(train_dataset, EPOCHS)- M% Y% C7 X# j. t
在刚开始训练时,生成的图片看起来很像随机噪声,随着训练过程的进行,生成的数字越来越真实。训练大约50轮后,生成器生成的图片看起来很像MNIST数字了。+ c$ i' i' ^. N. A$ x  a

+ v/ v# U) j+ Q
  W2 r5 Q6 F% s, U4 d! i( z
训练了15轮的效果:+ b, n$ P5 Z3 h2 {2 F2 }

0 i+ @, z* W( I

. C; E" M3 g0 n6 q, {$ Z7 M  o3 N: a9 ?9 b* R; g( E" I

: I$ ?7 K0 n% F& J  n2 W* w1 n. h/ ]* {& E4 z8 O; p
- F2 W6 N' @' k: y, d& a
训练了30轮的效果:/ t5 A. ]! B* O
/ ^' n( e, `7 y

- D3 H; j- a  T' ?9 i. \1 L; _' U5 t+ s7 C, G
/ r# I! }" s. X- \

, p% c0 Q& P+ D  {$ N" {
7 J+ M/ ?* m) ]  u
训练过程:
# o. m  V( |. N9 |8 @; Z/ k# e2 k7 y

4 r' p% b1 t  j4 r+ o5 l
3 X! F0 U3 O3 B

$ S! c/ q# P1 G' j3 }) d
: a$ Z( m" a* g9 E* w

3 Q: ^3 B0 V+ P/ F( @6 [0 l恢复最新的检查点4 c: U/ H7 X  q; s4 |$ O0 S+ Z
9 d) B* j0 M; E

& \- p8 I: B& @checkpoint.restore(tf.train.latest_checkpoint(checkpoint_dir))
, n# g1 J0 c4 b六、评估模型
: I/ D; p! V+ _' h2 k* s2 o3 m( r2 A这里通过直接查看生成的图片,来看模型的效果。使用训练过程中生成的图片,通过imageio生成动态gif。
$ l/ t4 I( m/ k/ S  r$ \9 {7 U$ S  `

$ ^1 r1 K! s7 f  |: H# 使用 epoch 数生成单张图片+ }3 }1 M5 v( W8 M
def display_image(epoch_no):2 h' r$ b1 q. F, O
  return PIL.Image.open('image_at_epoch_{:04d}.png'.format(epoch_no))
, z5 B8 C. C9 N+ l! C6 G
  ]4 v+ a. Y# f$ T! ]& U3 u0 Pdisplay_image(EPOCHS)  q) j0 l# I. K& l$ q
anim_file = 'dcgan.gif'
2 ?' z9 j' K+ T: ?
, E6 ^9 K0 z! y1 Rwith imageio.get_writer(anim_file, mode='I') as writer:
+ U: x5 B% y9 {* {" x7 t' }  filenames = glob.glob('image*.png')
) H' B0 \" H  [- Y) {- Y; g# {7 z  filenames = sorted(filenames)
, p+ [9 {( S9 E: O$ M1 p& G* F  last = -12 G) k1 k% X8 v4 Q; g
  for i,filename in enumerate(filenames):
- H0 I0 L/ e( |    frame = 2*(i**0.5)! J" l# U6 U/ n
    if round(frame) > round(last):+ M% L+ r1 d6 A/ j( E3 A3 V9 K
      last = frame* Y: I1 }! A9 g% O4 g8 p7 J
    else:( d' `, h2 u( N
      continue. S6 h5 Q3 r& u6 `. s
    image = imageio.imread(filename)
+ w8 _6 _8 _3 W/ O0 Y& r7 p1 K    writer.append_data(image)3 w! P! }, w! _. {2 N/ \2 O
  image = imageio.imread(filename)' @6 o: O$ e4 o* w' E
  writer.append_data(image)
1 j  Z( d$ }( e, N( V1 M ; k. M5 z! ?9 n+ E# j+ j" [
import IPython$ B+ ~$ B5 w) [1 \
if IPython.version_info > (6,2,0,''):
  O$ h! _5 |8 a6 T- j. J$ R7 V: d$ W  display.Image(filename=anim_file)  X0 S8 c" c$ w

# p- x! b& ~9 `

6 M) Z( ?$ w1 h0 R& @$ a$ e& U! b7 \% [: v! q$ V8 X1 a, u3 {
* F6 k# e' h, K9 V3 P2 ]
完整代码:5 O8 A1 O1 S* w# F# j
' n* W( ~+ Z. W  H# m! W
; m  \+ w( k: [+ @4 c& w; j: M
import tensorflow as tf! u# h* ?0 o7 v: x1 X% k2 Y
import glob- E  v, q) y1 Y4 |1 H! ?+ Q
import imageio
% K' E7 M1 Q  T3 N+ o% dimport matplotlib.pyplot as plt
5 R6 T+ R0 V. {: [0 O6 `5 cimport numpy as np+ a  @1 S5 x$ s0 m- {
import os! S' \% I* g2 ]9 h
import PIL
9 V0 R' w9 e) P8 Y3 cfrom tensorflow.keras import layers
6 |3 Q7 Q! x$ o/ K- U7 L5 Ximport time) t( X* w: w8 b) _
8 T5 o! E1 N$ s& R* j. _
from IPython import display, x2 n  c' n3 R0 f

9 U1 Q9 Z6 [8 G6 s% w* v% G(train_images, train_labels), (_, _) = tf.keras.datasets.mnist.load_data()
, M% k# f1 v; n" c/ _! m& T " K: l4 y3 i' [3 y6 M
train_images = train_images.reshape(train_images.shape[0], 28, 28, 1).astype('float32')6 Q1 i7 L( ^# u- ]& S8 M( G
train_images = (train_images - 127.5) / 127.5 # 将图片标准化到 [-1, 1] 区间内' L; n  g/ I) b# O0 B5 G
6 C0 N# ?1 x6 C! T# {' ]$ J  E
BUFFER_SIZE = 60000
! z9 v1 {9 ]0 QBATCH_SIZE = 256
% j" ~8 q; s/ @% |/ }8 u/ M/ k$ K 1 e; x' W1 A$ q9 u* X
# 批量化和打乱数据. U0 s  ~6 g( U  t
train_dataset = tf.data.Dataset.from_tensor_slices(train_images).shuffle(BUFFER_SIZE).batch(BATCH_SIZE)
  `  Z0 }( _( E- W8 m
+ Z% n6 n9 x' A/ N! f# 创建模型--生成器& h# o5 s  |! n# T7 T% h" Z
def make_generator_model():  q$ B8 z" Y* b5 s- r( r6 @, m1 E
    model = tf.keras.Sequential()$ ]# P9 L$ k! Z  c# O# u( R3 A
    model.add(layers.Dense(7*7*256, use_bias=False, input_shape=(100,)))' G& n2 q$ t$ ]# v& G
    model.add(layers.BatchNormalization())  @$ q6 U2 o8 [3 u1 G0 n: V  W
    model.add(layers.LeakyReLU())
5 E! s3 N* h9 N  N/ N7 K5 S
: ?0 A7 l! z2 R' x    model.add(layers.Reshape((7, 7, 256)))
; s8 x& D3 r' |+ I5 [    assert model.output_shape == (None, 7, 7, 256) # 注意:batch size 没有限制$ s( P; E1 C) v# o; y" X& k

# ?6 |. ~* B' K0 X    model.add(layers.Conv2DTranspose(128, (5, 5), strides=(1, 1), padding='same', use_bias=False))
  q- I5 P9 f/ j: T  v, b9 S    assert model.output_shape == (None, 7, 7, 128)
! e4 I! l5 ~# i    model.add(layers.BatchNormalization())
, b# [& ]2 Q+ {! N    model.add(layers.LeakyReLU())
/ E2 H3 F4 I( W- }9 V # g8 F' J0 _( e4 f5 W
    model.add(layers.Conv2DTranspose(64, (5, 5), strides=(2, 2), padding='same', use_bias=False))
2 q+ ~" w* I% u; A0 d* F    assert model.output_shape == (None, 14, 14, 64)! z- A3 Y. I( X) H& z6 O
    model.add(layers.BatchNormalization())  t+ {' l% Z# R
    model.add(layers.LeakyReLU())
. q* |8 ?5 X3 W& V- x3 _, _8 o$ H
8 d* x  a* G- I6 L+ j/ h; F6 K    model.add(layers.Conv2DTranspose(1, (5, 5), strides=(2, 2), padding='same', use_bias=False, activation='tanh'))- e* o$ V8 b* Y- p0 r0 |
    assert model.output_shape == (None, 28, 28, 1)0 G/ h% L( B! Y2 @2 o+ Z3 D
4 w: O2 V, |1 Z9 G/ p! e0 U
    return model
! \9 \: \- _+ D % N. p! F* y% q0 u
# 使用尚未训练的生成器,创建一张图片,这时的图片是随机噪声中产生。3 s8 r$ u& W, R" p
generator = make_generator_model()
! ~4 G  h+ S. d) B% s6 L $ d! Y# D9 P) m3 y& Y
noise = tf.random.normal([1, 100])5 j& [+ i" j- y( _
generated_image = generator(noise, training=False)+ m" {/ Q5 T. |! n. z4 I
: h) r8 X& n! [/ e
plt.imshow(generated_image[0, :, :, 0], cmap='gray')9 O' w4 {( g* \/ w7 n
tf.keras.utils.plot_model(generator)
( H/ E1 T2 e  d) p" ?3 \7 \
4 ~  z' O" L+ k+ g3 F$ [9 g# 判别器
  B8 I' _  w3 o5 G- c; Sdef make_discriminator_model():
9 Y5 B- x) q8 I! u2 @& }  p    model = tf.keras.Sequential(); x/ @6 |2 k* V9 s- K$ ]
    model.add(layers.Conv2D(64, (5, 5), strides=(2, 2), padding='same',
4 i3 n5 n, [9 p0 q7 Z  _                                     input_shape=[28, 28, 1]))
) m5 H4 E3 X  s: C2 e  v    model.add(layers.LeakyReLU())7 J5 }  N% }1 P! c  B( R1 x0 p
    model.add(layers.Dropout(0.3))2 S- t  ~% b# `5 M) T
* o4 p3 U7 S  N* h7 L
    model.add(layers.Conv2D(128, (5, 5), strides=(2, 2), padding='same'))
) v* h2 B: Z% U7 X# R0 ~4 |    model.add(layers.LeakyReLU())" R+ e, ]& ]: z' _! q9 O
    model.add(layers.Dropout(0.3))
# Z3 @- a4 R/ s" k' k 1 d, ^" N7 Q+ p! O, L5 T
    model.add(layers.Flatten())
% y) Z# b5 J3 Z/ i    model.add(layers.Dense(1))
5 h7 Z" l8 k5 B; \1 @
% a% I, N+ f2 {! g* t0 w    return model
4 u$ S9 i: v$ A& L! d ( B; }, K3 O! G6 z: T( t1 `+ t
# 使用(尚未训练的)判别器来对图片的真伪进行判断。模型将被训练为为真实图片输出正值,为伪造图片输出负值。
! i$ f: P& y: |6 Odiscriminator = make_discriminator_model()
9 _: W+ I, H( C3 z; e8 bdecision = discriminator(generated_image)0 \. _# A5 ~. l
print (decision): T, N- N1 m( \+ h
6 @2 E0 N* y0 p% ]7 G, L$ X7 Q# u
# 首先定义一个辅助函数,用于计算交叉熵损失的,这个两个模型通用。
$ c: D6 n5 {! H2 _1 Ecross_entropy = tf.keras.losses.BinaryCrossentropy(from_logits=True); ^0 w( a; c" h
+ Q& v8 r# o2 r% Q! H
# 生成器的损失和优化器
( A5 z; ^! d. V# ]def generator_loss(fake_output):+ J) D7 I1 I' k3 W4 P6 E# R3 O
    return cross_entropy(tf.ones_like(fake_output), fake_output)
7 `2 |! h6 [, W& @generator_optimizer = tf.keras.optimizers.Adam(1e-4)& y1 E( O. x$ g- G& s
# N8 W" t1 w7 o2 w3 ~
# 判别器的损失和优化器3 ]. ?* x, r! o, d
def discriminator_loss(real_output, fake_output):8 P/ b# r# B. u! g( t3 ~$ m% S
    real_loss = cross_entropy(tf.ones_like(real_output), real_output)1 L  K7 D; q- ], _
    fake_loss = cross_entropy(tf.zeros_like(fake_output), fake_output)4 {% C1 q4 Y: C3 `4 y; v: B
    total_loss = real_loss + fake_loss( o% k9 g- D$ V& b+ ]3 O) F( Z7 w
    return total_loss* n& c) C- F* g. Q- y3 c
discriminator_optimizer = tf.keras.optimizers.Adam(1e-4), _6 q0 o9 M+ a5 g# ^) J$ S1 @
; r5 F3 @6 K9 H! k1 N5 G: `
# 保存检查点9 x' [1 V+ d' z3 {& ^, C5 F7 q. L& @
checkpoint_dir = './training_checkpoints'0 Y% R% L+ r0 x  D" K+ v
checkpoint_prefix = os.path.join(checkpoint_dir, "ckpt")$ h) Y9 B1 ~: X
checkpoint = tf.train.Checkpoint(generator_optimizer=generator_optimizer,8 e! S8 m) w  a4 X) F& Z
                                 discriminator_optimizer=discriminator_optimizer,' g# [! B* Z1 ]0 t1 `
                                 generator=generator,
9 c/ }" _0 |! V$ R6 `7 a; K                                 discriminator=discriminator)8 r, o, E) N8 e

/ y3 B% _8 s5 s8 S' M5 i# 定义训练过程% K% H6 B/ ^, ~8 ]6 _! e; N. G
EPOCHS = 50& |  Q$ s& [, H2 g; `! E
noise_dim = 100
  C: V0 O3 M5 ~6 H. inum_examples_to_generate = 16
& Y9 \7 o& q9 v
2 ?7 ?! ~: e( u# 我们将重复使用该种子(因此在动画 GIF 中更容易可视化进度)
$ c9 Z) s# V! _& h% _4 aseed = tf.random.normal([num_examples_to_generate, noise_dim])5 R; P* g0 z% Y- P6 G5 ]- G. p  n1 h

2 {7 b' b) Q. d5 l2 e4 W# 注意 `tf.function` 的使用9 j6 c! G% o4 P
# 该注解使函数被“编译”9 u  P$ l( {" b& u6 v# o
@tf.function+ {7 I0 [+ v: Q. u7 t/ B% A3 O
def train_step(images):
4 r/ J" B8 K! y7 D4 F1 ?- M1 m+ B& ~    noise = tf.random.normal([BATCH_SIZE, noise_dim])
/ c- [& A/ p3 e . l: `1 g+ Z/ O! T4 g5 l4 \& D
    with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape:
- F; |- d$ e# H" U8 R. j      generated_images = generator(noise, training=True)
! g( m/ U, R" `) D% }9 | * F3 Z7 ^2 P3 ^% ?2 S( |/ W/ `
      real_output = discriminator(images, training=True)2 K0 D9 l, f/ K! d6 x& q/ d: U/ A; h
      fake_output = discriminator(generated_images, training=True)
) t( G' a$ o9 u; u" ?6 B5 s ! P4 t  ~$ n" G8 R! v
      gen_loss = generator_loss(fake_output)
+ G8 k  U6 U; s: P# o5 U% f      disc_loss = discriminator_loss(real_output, fake_output)
2 T9 f: A& q3 R6 x: e3 w" p6 B% E' u9 g - I5 d, J! x: i
    gradients_of_generator = gen_tape.gradient(gen_loss, generator.trainable_variables); H. u" {+ n: w; E( Q
    gradients_of_discriminator = disc_tape.gradient(disc_loss, discriminator.trainable_variables)
+ b$ z0 H: E7 F% z1 M6 _: Y5 s
! ~$ Q. m/ B9 m/ f7 `    generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables)), Q$ P! j! O5 r% J
    discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables))2 U% C4 O# Z8 c3 H0 c

/ s9 K0 L* q& Gdef train(dataset, epochs):
, h* a  I" k( m- l  T  for epoch in range(epochs):
" t+ j/ `# m- ]* o* W    start = time.time()& e5 I$ K8 f. u! c6 n  h
4 P* B2 B6 |' ]2 J/ V
    for image_batch in dataset:) ?7 x4 |6 S0 ]* W  c9 G! f
      train_step(image_batch)
( P1 j% O, r. ]7 J; t8 f* A9 l, M# z
- l3 x) c/ E7 I) y. I    # 继续进行时为 GIF 生成图像, T2 S& p$ ~2 S1 B
    display.clear_output(wait=True); `: F4 {* ^5 g6 M6 h8 g
    generate_and_save_images(generator,
, o( r$ F2 t% }* K- l$ o* ]! b                             epoch + 1,
' F  K% m' K, X* {                             seed)
( T  t3 ^& c& _7 [: J4 @ 8 Z8 |% I3 I0 p/ Y" |* @$ G
    # 每 15 个 epoch 保存一次模型# p6 n9 o+ k; Z8 i$ P4 N
    if (epoch + 1) % 15 == 0:# B& u# N" {" Q* \! m# O0 B
      checkpoint.save(file_prefix = checkpoint_prefix)$ A9 Y4 y6 x3 h+ X# t6 o. h8 X) l$ ~9 d
0 Z* ]9 o% H& n* i' s4 e
    print ('Time for epoch {} is {} sec'.format(epoch + 1, time.time()-start))
4 Q0 h) n( ?; {- z7 k! H7 @
0 O& J  _# t* Y1 L) q) k  # 最后一个 epoch 结束后生成图片. Q! v# U9 c8 ^0 [/ Z# J
  display.clear_output(wait=True)
4 I8 U5 e. u/ W  {, b  generate_and_save_images(generator,( `+ y' U; ]5 X  H# N
                           epochs,
0 Z5 m) T( C( e! L+ O2 @* ]                           seed)
0 E$ s! N5 T# Q : Y2 H- p& H0 o8 \" }. `
# 生成与保存图片$ m0 y5 i# ^6 h) _) O' ]
def generate_and_save_images(model, epoch, test_input):
& k+ f. E5 Z& N! p- Y  # 注意 training` 设定为 False) M: C/ l) \* ?" P
  # 因此,所有层都在推理模式下运行(batchnorm)。  E" P6 [4 d1 x: {  _
  predictions = model(test_input, training=False)# J4 s( V' Y, |/ I( D$ m

- S$ m! ^: g) |0 V  L3 _4 d  fig = plt.figure(figsize=(4,4)): J; j9 O3 r5 B& g7 d

8 q9 \0 `; f& z5 @4 H  for i in range(predictions.shape[0]):
- c* n7 i: R5 X# h* C# }7 w$ n      plt.subplot(4, 4, i+1)
1 Q; z3 g* z1 p) O      plt.imshow(predictions[i, :, :, 0] * 127.5 + 127.5, cmap='gray')
, O: d. y2 }6 J1 U! ?      plt.axis('off'): _- F0 ]: \/ s+ r5 D
8 T7 |7 k9 V8 |' `8 {
  plt.savefig('image_at_epoch_{:04d}.png'.format(epoch))
7 K- D" ?$ I8 M# J  plt.show(). b3 A* C/ A3 E2 a, }% m  y
4 U' F9 P8 U8 G6 U% ?+ _4 ~" f1 D* N
# 训练模型; y1 {& P0 R, u+ T7 J8 f, J
train(train_dataset, EPOCHS)
7 d8 d7 o" j, ?2 n& b: W) O
" ]# Q% X' Z" y( |3 I5 ?' o* N: q# 恢复最新的检查点
$ h& z: v+ a( c# u! Y2 k* h3 s: @checkpoint.restore(tf.train.latest_checkpoint(checkpoint_dir))
2 C' e" H% I$ [6 X4 y+ X - p9 K  N3 a; Y/ |2 U: f
# 评估模型/ f* e* N. c1 N# ^/ z+ [
# 使用 epoch 数生成单张图片, `( Y. U, ]! y* I# p
def display_image(epoch_no):  z3 ~0 u+ |" S: G
  return PIL.Image.open('image_at_epoch_{:04d}.png'.format(epoch_no))% i8 E7 ~" p5 `' M7 r+ e

" c# `3 l! y" ~+ y1 w  Q) D- Rdisplay_image(EPOCHS)" M2 ], w6 |7 a( q4 S3 ~/ D& s
* c5 Z, A, |5 n+ b3 r
anim_file = 'dcgan.gif'9 v4 K' x. x7 z5 d7 y/ @  r6 L
: p. P% X2 l# Y/ u, B! Z/ m: Q
with imageio.get_writer(anim_file, mode='I') as writer:
8 O7 S9 N/ a9 K9 u6 _6 C  filenames = glob.glob('image*.png')
7 L0 U. ?7 M7 R9 T, S* g4 n  filenames = sorted(filenames)! _1 K# T5 F# D% K6 G% G
  last = -1
0 f6 j% [' Y6 X# O% W, f  for i,filename in enumerate(filenames):  Q3 {# U/ G# e' E/ f% R- m( s
    frame = 2*(i**0.5)
" j. e0 F3 b/ G0 u$ G  i9 x) q5 b    if round(frame) > round(last):7 z9 G' i# {) R
      last = frame; [4 D! @" P% c* N  u  S: `
    else:
7 ]! q+ x' t) T  c0 N$ |$ R8 E; i      continue
: x2 J+ n; @% {( Q  \    image = imageio.imread(filename)
4 v  `2 ?( y; p: }9 E# W8 Q5 u# @    writer.append_data(image)
# U4 B2 M0 @4 r/ _0 i  image = imageio.imread(filename)
# J) S2 i, q1 ]& d  writer.append_data(image)
$ r5 x( h) D1 i7 J9 D. ]0 R - x4 x) }$ \9 e
import IPython
3 Y2 ]( L1 t4 U% Wif IPython.version_info > (6,2,0,''):
& T) x% M& g2 r/ q/ b. V  display.Image(filename=anim_file)8 z2 w, X4 J2 [8 X- v
参考:https://www.tensorflow.org/tutorials/generative/dcgan
! e) O0 W  l9 ?) J% R% K————————————————: K) b* r( a$ G3 o" F
版权声明:本文为CSDN博主「一颗小树x」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。
- U8 ^2 W: D) _4 V- z原文链接:https://blog.csdn.net/qq_41204464/article/details/118279111
+ @: m& S& w. ]# V) t5 i
5 @) e6 a6 Y& T* w% @* i. c0 L7 B7 z7 d, [* z





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