数学建模社区-数学中国
标题:
深度卷积生成对抗网络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. W
pip 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 w
3.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% `; ~- \* A
5.1 保存检查点
" A" p) o6 v1 _. H, M9 h% ~
) i- v' O' w! { d
! q0 G2 M" V7 ^. w( S1 j+ h$ W
5.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% W
train_images = train_images.reshape(train_images.shape[0], 28, 28, 1).astype('float32')
( g& j; Y6 t& i! S6 j, m1 u
train_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 I
BUFFER_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 q
train_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- i
8 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 K
9 ~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, I
generator = 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- v
cross_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 k
8 |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 m
discriminator_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; e
checkpoint_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! W
noise_dim = 100
# w% N1 D0 w1 W/ }3 q9 t
num_examples_to_generate = 16
1 ]- 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 U
def 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" F
def 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! V
5.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' G
train(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 n
2 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 P
display_image(EPOCHS)
q) j0 l# I. K& l$ q
anim_file = 'dcgan.gif'
2 ?' z9 j' K+ T: ?
, E6 ^9 K0 z! y1 R
with 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 = -1
2 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% d
import matplotlib.pyplot as plt
5 R6 T+ R0 V. {: [0 O6 `5 c
import 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 c
from tensorflow.keras import layers
6 |3 Q7 Q! x$ o/ K- U7 L5 X
import 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 Q
BATCH_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; S
def 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 O
discriminator = make_discriminator_model()
9 _: W+ I, H( C3 z; e8 b
decision = 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 E
cross_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. i
num_examples_to_generate = 16
& Y9 \7 o& q9 v
2 ?7 ?! ~: e( u
# 我们将重复使用该种子(因此在动画 GIF 中更容易可视化进度)
$ c9 Z) s# V! _& h% _4 a
seed = 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& G
def 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- R
display_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% W
if 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. c
0 L7 B7 z7 d, [* z
欢迎光临 数学建模社区-数学中国 (http://www.madio.net/)
Powered by Discuz! X2.5