数学建模社区-数学中国

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

作者: 杨利霞    时间: 2021-6-28 11:54
标题: 深度卷积生成对抗网络DCGAN——生成手写数字图片
% m5 i* w" K% T* R  T3 b& Q
深度卷积生成对抗网络DCGAN——生成手写数字图片) G: L8 Q# n) m! d' E
前言4 x0 P5 A: {, Z! |0 A" o; k! x7 \6 U
本文使用深度卷积生成对抗网络(DCGAN)生成手写数字图片,代码使用Keras API与tf.GradientTape 编写的,其中tf.GradientTrape是训练模型时用到的。
; \; j9 X. l) \
8 Q& U4 K. L0 O, r2 `6 z
5 J7 N: q+ A1 j. n3 r$ s% t
本文用到imageio 库来生成gif图片,如果没有安装的,需要安装下:0 Q9 y7 R0 W6 {8 B

3 x) J) ^3 N; ~3 C4 L2 T- E; b; o
% }( ]% t; I) S$ Y) G
# 用于生成 GIF 图片
( C# F7 `* Y9 S+ ]+ Jpip install -q imageio
5 J2 ^# n: p% x; q" _) }目录
6 x6 n2 T/ Z- c1 {6 O& |# ^" c- W) f5 l; f

$ r' k) Z! W+ H+ Q6 m. C% K前言
, x5 M9 p. }! S' j: T. h0 T  w' W! ~- P2 Q0 o, A

, ?1 w' G! B, t- p) {6 p1 u0 p一、什么是生成对抗网络?, L- |, N1 ^4 S, r

  {1 |9 q5 v5 `

# N. `1 P7 q1 V3 y6 i二、加载数据集) m' P7 Y+ q8 ]5 }: ]' T6 E/ n

0 ]8 B; U& t; E! [& a
+ T% V7 L1 z$ R( |/ e7 }5 b/ c
三、创建模型5 W# L5 U( ]- y/ U0 F! A( y$ z
, N3 p- g2 @0 ?2 P7 F( I
- l' g7 i# a4 J9 O$ ]
3.1 生成器
5 H$ _4 A% |2 L% k% \; W0 U, G: U
: S  @; F6 [* }# L
3.1 判别器
( z" B, L! c' ^  M" c* d2 G  P8 O- r% c
5 V# \. x  T, ]6 X7 o
四、定义损失函数和优化器
- |- v  {5 S& i# G5 {, I: N. A7 T* b8 w% @

3 z  e% C& @5 V6 I4.1 生成器的损失和优化器: U( u2 |5 ~* F3 p; |( z  O7 v
5 I; j3 Y/ U6 N7 f2 @: h! }4 L. W
( V3 f( V/ H0 M1 Q
4.2 判别器的损失和优化器
7 k6 {+ j5 Q+ n! r+ z+ [6 J" w5 j" \7 }# m; `) \

" r4 T6 l# o) r! t2 Z; z" p7 M: ^五、训练模型
$ V& I* p+ ?+ f: j# m0 T8 A2 q. t0 D
( m3 S) |( O- x2 ?9 z3 ^

, t4 A/ l7 `, K) n% h) y/ M8 T+ o5.1 保存检查点% R% G# c) L. R" P; \

7 t7 S5 A2 ?  C

! v, I& a' `  a8 Z0 |# _9 W5.2 定义训练过程+ C- E7 N; {( B* M1 R
  g2 B0 p2 t7 y1 |7 Q8 A7 B( z

; c  B$ q  c8 I6 k. h3 k8 y5 }2 E5.3 训练模型: P+ e. t% G, |

' H. m$ f5 V3 |7 a5 E! N

8 l% ?/ |2 T2 E* Y6 }六、评估模型
8 \# f7 p2 K: q3 ?) e
7 |0 _  j4 G) f3 S( c  ]

; @: _5 x' ~# \- F一、什么是生成对抗网络?3 o6 `% J& |' s7 \
生成对抗网络(GAN),包含生成器和判别器,两个模型通过对抗过程同时训练。8 Q8 I3 q3 X+ h6 z+ O$ I
+ W# t( G: ]! B4 w+ E4 B, O$ c/ @
8 D; ~4 W4 c1 `# l8 u
生成器,可以理解为“艺术家、创造者”,它学习创造看起来真实的图像。7 r2 Q2 k# p9 o. W6 N9 L$ A7 y

& g- Y0 y% w9 W* o$ d0 t, L; J
; @% o5 `8 ]& o3 s
判别器,可以理解为“艺术评论家、审核者”,它学习区分真假图像。$ ~# O- d8 D6 u, u3 P0 e0 ~
! P/ _3 h/ x" [, G: O

1 u, D, l/ O1 J/ b' |& s训练过程中,生成器在生成逼真图像方便逐渐变强,而判别器在辨别这些图像的能力上逐渐变强。  C3 h" ~% S: G8 f7 X) `8 {

2 r; U" I* ^5 d0 E* i
- J8 @' ~" z% y* @; O' A
当判别器不能再区分真实图片和伪造图片时,训练过程达到平衡。2 k3 Q4 \( W  S3 E$ h2 w! ?7 j% q
- F3 f! X8 u0 n, I4 f

4 P, [' ^" w! V( x+ a( b本文,在MNIST数据集上演示了该过程。随着训练的进行,生成器所生成的一系列图片,越来越像真实的手写数字。$ R; O: Z4 E' y3 r

3 j) Z, }3 B: P" k' v& h5 w
! `( t2 ]% C# Y
二、加载数据集
) M3 T4 S7 Z- y  _3 H7 u: e: s使用MNIST数据,来训练生成器和判别器。生成器将生成类似于MNIST数据集的手写数字。: C$ @# C( O& g3 J

% h) p7 q, _! w/ [

6 l+ ]% p" K/ d' o- @# R(train_images, train_labels), (_, _) = tf.keras.datasets.mnist.load_data()6 m/ A4 U' z" ]2 p: s6 d

: @1 N8 P0 w# etrain_images = train_images.reshape(train_images.shape[0], 28, 28, 1).astype('float32')$ X5 h: I# T5 l# S% u
train_images = (train_images - 127.5) / 127.5 # 将图片标准化到 [-1, 1] 区间内
) ~, d  s0 {; z) y
% \3 r: I1 f' a8 O6 \( {! LBUFFER_SIZE = 600006 O& Q: I) _  Z2 Z& |" s
BATCH_SIZE = 256
% |# G6 i4 C/ f+ e  H2 m
; y3 \( s+ T# V# y( D+ O7 R+ g# 批量化和打乱数据
) r% B8 ~3 a$ vtrain_dataset = tf.data.Dataset.from_tensor_slices(train_images).shuffle(BUFFER_SIZE).batch(BATCH_SIZE)
  `1 B7 a% I- O, \4 c, m; G三、创建模型( \$ `7 |, }( I) p$ \3 ^
主要创建两个模型,一个是生成器,另一个是判别器。! K, ~; q3 s1 n# J

) Q% o. }4 ^! ?1 a

) c9 u8 J$ |2 @2 Q7 K5 h3.1 生成器
2 v, Q0 Y3 ^* {5 P( M; ^4 o生成器使用 tf.keras.layers.Conv2DTranspose 层,来从随机噪声中产生图片。
6 J7 |, O% o1 @8 F5 ^6 w% P4 y4 S- e) M( U" `: {% L
! f- A  M; v8 F6 \9 q# s
然后把从随机噪声中产生图片,作为输入数据,输入到Dense层,开始。( @8 v3 ]  K6 z# F# n; T
5 x  a8 F7 `9 h' ]' q8 Q2 Y8 I
% Q" @7 Q. p5 g$ p) f* {
后面,经过多次上采样,达到所预期 28x28x1 的图片尺寸。1 @  }" ~# r. d" K( G* l

( }0 H- ]5 A; R$ r. M9 o% m
0 x, j$ \7 u. A1 j- u
def make_generator_model():
0 _0 S' c2 b7 g7 `! K6 z' J4 M* \    model = tf.keras.Sequential()
& O3 A* z0 C0 s8 t& n8 O    model.add(layers.Dense(7*7*256, use_bias=False, input_shape=(100,)))
! ^- c* p: d. I# y" r    model.add(layers.BatchNormalization())
  j/ f4 V& R" ]( `6 D/ ~    model.add(layers.LeakyReLU())
0 p# b1 e5 o9 O" @5 C* C % X- n/ V  H! b
    model.add(layers.Reshape((7, 7, 256)))
0 J0 O% f; S5 k" R5 z    assert model.output_shape == (None, 7, 7, 256) # 注意:batch size 没有限制
) W+ j+ z8 d+ ~* [5 @, N" {1 r
% o+ m! R0 I$ B& w0 w4 D- f    model.add(layers.Conv2DTranspose(128, (5, 5), strides=(1, 1), padding='same', use_bias=False))  R: |5 i* ]8 V' I) U/ \
    assert model.output_shape == (None, 7, 7, 128)4 g: W1 a4 k- p( h4 t5 E+ T
    model.add(layers.BatchNormalization())
: i' ~1 B/ \' W    model.add(layers.LeakyReLU())1 ?9 P& x. A9 u
8 g+ u* y& X4 ?, c) [
    model.add(layers.Conv2DTranspose(64, (5, 5), strides=(2, 2), padding='same', use_bias=False))
9 z' i# ]1 U, V0 x' e/ M    assert model.output_shape == (None, 14, 14, 64)9 a1 Y! _. a( r9 F! C9 T
    model.add(layers.BatchNormalization())
" [% O3 ?3 N  j    model.add(layers.LeakyReLU())1 }: i5 ~  F$ @

$ S! G" z& w6 a. B6 B/ p8 J! ~    model.add(layers.Conv2DTranspose(1, (5, 5), strides=(2, 2), padding='same', use_bias=False, activation='tanh'))  W( ]/ v, Y: G/ J- @, }
    assert model.output_shape == (None, 28, 28, 1)  }2 }& ?& @% p% U  ~& n
+ D* c3 X8 w( h) x9 s7 g
    return model+ R! o9 `2 }  H( ^% k
用tf.keras.utils.plot_model( ),看一下模型结构$ Y" V3 U# J4 n( [* `
3 a$ F$ a' N& i- ^
0 N' y' J) U; ?' c; j: g1 V- S

( Q- m- p4 S) M9 @- B% ?7 L! X6 d6 ]: t
5 x$ [3 v$ |) N
用summary(),看一下模型结构和参数
. I8 _7 {! b! X* ?
+ ^2 Q6 }% o; h: n3 F
5 ], Z8 [2 c; {) h7 x* P1 V# T2 m7 K
" c3 N' w9 f% c/ i2 o" W. m
9 m; u% Q( }5 S

  g, J- S$ H3 {. m' z5 j2 m
3 l5 i  q3 c- d2 U7 w2 W0 b* Q
使用尚未训练的生成器,创建一张图片,这时的图片是随机噪声中产生。, e0 Z: j9 T" _# z) x
' J% |) M/ q) D8 ~9 P6 j, d

; t3 Y: N1 c' [. S* M" Y% lgenerator = make_generator_model()' y' S$ G* T( Q8 _" s
" M5 |) d  \+ H* @& Z% \. G
noise = tf.random.normal([1, 100])% m0 q  O" ?8 I3 X5 |7 D# K
generated_image = generator(noise, training=False). {& F: w& i! U. W% M: P
+ s/ P7 b8 O! Q* u6 |6 ~
plt.imshow(generated_image[0, :, :, 0], cmap='gray')
! M" h9 i. E; Q" ^$ I7 N0 A( y8 p3 H% \. `3 L* o
1 x# S: J% @- O% ?
" w' K4 T5 h1 w' p

, l" S. n  n+ s) I" i. a7 q3.1 判别器
; v: v+ l" s5 A2 @2 A7 u( J判别器是基于 CNN卷积神经网络 的图片分类器。7 k+ k: \& u" H" L3 G2 M
$ H- Q% w, Z# [4 J, g$ N! ]
5 k0 P% g# J/ j- z1 j, R) }
def make_discriminator_model():' n# x( w; S6 r$ W
    model = tf.keras.Sequential()
& f3 Z9 j7 _! f/ T2 x1 Q6 z    model.add(layers.Conv2D(64, (5, 5), strides=(2, 2), padding='same',/ t' B) v6 L: [
                                     input_shape=[28, 28, 1]))1 E# A& M9 [& a( D5 f% Q: p; C
    model.add(layers.LeakyReLU())
; g* d2 k! N6 Y+ {    model.add(layers.Dropout(0.3))# B3 h* N, E8 x; e4 y, t

5 h2 O9 ?( H+ Q& z$ x    model.add(layers.Conv2D(128, (5, 5), strides=(2, 2), padding='same'))4 _0 n6 \5 P9 v% O- x
    model.add(layers.LeakyReLU())
" w5 W6 `) o+ W    model.add(layers.Dropout(0.3))
# q3 o! U) [3 g% Q. F1 [ ! G$ b) H! e, P3 s# n
    model.add(layers.Flatten())) [$ A! s3 W: T' _) R) ]' O
    model.add(layers.Dense(1))5 [+ V% q5 n8 |: v4 T8 t

  ~1 v6 c* x& C) B    return model; Y% L3 G% q2 J
用tf.keras.utils.plot_model( ),看一下模型结构
( Z# P6 K* t1 ]/ |( Y- t- X
" Q$ L% z; c$ I8 }0 |2 T

- Q5 s; d; a# F4 T; j; R! J$ V! W% j( P/ G

. X  `" X; f; u+ G4 D5 S. ~
- Y* E0 }; [8 |; u5 ^' q! J
& J7 v/ ^1 G$ t- K
用summary(),看一下模型结构和参数
: w2 _0 }9 X* r) Y- |7 G2 p& ~
8 r0 T# c  k, @
* p1 b  V" l5 u- \6 Q
+ U- y1 `; |, T9 x& n  j
# f  X' K! R5 W2 e% I

- N' }" [/ t+ u
/ K+ |5 D8 r0 U+ W
四、定义损失函数和优化器, ^6 }+ X7 |* L
由于有两个模型,一个是生成器,另一个是判别器;所以要分别为两个模型定义损失函数和优化器。, m' D' U- }  i0 x

9 q; o" k, ~. Y6 V, l, y
' Q2 v% Z0 W) u- l
首先定义一个辅助函数,用于计算交叉熵损失的,这个两个模型通用。
/ Z+ Y  J5 \3 n) x7 N0 K  L& P% H$ g

- D+ a0 @3 p6 G# G, Y# 该方法返回计算交叉熵损失的辅助函数
) U) f/ N' ?5 |) C# Hcross_entropy = tf.keras.losses.BinaryCrossentropy(from_logits=True)
: r& c0 @+ i* j& ^. W3 n% L1 S% t4.1 生成器的损失和优化器
9 T5 Q( W! b# _) _1)生成器损失
. R& d% \+ Y; v9 k7 {, T1 T: O4 I# X( `# s+ z
& A5 m  b  Y6 d9 W9 F4 h4 ?& \
生成器损失,是量化其欺骗判别器的能力;如果生成器表现良好,判别器将会把伪造图片判断为真实图片(或1)。
- H' X; m. _6 `) y; J- p5 ^
6 _5 ]& y8 `' Q. x7 R

( Q, h/ L4 G4 p& K这里我们将把判别器在生成图片上的判断结果,与一个值全为1的数组进行对比。, f1 Z; C5 c4 u( l$ _
2 e4 i/ y: B6 }1 ~% F- H5 N5 H

! O" c/ y1 k+ adef generator_loss(fake_output):, u5 w. t# H* E9 X7 _
    return cross_entropy(tf.ones_like(fake_output), fake_output)
" B8 z) i0 S/ }7 G$ e" n2 r+ G7 S2)生成器优化器' Y, \# D# ?7 \
7 D$ U' R" I$ B& c. G% S

& 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 @

+ l: F4 B4 T, O" p% j/ {1 o  plt.savefig('image_at_epoch_{:04d}.png'.format(epoch))% X3 R9 t* [* Q4 P) L
  plt.show()( s1 c" L7 U; B( y1 A+ y/ ^3 }
5.3 训练模型, A3 z2 E  d) Q
调用上面定义的train()函数,来同时训练生成器和判别器。6 W5 d; f% S0 {1 e/ V3 P0 |; o
* U& m- S! ~0 `& s; K" d
) U" w+ y$ Z; k7 F/ j8 U7 F! N4 r0 q
注意,训练GAN可能比较难的;生成器和判别器不能互相压制对方,需要两种达到平衡,它们用相似的学习率训练。" \; c8 H1 _6 F9 g% I

3 w! e9 ?- x7 i# l2 o% W' t

+ J, g/ u" K5 r" f8 K3 ^%%time
8 i  u* p6 P4 Utrain(train_dataset, EPOCHS)3 d+ ^' n) I5 ?2 o1 s( \6 X& P& c' B
在刚开始训练时,生成的图片看起来很像随机噪声,随着训练过程的进行,生成的数字越来越真实。训练大约50轮后,生成器生成的图片看起来很像MNIST数字了。
" ?$ j- t$ q. _. l" ?  a- h0 h2 X- t- a, k  ~

4 D5 u7 W! }( _6 p7 t9 d; W训练了15轮的效果:
- X) S* U+ ?  }, O: n! D
& ~0 h! `: M. S4 M& z: _1 o4 S

; J7 H; W  b% L3 z( g" |4 u
& ?& g, E3 j, R# y8 K1 ~) Z+ [3 K
4 S8 w6 Y% X2 S
1 O$ N9 O- w  X2 G' }$ A

. V: w( Y# t9 P5 ?: Z- _训练了30轮的效果:  R  d6 f: m, a. [( ]0 G

' B; L3 E% w0 R. ~3 }
) r* O8 q+ N+ m4 M1 p' S

! k. @, |- _) H, D6 D
$ f* N# I2 G) Z0 P8 p; M4 U

  C8 `7 ~6 _: ^7 e0 E

% B% Q! B) m8 f, x3 _$ q训练过程:
6 ]; e2 W& j  y; Y  `( w! I2 u
# H* o8 i2 c2 k5 l: T
; l( Q. u( k" a" Z  S

3 Z6 f9 P) S* A! H, W+ R

) J+ X4 R4 B9 j; ~6 H  x1 p  x6 g
) D7 B0 M2 O4 q; _
恢复最新的检查点/ U1 r+ ^6 v/ ~9 D" k
: {& j( F3 n% k$ K
4 A+ ]% F3 Q2 Z
checkpoint.restore(tf.train.latest_checkpoint(checkpoint_dir)), e" g/ N7 k' r" Q  y) {
六、评估模型
% f: O" x( c) j这里通过直接查看生成的图片,来看模型的效果。使用训练过程中生成的图片,通过imageio生成动态gif。6 X8 [/ a4 _% V8 x2 f" E
+ D6 @+ j6 e4 x1 {- p! m, E

9 C' m9 T/ ]: \" u6 o1 R# 使用 epoch 数生成单张图片
$ `' J8 v- t0 s+ b6 C9 `+ @$ `/ wdef display_image(epoch_no):9 |: }. m: F- R# l! N8 T5 ~
  return PIL.Image.open('image_at_epoch_{:04d}.png'.format(epoch_no))% h$ E& |: y7 L" _. X1 `7 S
- A/ b7 \& }8 R, W- G, I
display_image(EPOCHS)8 M/ ?( c3 X/ R+ C+ g
anim_file = 'dcgan.gif'% v) h3 I/ A$ a/ n
& Q, P0 c5 u6 ?" J& _, i
with imageio.get_writer(anim_file, mode='I') as writer:+ g5 k& E4 r( Q7 V  P) K3 f) g
  filenames = glob.glob('image*.png')
$ J: ]1 x' z' X2 R7 @! U  filenames = sorted(filenames)
$ X. v% P) Y1 m% K2 v0 r  last = -1
6 d  w! G7 s8 B  for i,filename in enumerate(filenames):
. V1 h* _: `$ C; c- O    frame = 2*(i**0.5)
" J+ q) Q' x5 v. H    if round(frame) > round(last):* s# V* s* D. e" b$ c
      last = frame' {0 _7 C8 s: \% ]3 A4 e
    else:
' Q4 k) A0 z7 z/ `8 W      continue0 @; F/ J/ q! \; E$ S& k: V/ t7 ^
    image = imageio.imread(filename)7 ?, \: \2 ]% {0 E
    writer.append_data(image)
# D5 Y1 @5 K: G. W) _. Q$ l  image = imageio.imread(filename)
1 |1 l+ K1 |3 k4 {2 @, m  writer.append_data(image)8 W. |- J+ Q' _$ ?+ D6 s  i

6 j0 ~: W- v5 K9 [5 Z6 Mimport IPython
9 ?! P$ l/ I( |- D$ V! O5 vif IPython.version_info > (6,2,0,''):( U7 K; v1 R" y$ w+ a1 f$ x! w. ^0 X! h
  display.Image(filename=anim_file)
- y9 P5 X+ w; D& I/ R( T2 k! H. w6 Z9 Q

0 w. F. }; h0 D1 g/ j0 @5 p$ H5 s! l2 H- Z* P4 e! A2 ?

) e& O3 \; i6 k* @( r" d1 V; F# }完整代码:2 q$ t! y9 a! |$ t  v
( M) N9 q( S( P' K( Y# T

/ A5 M* S5 D1 U- `import tensorflow as tf; n4 ~2 a3 ~0 ]* T* n
import glob
2 C, A5 V; M% c0 I7 wimport imageio1 E& O! \5 i6 C" ]) Q. H  o( ~
import matplotlib.pyplot as plt
2 s4 y# m8 b0 S- J8 q4 jimport numpy as np
+ D/ \/ T  @) m" [import os6 k' x& \7 E5 U9 y! d7 y! `1 H+ e
import PIL
. C+ Q& O/ a* e" s* j# ~0 ~from tensorflow.keras import layers
  r# F% ?9 k# W) @6 V0 fimport time% |2 v/ Y$ K; Q& O
5 `7 ^* S+ M+ E6 q( ^+ v3 l
from IPython import display
9 M7 M( Q/ ?+ ^7 {8 ^: x   }3 g7 |. d/ u' B* _: q( W
(train_images, train_labels), (_, _) = tf.keras.datasets.mnist.load_data()0 ^) w5 q" [5 l2 V0 N9 O

( q" F+ e7 e( ?( ~) _8 R6 }train_images = train_images.reshape(train_images.shape[0], 28, 28, 1).astype('float32')
# J, J* s1 p7 P; `2 k9 v  g7 u1 Mtrain_images = (train_images - 127.5) / 127.5 # 将图片标准化到 [-1, 1] 区间内: @, U  F! a) \# H

) D+ U9 P0 s+ M  t+ T6 W5 eBUFFER_SIZE = 600007 E! G: I% z8 ]5 t! M
BATCH_SIZE = 256. p' _, c3 f) }; l+ z# b
8 t  T9 S$ k& W  E0 D
# 批量化和打乱数据; o9 m* s4 z5 O( C" P8 O6 t% k0 ~
train_dataset = tf.data.Dataset.from_tensor_slices(train_images).shuffle(BUFFER_SIZE).batch(BATCH_SIZE)- E- A' Z! m' b6 z. q% C

0 D  A& ^, L3 ~1 L# 创建模型--生成器
3 q( l$ Q  j6 Z6 f5 a( H. @. pdef make_generator_model():
7 W/ O, B! o- t/ B8 `    model = tf.keras.Sequential()" g8 n% b$ g0 f6 i& `5 Q
    model.add(layers.Dense(7*7*256, use_bias=False, input_shape=(100,)))
" B4 {/ U8 u$ G/ w    model.add(layers.BatchNormalization())
$ ^) j1 n7 G/ F    model.add(layers.LeakyReLU())' p$ h1 y3 Z- i' y( E2 C& b0 G  o

$ q' j% [! o& P$ L- }    model.add(layers.Reshape((7, 7, 256)))5 y) F: F& X6 h5 q* Y7 X* x. m5 x
    assert model.output_shape == (None, 7, 7, 256) # 注意:batch size 没有限制9 i$ Q1 e* l% [$ `  ?" @
& i. `! S# B& {$ q' _3 {
    model.add(layers.Conv2DTranspose(128, (5, 5), strides=(1, 1), padding='same', use_bias=False))
$ c: ?& X* P" L4 U" a* o    assert model.output_shape == (None, 7, 7, 128)
6 r9 k( n1 R" W) e% x  I8 u    model.add(layers.BatchNormalization())
. C6 n: M! Q. s    model.add(layers.LeakyReLU())
9 g( `- T( C0 y 6 o% t/ o  Q+ t3 a  ]) ~( m0 @# h
    model.add(layers.Conv2DTranspose(64, (5, 5), strides=(2, 2), padding='same', use_bias=False)); L1 U6 r% L, E: A% I* Y0 l/ Q
    assert model.output_shape == (None, 14, 14, 64)$ O7 ?2 J# K" _1 M
    model.add(layers.BatchNormalization()), w9 l7 b. k$ [: c  y/ v
    model.add(layers.LeakyReLU())
5 i$ n' z+ i6 D0 l
6 r6 R2 Q8 F5 Q, c    model.add(layers.Conv2DTranspose(1, (5, 5), strides=(2, 2), padding='same', use_bias=False, activation='tanh'))0 F5 ]" C. {9 `( x' O$ {% j
    assert model.output_shape == (None, 28, 28, 1)
: ?& B6 a6 D) y& q$ B
! a# c! V: l1 Y: _  M    return model
8 o9 G% g1 T# W, P$ K; J: X1 o/ L/ c % J9 _- @* V- x+ x. Y
# 使用尚未训练的生成器,创建一张图片,这时的图片是随机噪声中产生。
3 Z$ |- d2 \$ `- d+ q# l9 Rgenerator = make_generator_model()
( j& i# j& ~( h( H 0 c$ u& B# \' S9 E
noise = tf.random.normal([1, 100]). y) W% {# R- A, ]5 |: J
generated_image = generator(noise, training=False)
% Q# k1 k+ s$ C2 `
- ?. `. b: k. H8 L7 r1 ^plt.imshow(generated_image[0, :, :, 0], cmap='gray'). N' B; H4 V$ N; A
tf.keras.utils.plot_model(generator)
1 B$ A' `" ?; @% [- {1 b 7 U6 v$ [6 J0 O4 m6 i$ @: k
# 判别器& H' Q. O% r- a4 ]# t
def make_discriminator_model():2 [! Z" G- u0 J4 L0 {
    model = tf.keras.Sequential(), `2 U2 j1 `) m& R; ^
    model.add(layers.Conv2D(64, (5, 5), strides=(2, 2), padding='same',2 Y  R  ~8 f% `( _# R* m6 I
                                     input_shape=[28, 28, 1]))- W% N6 X: N( k6 ^
    model.add(layers.LeakyReLU())
- S' d7 |: Q# H, C    model.add(layers.Dropout(0.3))
4 A) Z& j1 f+ n% s) I
, B: I  D# L1 D3 K+ n+ t    model.add(layers.Conv2D(128, (5, 5), strides=(2, 2), padding='same'))9 |; E4 z* L+ L6 _# o
    model.add(layers.LeakyReLU())4 j# V, Q/ Z/ z8 z9 P" T
    model.add(layers.Dropout(0.3))
1 T* y2 x1 R& a/ m2 _- ~3 S8 b0 A
& K% B. O: Y: p6 n( i3 @2 |: i    model.add(layers.Flatten())) H+ r8 H/ N0 p. n' C7 ]1 J+ X
    model.add(layers.Dense(1))
" H5 p  v4 ]; V, l+ K7 G 4 Y1 f6 n7 J# T9 Z6 ~) ^' V
    return model8 e- `+ M9 a9 m7 K( Q2 u; l6 J% e

7 b$ B* \% m* t5 [, e* t7 ?# 使用(尚未训练的)判别器来对图片的真伪进行判断。模型将被训练为为真实图片输出正值,为伪造图片输出负值。
! I+ S/ I8 S1 k( g# adiscriminator = make_discriminator_model()* Q0 b1 V5 D+ K, m4 w% W/ ]
decision = discriminator(generated_image)7 [% U3 S+ c6 R% e0 q
print (decision)
& m6 T9 Q8 d+ v+ w5 G3 [ , |* m4 N; D. R
# 首先定义一个辅助函数,用于计算交叉熵损失的,这个两个模型通用。
9 |2 T9 M. R1 v3 c! Xcross_entropy = tf.keras.losses.BinaryCrossentropy(from_logits=True)
1 F) F$ }; f" O; e4 c* w# k  q. o
+ t1 C$ \0 _& }- r; ]. P# 生成器的损失和优化器2 P( |: A& d! ~
def generator_loss(fake_output):
, }" \' M1 u( R% l# T$ f    return cross_entropy(tf.ones_like(fake_output), fake_output)7 h; j6 s8 l+ l
generator_optimizer = tf.keras.optimizers.Adam(1e-4)
7 Z* Q# c/ ~1 j% `+ K
+ G# ]. s- l0 d% V" F# 判别器的损失和优化器3 f7 A$ g2 ^( _: G4 L' l7 R6 q
def discriminator_loss(real_output, fake_output):: O8 u* G3 U2 Q2 j; H# h/ J
    real_loss = cross_entropy(tf.ones_like(real_output), real_output)- F! j5 r/ C) F3 b
    fake_loss = cross_entropy(tf.zeros_like(fake_output), fake_output)) u/ ?6 c# N/ k# G% V
    total_loss = real_loss + fake_loss0 z/ w3 j: v2 m6 Y& S3 F; l% c
    return total_loss
0 R9 J  ~5 Q6 a+ {$ e# h5 U% adiscriminator_optimizer = tf.keras.optimizers.Adam(1e-4)
4 |2 G6 D5 i) a: X) y1 ~ & l, Y2 R3 D  H! x& N2 p+ e
# 保存检查点
# R+ ^; n' h% {6 echeckpoint_dir = './training_checkpoints'% D. r$ M) Z- s; s0 i
checkpoint_prefix = os.path.join(checkpoint_dir, "ckpt")& T  a# O8 W5 P. E$ K
checkpoint = tf.train.Checkpoint(generator_optimizer=generator_optimizer,4 U8 r0 ~0 Y' U: r
                                 discriminator_optimizer=discriminator_optimizer,
; a. Z4 y6 `- \$ U0 t* Y                                 generator=generator,
" F" [3 J0 M, M/ V                                 discriminator=discriminator)
( H+ v; ~! E: F 2 \" O; z+ g: V: b+ Q, D8 U
# 定义训练过程
, n2 v: ?7 Q7 qEPOCHS = 504 s, a5 |( q9 s% Q- n- T
noise_dim = 100# F# ~* ~  l/ m$ }, N6 G7 S" j7 S
num_examples_to_generate = 162 k' }0 A) {: M( H8 @/ W: M

/ ]  l6 {9 U0 x% M$ f' _- f) ?# 我们将重复使用该种子(因此在动画 GIF 中更容易可视化进度)
. m7 z- O) m6 Z5 V+ x% p$ g! ~+ [seed = tf.random.normal([num_examples_to_generate, noise_dim])
6 F& s5 M- k9 A5 E9 K" X; _' ?- }* G
8 {" _2 A4 b/ n( G! K# 注意 `tf.function` 的使用+ ]7 n, J- ]$ g& B" `
# 该注解使函数被“编译”5 j. L, R9 h7 E% A$ M* R+ n+ L
@tf.function
- ^6 s+ [% T5 |/ n* o5 F/ B/ Edef train_step(images):
  D0 r$ H1 c6 [* H8 c. s+ {- }9 m    noise = tf.random.normal([BATCH_SIZE, noise_dim])
* ?" v( a7 L' O ( k  \" r6 ?2 S
    with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape:" _4 H, X  P. J) }& G. [
      generated_images = generator(noise, training=True)0 B3 r& q+ }3 z( C
& b. h6 I+ t6 R5 D+ t
      real_output = discriminator(images, training=True)% i; g/ W( |5 w/ ]2 D
      fake_output = discriminator(generated_images, training=True)# }4 Y/ x+ k' C, w9 M+ A
/ u+ Y  E! Q$ v1 l+ y# @/ ?  p* Y
      gen_loss = generator_loss(fake_output): q; i  i, c( |
      disc_loss = discriminator_loss(real_output, fake_output)
2 |$ m6 z! ?/ r 2 |' T0 v- n/ E& r* q
    gradients_of_generator = gen_tape.gradient(gen_loss, generator.trainable_variables)
8 Y( r% l% |) O) q# ^3 N3 ?/ ]    gradients_of_discriminator = disc_tape.gradient(disc_loss, discriminator.trainable_variables)
: g. G+ a: y6 k! p$ m
/ S  t5 V' n& z5 W# d5 |! }0 ]    generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables))
6 w* u! d% g! Y8 g6 c/ Y: Q    discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables))
/ E2 G; f9 J7 c, Y* r
# @; Y9 V0 \8 ~3 gdef train(dataset, epochs):8 \9 B) J" {1 p2 O% D
  for epoch in range(epochs):, D  s5 o) T: N
    start = time.time(), n- v6 F- I" V: \

; j% _" y" L1 O  V+ M3 S    for image_batch in dataset:) r' s1 h- E( j+ q
      train_step(image_batch)3 {9 T% m, ]" d" E

6 }6 l4 b/ f: a+ W" w1 ~    # 继续进行时为 GIF 生成图像
8 K# {9 f* ^4 @* C    display.clear_output(wait=True)
, d9 t# ?- y  I! l& n, @" u& m' h1 I    generate_and_save_images(generator,' s' f$ H3 L8 v/ ]* {4 G- u
                             epoch + 1,' U! w' h6 m5 u7 Q+ z
                             seed)
) _$ ]; P4 q) I) R8 ~0 H
" G( `  P& P4 m( j    # 每 15 个 epoch 保存一次模型. Z6 F5 v; u' D4 p
    if (epoch + 1) % 15 == 0:8 S! E& o9 d. Z$ E$ F5 \
      checkpoint.save(file_prefix = checkpoint_prefix)
  |  |% R, a3 [) i
9 v% l" a& E9 M; R    print ('Time for epoch {} is {} sec'.format(epoch + 1, time.time()-start))
2 I- G  i% s% J  k
' y% \; Y' d: t$ {2 L1 ^/ U  # 最后一个 epoch 结束后生成图片
9 ^3 F' y6 @5 |  {7 O, V0 {  display.clear_output(wait=True)
) r" ?$ ^$ d/ o4 H  generate_and_save_images(generator,: \8 ]8 v, ^# Y9 h/ D( Y' B$ f
                           epochs,# U' M1 f: A$ c& k: _( }
                           seed)
- r. A+ Q$ S8 s+ A2 ] , C- K2 H1 H* X9 n, |$ Q
# 生成与保存图片
" y% @8 K, ]& hdef generate_and_save_images(model, epoch, test_input):
9 k7 g+ w1 s9 I* o; w, E1 ]$ I/ r  # 注意 training` 设定为 False0 V1 \0 F9 V5 c0 H
  # 因此,所有层都在推理模式下运行(batchnorm)。
2 ~6 e. {3 Y! l; e# X, \, Z  predictions = model(test_input, training=False)
2 n, T1 q+ D; z5 n& E6 ?
) \; G# M4 @* F& I  fig = plt.figure(figsize=(4,4))
/ M& N# ~6 t+ `) m 6 E% z5 I0 H( ~2 `
  for i in range(predictions.shape[0]):0 L! F# R$ Y! J+ S
      plt.subplot(4, 4, i+1)
4 h: U4 q2 e2 s  G! r      plt.imshow(predictions[i, :, :, 0] * 127.5 + 127.5, cmap='gray')- K7 k9 r5 `! w9 L" H- M
      plt.axis('off')7 S  N7 l4 Z! u

' ]4 \/ t) n: b2 G+ M* x* _& J( f  plt.savefig('image_at_epoch_{:04d}.png'.format(epoch))
( O# V; g7 `. V% d; u  plt.show()
/ Y# W3 }5 V+ W- X 0 I2 ]' k: j3 p2 p; r! `
# 训练模型
: I. ?/ `) L$ G2 Q0 Y2 k) b& ctrain(train_dataset, EPOCHS)3 S( a1 s8 P4 s& d7 A" H- x4 z

) y+ l$ R. U. P  L! ^# 恢复最新的检查点
6 k( v9 ]. J( y  e6 q/ v( {2 T+ j/ L& {checkpoint.restore(tf.train.latest_checkpoint(checkpoint_dir))2 }1 D2 G- D, }$ @0 @2 w

, V$ I1 I' f5 ]& E1 |# 评估模型
+ l5 `7 p; ?0 K. s# 使用 epoch 数生成单张图片
0 v! T) [( B- o( X9 p0 jdef display_image(epoch_no):; F- |4 e8 G) m: H; l5 Y  M
  return PIL.Image.open('image_at_epoch_{:04d}.png'.format(epoch_no))
8 Y7 p# D8 j3 M
* [' K( d& N( _5 U" V" U9 Tdisplay_image(EPOCHS)5 D& o  Z( ^! z/ J
8 q; k; s' R1 `/ T
anim_file = 'dcgan.gif'
1 Z2 u, Q+ V4 _4 N, d3 [( j9 @# F 2 ]  I( ?: k# m' S( j
with imageio.get_writer(anim_file, mode='I') as writer:" O- O* R2 R/ ]
  filenames = glob.glob('image*.png')2 w5 S8 w5 l& H
  filenames = sorted(filenames)
( k4 z8 j% R* ~& l0 G* P) f6 y  last = -1/ Y( c. B* n8 u) L
  for i,filename in enumerate(filenames):
: q! ?, N: K6 H, s) _; `$ M! ^* n* G    frame = 2*(i**0.5)
- Q: r! q* `* b5 A# `( t- y    if round(frame) > round(last):$ T8 t7 Z3 z  `$ q% n% z
      last = frame) e; a# n" n; v3 q
    else:
5 a. [' V  v6 b7 o5 |      continue# ?. l2 P. \" {
    image = imageio.imread(filename)
1 K$ [7 d: r1 U# O3 N5 m; F    writer.append_data(image)- Y  f* t% d5 f! h8 F
  image = imageio.imread(filename)
; N9 G5 _6 c) q& }! T0 r2 d  writer.append_data(image)$ T* r+ s6 z5 j7 ~% {+ T, V0 v) [
2 \" P7 [* J8 w$ k' n# S
import IPython
% [: j. w% P8 v; U( b- _: @if IPython.version_info > (6,2,0,''):/ x4 m' [5 F( ~+ b3 v
  display.Image(filename=anim_file)
" y0 J7 ]" j( W# M8 j" |参考:https://www.tensorflow.org/tutorials/generative/dcgan
5 n& R2 v9 g/ C; K7 g* M$ [————————————————
$ g' |4 u. z& n4 I版权声明:本文为CSDN博主「一颗小树x」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。
3 B4 ]! W  p+ x/ J- J& l9 w原文链接:https://blog.csdn.net/qq_41204464/article/details/118279111
# c1 Y/ {+ @, W  D) j6 g" A# {% i' ]! J1 \9 m

( V/ d5 R. j: k- G# }4 D4 B* }




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