在线时间 1630 小时 最后登录 2024-1-29 注册时间 2017-5-16 听众数 82 收听数 1 能力 120 分 体力 565671 点 威望 12 点 阅读权限 255 积分 174924 相册 1 日志 0 记录 0 帖子 5313 主题 5273 精华 3 分享 0 好友 163
TA的每日心情 开心 2021-8-11 17:59
签到天数: 17 天
[LV.4]偶尔看看III
网络挑战赛参赛者
网络挑战赛参赛者
自我介绍 本人女,毕业于内蒙古科技大学,担任文职专业,毕业专业英语。
群组 : 2018美赛大象算法课程
群组 : 2018美赛护航培训课程
群组 : 2019年 数学中国站长建
群组 : 2019年数据分析师课程
群组 : 2018年大象老师国赛优
U9 }, ]) L9 G% M
深度卷积生成对抗网络DCGAN——生成手写数字图片
, R5 ?8 v$ _+ F. O 前言 9 \4 {$ e- u4 B/ `( u: k8 H
本文使用深度卷积生成对抗网络(DCGAN)生成手写数字图片,代码使用Keras API与tf.GradientTape 编写的,其中tf.GradientTrape是训练模型时用到的。 1 _1 u" K3 U# V( v1 z8 h
% `, B3 {6 S+ M$ g 0 d' u6 w- h/ |
本文用到imageio 库来生成gif图片,如果没有安装的,需要安装下: - o% Y( T6 B# G% J6 Z
3 O- j: J. S8 ? _: F+ u
% N$ v" G; V* Q2 j) w$ ~$ c
# 用于生成 GIF 图片
5 z( P3 ]0 b, z" I+ `+ y pip install -q imageio 6 o1 f( d& [7 L. O
目录 $ u" X4 |7 |6 t3 q0 C
6 e7 P' d u% C8 }: i- L
3 Q e1 d7 l* G: ?7 K. r3 K
前言
! ]. i0 R2 ^* A 8 ?% {9 Y' i* W1 M
, R( K8 Y; n9 E; y$ @. q
一、什么是生成对抗网络?
2 B: N1 }$ N& j: ?+ _
( q6 `- L+ L0 x& m
1 ~- J8 q1 K; W+ s8 { j 二、加载数据集 6 @7 z: [( ?! L. S
$ I; V' M2 u8 }) E
; G# o4 @) s3 I$ E. }+ f" [ 三、创建模型
: d+ b# s0 M$ a$ h
5 }/ P. d6 M5 P* w 4 M/ L5 A) l; e9 @8 U: f9 ?
3.1 生成器 8 ^3 k4 j9 q% X+ j6 H
: k7 M7 H9 g T5 z9 m0 U0 L- e
; V; j( N$ m0 e& ~
3.1 判别器 - B @: v9 _! k
" J9 L; K) n7 n6 v. Y1 f
; Y z% l. V/ ?5 D) t$ R 四、定义损失函数和优化器 - b2 s4 V* ^$ X( X* ~) x! X8 A
7 N5 r3 t4 Y) |8 H
1 H q0 r! g; b& Y$ Q) I9 N& l 4.1 生成器的损失和优化器
4 _' s Y P$ l' d
! [$ F/ U" u& \! g/ a, ~9 |+ w9 O' S6 H
6 i/ t0 [8 G& ~; a5 s 4.2 判别器的损失和优化器 9 F8 _" j v. n, _9 E" A
; G: e' Z K$ w9 U$ X
: W1 Q$ I9 O1 J& E& ~ 五、训练模型
9 A5 @ ~( [( b ]; e+ P! x
6 D, W7 N% X$ z5 Y- @ & c* ^+ {, m9 v ~) h
5.1 保存检查点
+ x& j# E& @8 I! M 1 g% J: g& u" F, n3 E
. A4 G0 w( g/ n0 h7 R 5.2 定义训练过程
" o: ^* x3 \: F) R 0 }, Z' s- b+ o) l4 l/ E. C! i
3 \, h' p; g7 Q( R, W 5.3 训练模型
; [+ A# B1 N) Y! n, x8 }5 c! E
% h+ v# k" \1 s* D8 y + L9 U4 ]0 {6 b, B$ y
六、评估模型
9 t- N2 v- e# @2 K
3 `7 @0 s, D4 \1 J4 Y1 |+ Z
1 P$ m9 F% l2 k8 }9 x 一、什么是生成对抗网络?
7 _8 K, P8 p( X- S 生成对抗网络(GAN),包含生成器和判别器,两个模型通过对抗过程同时训练。 / h4 q& S2 A$ u
0 C. E; Q" e6 b' a5 R5 G0 _ : }) u, l/ Z: v2 S' }$ ]
生成器,可以理解为“艺术家、创造者”,它学习创造看起来真实的图像。
( ]! |- p; D! @# b 9 T l, E5 Z5 c9 H/ D l2 R' `
2 V J' p) e3 x. ~
判别器,可以理解为“艺术评论家、审核者”,它学习区分真假图像。 0 L7 C2 N$ B: [) r; O/ y4 v: X" Q
8 t0 u. |9 x( J, ?/ V# H
4 l1 i% @- S$ V( a& | 训练过程中,生成器在生成逼真图像方便逐渐变强,而判别器在辨别这些图像的能力上逐渐变强。 . ?+ t/ H! y" Y4 O
8 t1 m+ E1 ?2 T6 Y F
/ w. i6 d! z, \; { 当判别器不能再区分真实图片和伪造图片时,训练过程达到平衡。 4 F' K3 I( \# R, f# B' _
. x8 h N' Y q3 t* Y$ h3 ~
& u3 \0 z X% O/ a( `& X4 R4 S 本文,在MNIST数据集上演示了该过程。随着训练的进行,生成器所生成的一系列图片,越来越像真实的手写数字。 ; _% i' K5 ]5 F; D
! V' y# z) @* A& g0 P: [' _
% I) ^ H1 e6 I% J, l- f
二、加载数据集 & {6 m2 A; O1 _! v1 h) [3 F. C* k
使用MNIST数据,来训练生成器和判别器。生成器将生成类似于MNIST数据集的手写数字。
- ~8 s: M0 q3 g, K/ R 2 W0 K' L: D- S* f- X6 G+ q
* [1 X1 Z) M# Q0 ~
(train_images, train_labels), (_, _) = tf.keras.datasets.mnist.load_data()
5 U2 h+ ^& ^ p! Q. Y ' g/ [/ f2 k W
train_images = train_images.reshape(train_images.shape[0], 28, 28, 1).astype('float32') & @3 x7 o' r1 v, ?
train_images = (train_images - 127.5) / 127.5 # 将图片标准化到 [-1, 1] 区间内 . ~$ N9 R% G' k, E+ u
/ f) v/ I7 q- ?6 `. ^% x
BUFFER_SIZE = 60000 1 U4 U! \6 U8 Y/ B4 ^
BATCH_SIZE = 256 8 G, z* v& Q& h8 ?. G L- N/ v
, r& f( k1 e% \# A* Y& I, V # 批量化和打乱数据
! H4 R' Q: C5 Z2 f train_dataset = tf.data.Dataset.from_tensor_slices(train_images).shuffle(BUFFER_SIZE).batch(BATCH_SIZE)
Y+ L# f- ~6 H2 m; J& v4 i 三、创建模型
3 u3 p4 M! W; J 主要创建两个模型,一个是生成器,另一个是判别器。
! V7 a2 l/ _7 U
5 v% C) M; a' I( r; f4 A% H. X4 |
H, W; S, h6 P7 @; [( b) [ 3.1 生成器
, b, M. q( }( y6 S 生成器使用 tf.keras.layers.Conv2DTranspose 层,来从随机噪声中产生图片。 ! _; ^! u" Q8 y6 r* K
( O' t! L& P' O8 Z3 `* E # V$ S/ ]2 q8 [1 W \
然后把从随机噪声中产生图片,作为输入数据,输入到Dense层,开始。
) L- R9 X" I2 E5 n# N+ N X
. J) B J' g+ k* k P2 @1 v 7 W$ D0 m) K" z( U' L
后面,经过多次上采样,达到所预期 28x28x1 的图片尺寸。 " M& S1 O# ^* q% [' ]8 Q& \3 Z) Z) G+ V
% L1 W! ]* B# h1 N/ Z
% D2 a) c9 z2 a def make_generator_model():
& ^, h/ M, `1 F9 j& ]( R# O& c, } model = tf.keras.Sequential()
1 {+ l7 ]" f& @: Y$ |. X model.add(layers.Dense(7*7*256, use_bias=False, input_shape=(100,)))
' k4 i7 c3 a2 ]8 k2 j model.add(layers.BatchNormalization()) % d. o1 y6 J; k. L' m# |
model.add(layers.LeakyReLU()) ; K3 q1 C5 F. S/ |0 k
2 [2 [' v# a& _& K* P1 M* d6 t
model.add(layers.Reshape((7, 7, 256)))
. O/ [ A+ i& {7 e3 D* i4 P assert model.output_shape == (None, 7, 7, 256) # 注意:batch size 没有限制
9 ?2 K& F1 Q! S8 H' p- D) u6 X ) X: f! q" o* j3 e
model.add(layers.Conv2DTranspose(128, (5, 5), strides=(1, 1), padding='same', use_bias=False)) % T8 {* R) D6 `4 T& f, V
assert model.output_shape == (None, 7, 7, 128)
x6 c4 A* F9 m8 t2 M model.add(layers.BatchNormalization()) j6 A k" c& ]3 g5 ^: r
model.add(layers.LeakyReLU()) 3 F: n7 e, @9 |, ]
7 `! K$ ]+ _5 i& h2 E) O3 w; G$ O
model.add(layers.Conv2DTranspose(64, (5, 5), strides=(2, 2), padding='same', use_bias=False)) + U5 L( l2 C) p U; V
assert model.output_shape == (None, 14, 14, 64)
1 t) ^1 N- h; t5 j: W model.add(layers.BatchNormalization()) 3 N; A" w' Z: v9 P2 c
model.add(layers.LeakyReLU())
. P1 c* U$ J& U) K5 ]1 L' [0 { # I. T# n1 @4 \6 Y: @0 c3 F4 l2 m
model.add(layers.Conv2DTranspose(1, (5, 5), strides=(2, 2), padding='same', use_bias=False, activation='tanh')) 3 @2 p/ f4 N+ s3 M7 D
assert model.output_shape == (None, 28, 28, 1) 6 X. X2 R( P: E. I% f
% o0 f& E5 |3 a$ j/ l5 o' S
return model
, `0 U8 Z( F" ]" a: \9 w( D 用tf.keras.utils.plot_model( ),看一下模型结构 2 q/ I9 ?: z t
* }2 E7 L1 ]5 S, B: y4 y
% _$ ^! ^3 B, U7 L( E
: V2 t2 A. V @* l% X, P0 c% s6 i9 G
9 j! \3 D$ I |2 l, D: Y6 a
7 h, q) S6 X) i4 ]0 r 用summary(),看一下模型结构和参数
" S0 B! q" Q1 p: d& { * B( N, z4 d, A4 V2 ^$ |
: B/ O. C; ?' q! l0 X9 [ 2 I: B1 |2 `" {5 n
1 _- S7 _7 ]: m$ c' R% [
) M b! W+ y7 Z Z
6 s9 x& Z) Z9 F8 ?& U, w4 b 使用尚未训练的生成器,创建一张图片,这时的图片是随机噪声中产生。 % z1 r2 J$ s( q0 H
5 O3 D* z M8 W0 Q1 D$ E
7 Z) Q& }4 m' Q( \8 Z2 s$ {. U3 @9 y/ V generator = make_generator_model() , C9 N# `8 w7 {& Z1 V+ x
" L$ y( q) _% v0 y* m# \
noise = tf.random.normal([1, 100]) 2 m5 z+ h* `/ O7 q$ K0 v
generated_image = generator(noise, training=False) % Q4 w/ j" D: h0 i8 S
+ `) Z( m) z- A# G x
plt.imshow(generated_image[0, :, :, 0], cmap='gray') 2 j& B- c4 _8 y* h( o( S' E4 v
6 B+ |4 s. g) M. W+ a
9 Y3 U! \$ A; L
* i( y9 A% n6 |8 ~. r' |0 w; l * S/ ?# v& T5 n8 m- }1 x( \
3.1 判别器
2 G: {3 V" k. H( y* M$ U1 V( H3 m 判别器是基于 CNN卷积神经网络 的图片分类器。
! R: o6 i% E' v" `4 z' b) i
3 h9 a# F$ Y$ H ) I, L1 q) b; Z1 _6 d( s
def make_discriminator_model(): ; [/ `1 ]8 D) w; A
model = tf.keras.Sequential()
) u6 W1 f1 P) F( e+ h. g6 H model.add(layers.Conv2D(64, (5, 5), strides=(2, 2), padding='same',
4 C2 g, G; @& ]- @2 N. r input_shape=[28, 28, 1]))
5 n k9 s0 X$ u4 C+ B% O model.add(layers.LeakyReLU()) + ~! \. M" r. {9 w- {
model.add(layers.Dropout(0.3)) [! e' W# M2 W( T9 m2 o8 ~
7 r& h' [1 i% }+ P3 b% @4 M! Y1 @ model.add(layers.Conv2D(128, (5, 5), strides=(2, 2), padding='same'))
5 q+ Z2 d9 L$ a* Y" y model.add(layers.LeakyReLU())
- n- z/ R" ]4 ~3 ? model.add(layers.Dropout(0.3))
- O( @! R: K1 |5 W" {/ N
% P! a( m% P. G* j9 i. v2 T model.add(layers.Flatten()) 6 |% L3 d# Z0 ~# s: B
model.add(layers.Dense(1))
5 }& \7 L+ ?& N% l5 I
) d3 I! O/ c: Q9 `+ q5 r return model
( v- k8 n7 G6 o& D" h5 O5 n 用tf.keras.utils.plot_model( ),看一下模型结构 % v8 I: p) V# h- B+ o; ]! t
) t! J9 W$ ^: ~2 e$ [% C; E* b
/ G+ h. w' N& y6 O" t; o" i2 ~ # q6 j: M" ~ I
1 ]) p' {- Y: X' |
9 V( ?& {$ T% ~! j. F6 A 7 Y6 H5 Y7 {3 ^. m
用summary(),看一下模型结构和参数
; h8 k( m2 C( q" { : t" `/ i, X! F$ |
( Q" q" X- K& F7 v
; I" R9 w7 Z0 C& N: f4 c 4 w6 i( }/ k6 z" @0 w
0 e' Q$ Y$ G2 D/ I7 S g* i) f
5 ~2 o/ |4 e2 h 四、定义损失函数和优化器
6 t) r! W7 a9 T# ^ 由于有两个模型,一个是生成器,另一个是判别器;所以要分别为两个模型定义损失函数和优化器。
v/ ?- w% G, o8 K0 T ( h9 y' I) t: b+ o' `
h% K5 Y% A6 T( e4 ^+ p 首先定义一个辅助函数,用于计算交叉熵损失的,这个两个模型通用。 ; i. H Z. B S0 C0 E- P
# _: N+ ?, p. H2 r5 J
6 P* L H6 ?& t+ `' T
# 该方法返回计算交叉熵损失的辅助函数 1 _- O: c6 X9 T2 y E
cross_entropy = tf.keras.losses.BinaryCrossentropy(from_logits=True)
) o- j6 X0 w0 T 4.1 生成器的损失和优化器 ! l" h5 s- t/ o' M" ?: E& M
1)生成器损失
) s( y$ O% Y; N, }; a9 v% R 5 D6 A6 Q; V2 y1 T
9 Q7 T2 _* l. y& S6 i
生成器损失,是量化其欺骗判别器的能力;如果生成器表现良好,判别器将会把伪造图片判断为真实图片(或1)。
! r+ J" R5 n2 q t/ G8 @' H7 F
1 L9 i# r& y( x5 S( @' o$ K0 d 2 W& p& v4 _9 ?$ t9 A
这里我们将把判别器在生成图片上的判断结果,与一个值全为1的数组进行对比。
) ^. @, w3 E2 } v) T& @
; t0 Y. k. h5 k2 {+ w' @; m5 g 2 h: J" g! r. A H
def generator_loss(fake_output):
0 d# P' {0 ^0 _( f, a return cross_entropy(tf.ones_like(fake_output), fake_output)
2 Y& F0 h& L# O0 o% A( ~5 v. @% R% H 2)生成器优化器
* o$ {; B" |8 W" T; K " a; }- c3 p1 Z, V5 \
z" u4 z4 D r# M8 N% P generator_optimizer = tf.keras.optimizers.Adam(1e-4)
4 u, s m* w; R8 x7 e 4.2 判别器的损失和优化器
; S- A' E# T$ g! s 1)判别器损失
6 h$ X" I) B% w- Q9 M" Y
/ x# t' y" t7 d# ` 4 m5 p5 O- d" I8 z: v
判别器损失,是量化判断真伪图片的能力。它将判别器对真实图片的预测值,与全值为1的数组进行对比;将判别器对伪造(生成的)图片的预测值,与全值为0的数组进行对比。
6 `2 n8 s1 K# U! k, k / Z) H& U3 E* [/ o
! t, \) _* ^$ l8 p; U def discriminator_loss(real_output, fake_output):
) w8 S! H3 e& l& z real_loss = cross_entropy(tf.ones_like(real_output), real_output) 9 U: L* _2 i9 k
fake_loss = cross_entropy(tf.zeros_like(fake_output), fake_output) ) i# P- n9 R. t
total_loss = real_loss + fake_loss 9 \/ L$ R& G, O% y0 u
return total_loss 2 M2 _7 ~% ?' `( U2 b. b- R
2)判别器优化器 8 m8 I' }$ j$ X9 U! m* h! x
, U ?! b+ w, z8 ]6 a9 s) v; S ) x6 q3 B9 l H: I" S* p* O8 @& ~! ^
discriminator_optimizer = tf.keras.optimizers.Adam(1e-4) & S( O ` ^7 r& ^) c9 ]- L
五、训练模型 : d1 D4 V+ W) U
5.1 保存检查点 - K. c. g1 q8 C3 [+ P; n |
保存检查点,能帮助保存和恢复模型,在长时间训练任务被中断的情况下比较有帮助。 8 O! }, M4 a- \
6 N# ^' b# Z0 c! l) u7 E- |
6 Q+ _2 n* Q5 W7 z0 u/ b checkpoint_dir = './training_checkpoints' 6 c3 l; R* a: X( `# Y& o! `7 I
checkpoint_prefix = os.path.join(checkpoint_dir, "ckpt")
+ e. U5 m6 k5 l checkpoint = tf.train.Checkpoint(generator_optimizer=generator_optimizer,
# j/ i9 S: n4 U- O discriminator_optimizer=discriminator_optimizer,
1 h" ]' s; Q2 \$ M generator=generator,
6 S$ t4 B* k% T% s# }# N3 q discriminator=discriminator) 0 Z! [, R: ^2 R4 | M$ p/ M
5.2 定义训练过程 # G' m. E! A2 X8 x
EPOCHS = 50
0 q1 n, A3 h. d& y! X noise_dim = 100
8 g" ~: A5 \$ A num_examples_to_generate = 16 2 |- g' ?+ {1 n- n4 [
6 p. {0 ~' N6 Q2 l w, T' O
2 }7 P; N" h9 L # 我们将重复使用该种子(因此在动画 GIF 中更容易可视化进度) : ~/ h7 ~3 i1 \6 U) N6 e) s) Z
seed = tf.random.normal([num_examples_to_generate, noise_dim]) ; G W7 T0 B- e# M1 W! U l
训练过程中,在生成器接收到一个“随机噪声中产生的图片”作为输入开始。 0 i* U/ m" K s# _* F( h" _' D! ?, Q
" F* B2 @7 X& m* H4 X. ~' `
1 T x" S6 L* x6 r9 e# y+ y 判别器随后被用于区分真实图片(训练集的)和伪造图片(生成器生成的)。 ( N/ u; K# [. o3 q
6 N, S0 V- F+ I0 j" ^/ i8 o 8 P, t6 N. g- K B8 r2 }+ K' V n
两个模型都计算损失函数,并且分别计算梯度用于更新生成器与判别器。 + ]- _0 z1 u7 }0 q' F
; V7 Y0 ?- R6 r7 I! W7 G# u# \, f9 b
0 e& s" T9 \3 q+ ?- c- P8 E( q
# 注意 `tf.function` 的使用
& \6 @6 q3 u5 b% g! _5 a9 N # 该注解使函数被“编译”
u& F6 D% \! f6 Q6 ?: O @tf.function
T9 P' L/ ]* \6 M3 F) j1 | def train_step(images):
0 C9 W3 s+ {, u7 u v noise = tf.random.normal([BATCH_SIZE, noise_dim]) ! `& A) K4 ]& n* Y" \
" N' ?- J6 P; M- N: j/ T( v
with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape:
9 \5 Y: H" m+ E" }: R t' w8 J6 S generated_images = generator(noise, training=True) $ h$ _4 y& l8 v9 ^7 O( }1 X
% K9 c# }' C' f
real_output = discriminator(images, training=True) 5 T3 E$ N) a& v* P1 ]7 z" k, N6 @
fake_output = discriminator(generated_images, training=True)
- P1 w: |8 z) W+ f D
( U0 E+ V, F! y. K: P gen_loss = generator_loss(fake_output)
- y+ _5 f% h# ?3 F T$ O disc_loss = discriminator_loss(real_output, fake_output)
/ W) W [0 P/ ]$ ^. a6 u% ]. H 3 c2 D$ \' A' p( v, ~1 A, i
gradients_of_generator = gen_tape.gradient(gen_loss, generator.trainable_variables) # t6 [( o& w( A$ K
gradients_of_discriminator = disc_tape.gradient(disc_loss, discriminator.trainable_variables)
, a6 e0 E9 _6 M) S
, w) V" ^( @/ B# D- I generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables))
p) n/ E5 F: X; N2 b0 g7 T8 H1 n& A( L discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables))
. \! n ^0 x: \3 l$ l! a 4 s$ l- ], J$ `
def train(dataset, epochs): # }" @% S3 ^0 W: ^
for epoch in range(epochs):
. @' y6 N2 J3 d) ]$ W( m2 _" R4 F. I start = time.time() / `. v9 N. S7 s5 W% B
" } y8 v; V; H3 E! |; m& d4 F for image_batch in dataset:
/ W0 }9 H- P8 q/ v7 l: O train_step(image_batch)
7 ?8 J) T7 j( q& g # Y, G1 P: Z z
# 继续进行时为 GIF 生成图像 & Y1 G' a3 b1 W A8 S# I
display.clear_output(wait=True) ( Q* a% x& ^2 z
generate_and_save_images(generator,
% D4 L6 L, @$ w. L9 Y+ ? epoch + 1, # K9 X( B" m- t3 A# ~
seed)
3 s$ {0 T! H7 d2 w8 d% b) @
* I& G# A$ f: [! Z" C # 每 15 个 epoch 保存一次模型 " G" V/ Z; Z3 K
if (epoch + 1) % 15 == 0: ) N) k4 c% k7 O$ c& q" j0 Q2 b) M
checkpoint.save(file_prefix = checkpoint_prefix)
' ]7 _) W; ~- o 8 l% `, q4 Y- _; z
print ('Time for epoch {} is {} sec'.format(epoch + 1, time.time()-start))
9 ]! E1 d6 O2 r2 e9 x, V6 X) a4 E8 a 9 W! y4 ]" c6 ^! N+ {
# 最后一个 epoch 结束后生成图片
3 y, |* \# Z x7 |2 Q: A' d6 V4 P display.clear_output(wait=True)
# ~9 ~; X7 s5 D generate_and_save_images(generator,
; j' }% x" y$ M5 ? Q1 ] epochs, 0 w$ m$ V) p) }! W% x' P
seed) * h2 u" Y* o+ P9 p
, @4 ?3 L/ M9 T2 w) X/ e9 |$ J # 生成与保存图片
/ n: [, i1 l" n( N def generate_and_save_images(model, epoch, test_input):
; Q) |5 P& V) l* s# j # 注意 training` 设定为 False . k: k, G) ?' x, @* s
# 因此,所有层都在推理模式下运行(batchnorm)。
' t: w5 ]1 t/ H1 T+ z6 p5 z2 H! ~ predictions = model(test_input, training=False) * M: q; D. A, T- |" a& _# m* \3 T
/ I$ \- I h0 X' @* Y fig = plt.figure(figsize=(4,4)) 8 t2 E5 ]! V! S! d
$ F; ~6 h5 T' X& u; X
for i in range(predictions.shape[0]):
6 Q& O/ O4 S0 r. N. b plt.subplot(4, 4, i+1) # J( S- |- W* d. M3 w. K; N8 M
plt.imshow(predictions[i, :, :, 0] * 127.5 + 127.5, cmap='gray') ' D7 D; T$ B* P
plt.axis('off') , @/ w9 F3 \+ U8 K9 b/ E
3 V+ d, R8 P5 E8 @6 h5 q1 y plt.savefig('image_at_epoch_{:04d}.png'.format(epoch)) % }5 W+ x& T7 y6 @
plt.show() 1 Q! e2 P$ j- i+ Z- ^
5.3 训练模型
% U. i+ \, W9 Z* R 调用上面定义的train()函数,来同时训练生成器和判别器。 ' ~: c' ~) s- t0 ]' e0 a! b1 E
4 l8 N4 y; s% o" ^# |& ]' G
& ?8 O: O2 e6 c$ K* H9 A 注意,训练GAN可能比较难的;生成器和判别器不能互相压制对方,需要两种达到平衡,它们用相似的学习率训练。
, T- z$ w3 \ c
u$ N+ \' K) P, g* b$ E 5 R& d+ l* Z' i3 J4 F+ t2 _
%%time
: |! \5 p. O3 e7 J: @7 p. F train(train_dataset, EPOCHS) " f, R- I- O& a* H' d* b6 ^1 g
在刚开始训练时,生成的图片看起来很像随机噪声,随着训练过程的进行,生成的数字越来越真实。训练大约50轮后,生成器生成的图片看起来很像MNIST数字了。
- K1 B) Z+ v$ @3 w d $ d* M H* {! U7 K5 A: `
, H4 R. Q" n" [. J$ c/ f 训练了15轮的效果: 8 U1 L1 _* r. Q; y4 ^1 v3 f
% U2 |5 O) k% _& i0 [* C4 g
1 K" f' C: y! o/ T6 w
6 \+ O2 T$ p2 s; W+ v + q ?2 g* g* z8 ~- `
3 s5 k! k0 U6 z/ ^! ^ ) I' S, \4 a. R4 t
训练了30轮的效果:
/ ^* ]" E [9 Z# d , M3 }% G# k P! V2 L- q
, t: z: l' {, S4 I3 M & v: H. p5 d. r1 X& A( P& z, {
3 b4 Y0 d5 L3 a& P& O& m) g! R* ]
& y0 D8 l8 i1 Y
: r2 e9 l6 E7 y 训练过程: 5 f$ Z5 ~: u$ o: @" ]3 B
" u$ d) o6 z0 b! ]8 @5 O " i9 v- L' O6 [9 q6 m3 |# N
8 a. n4 V7 l+ @% W5 v/ P( w ' N7 z1 _' p; G: L" }6 q, R b) b
" x% [; t0 Y; S% T. R7 e0 G1 }5 m" Q
# {( _1 b8 M' r' a2 a 恢复最新的检查点 ( V9 V. a& i' S* X( C' |
$ {8 c6 }4 s- L8 c6 \* H) F5 B
1 P* L+ \+ n/ L4 O) K! L% c
checkpoint.restore(tf.train.latest_checkpoint(checkpoint_dir)) , H; \8 N+ a1 ?+ Q8 l# Y6 @ N" o
六、评估模型
* u+ b( d( Q4 o0 O& ]2 P- H& a 这里通过直接查看生成的图片,来看模型的效果。使用训练过程中生成的图片,通过imageio生成动态gif。 0 i* w: B9 J$ g Z# ?% \
+ M; n( m% A: ^, z; a- g3 C/ a7 k- Q , U( Z! n9 O- \1 l5 t! N
# 使用 epoch 数生成单张图片
) t3 U5 c8 D d# k5 A1 I) X' f/ d def display_image(epoch_no): ! _. k* h& ^' c) a+ I- M! ~7 f) ^
return PIL.Image.open('image_at_epoch_{:04d}.png'.format(epoch_no)) 7 _! x+ }; G! V( T3 v
, I2 i! n- Q' i& `; T4 y
display_image(EPOCHS) 0 v5 v/ X L, x# i8 R. |, h
anim_file = 'dcgan.gif'
6 G) |( I6 R7 N; h, \
3 v- a% p! K$ r8 ]7 a* o6 x with imageio.get_writer(anim_file, mode='I') as writer: 0 m3 _: P* s2 A0 V4 m" @5 G0 {
filenames = glob.glob('image*.png')
8 X/ R$ C1 R5 g v/ T filenames = sorted(filenames)
1 q; d8 ]* w* s3 l last = -1
0 M' }; F6 i, L, X! d2 e for i,filename in enumerate(filenames):
% H7 j5 h3 N+ d+ | frame = 2*(i**0.5)
2 F/ X- L( q: k- H% L if round(frame) > round(last): + b% h$ d5 v; p9 O& B
last = frame
/ z/ x: R% Y0 T3 b: z6 D else:
5 C" q2 h# _& F- j continue + M! l5 y/ G2 w, U* D% j
image = imageio.imread(filename) 8 v8 C+ g& f- M8 X4 \. }" ~
writer.append_data(image)
" x7 s' ^% Z. H; Z( r8 z image = imageio.imread(filename) 0 t' N' A, B2 t& X; X) c# H2 A6 T0 j
writer.append_data(image) 1 h, g& }: i: p1 x. t
# ^2 ^* s) S( o. }7 G* k4 t import IPython
+ x. U& }$ o; m) V) M7 @* Y if IPython.version_info > (6,2,0,''): + m) Q" K# n: i. C2 V
display.Image(filename=anim_file) 0 Q; b0 |6 r- u
1 Z" z A9 Y' N8 O $ U$ Y2 z2 a- C4 w% W5 j U
, G m! m0 r! {7 K( X$ d 3 R1 b ?* Z$ s* M. D1 r1 r) r
完整代码: . k) l6 q- s5 Z
# i- u. k1 i, N
/ M3 m3 h7 Y, ^7 `" k0 o4 ^8 [ import tensorflow as tf ; q3 `8 T7 c' V3 y- d
import glob
& q% Z- W+ z$ \ import imageio ! D U! X8 D+ D: }2 P5 `
import matplotlib.pyplot as plt & s( C1 `; @2 p$ b
import numpy as np - T4 }# l+ a! F
import os
e B6 Y8 q$ D$ C. q" l import PIL
' g* E* Z C2 `' K- \& W. S from tensorflow.keras import layers 3 p6 [% x$ F9 D) j8 u% W; n
import time . g5 s! H4 p3 Q0 [$ N# Q7 E
5 z; A) u- h3 N
from IPython import display
$ ]- }9 d3 M8 X ]+ b4 m2 _5 W" q5 g
(train_images, train_labels), (_, _) = tf.keras.datasets.mnist.load_data() $ F# k$ t) @4 N
, o, Y6 t0 G6 V% p3 O" ~6 K
train_images = train_images.reshape(train_images.shape[0], 28, 28, 1).astype('float32') 5 a. a' Y& p" K0 a) O) k
train_images = (train_images - 127.5) / 127.5 # 将图片标准化到 [-1, 1] 区间内
5 H* Y( {8 q! v4 t( v. y: j " \5 I5 C9 w8 `( I4 m, N
BUFFER_SIZE = 60000
0 j3 i% u% f6 ]8 H, o/ R BATCH_SIZE = 256 1 o5 a* T$ G+ J. T) Z
# d$ W) q7 r% m6 U # 批量化和打乱数据 + {' {; O* Z: N7 |5 Q2 ]8 |
train_dataset = tf.data.Dataset.from_tensor_slices(train_images).shuffle(BUFFER_SIZE).batch(BATCH_SIZE)
! V+ M9 j. p) l8 z; j4 D7 z/ a8 I & i1 x& j+ f: A# o
# 创建模型--生成器 8 y1 s$ z; D4 Q3 f& [4 N, c" P
def make_generator_model():
. [5 X- G' m, W$ _" o- G model = tf.keras.Sequential()
+ y `# J( |! ~5 g0 ?5 ^& U model.add(layers.Dense(7*7*256, use_bias=False, input_shape=(100,))) 5 I, i! P0 N1 q, `) }
model.add(layers.BatchNormalization())
7 G: n6 V' S' @4 T model.add(layers.LeakyReLU())
& L, m+ ~' Z" [6 e, A) l: @ 4 X# z8 a5 ]& k$ a1 [
model.add(layers.Reshape((7, 7, 256)))
8 A, R" j* B1 j0 U: F assert model.output_shape == (None, 7, 7, 256) # 注意:batch size 没有限制
! c$ w, z8 Y4 S% B : J" Z; `8 \, S
model.add(layers.Conv2DTranspose(128, (5, 5), strides=(1, 1), padding='same', use_bias=False)) - \& `$ I4 z5 `- R( S) r: ^3 Y, f
assert model.output_shape == (None, 7, 7, 128) 6 j7 n) O* T" p, \6 B
model.add(layers.BatchNormalization())
( V6 [) I$ H6 u" A4 m model.add(layers.LeakyReLU()) 6 m; W* n+ H7 B7 z. ~7 |) B$ _
' M* P+ u J( C; x4 G( C% t. v model.add(layers.Conv2DTranspose(64, (5, 5), strides=(2, 2), padding='same', use_bias=False))
+ f# O: c- F: S/ @- B8 _ assert model.output_shape == (None, 14, 14, 64)
; A1 g# c5 |; x/ q2 j. K model.add(layers.BatchNormalization()) " }* G; g% o. J8 y7 P
model.add(layers.LeakyReLU()) 7 o7 _, Q6 \3 Z5 H
8 t. Q$ M, S5 K
model.add(layers.Conv2DTranspose(1, (5, 5), strides=(2, 2), padding='same', use_bias=False, activation='tanh')) ) l5 A' t! i' x+ u6 ^
assert model.output_shape == (None, 28, 28, 1)
: M0 ?3 t1 w& R/ X2 U( L
" _) F6 {) q- Q* A5 r return model ' d: Y# k* l7 A4 o( w6 L j
- o9 p# B9 z0 e7 m& K" E # 使用尚未训练的生成器,创建一张图片,这时的图片是随机噪声中产生。
* F' W& o1 M' r( R& d generator = make_generator_model()
7 a+ f) P3 P3 m N# s" [2 h, T, h & g" B: E4 N& v4 ]
noise = tf.random.normal([1, 100]) 7 A S( {' }8 G v M6 M! H
generated_image = generator(noise, training=False) + |8 I, R$ D2 _4 a
; H) ]! o1 e( F7 @
plt.imshow(generated_image[0, :, :, 0], cmap='gray')
s$ E" T2 }5 @+ c% W" F. R- | m tf.keras.utils.plot_model(generator) 5 B- s/ ~ n- e0 e
2 n) O/ [- b! \9 l% s% z& z
# 判别器 " C# b1 J/ x+ E
def make_discriminator_model():
, Z$ t1 P5 p3 |1 R% p/ J3 ~ model = tf.keras.Sequential() $ v5 T* k% q S0 V- T. C$ Y
model.add(layers.Conv2D(64, (5, 5), strides=(2, 2), padding='same', / ?& y/ E! }1 |2 S: n% H
input_shape=[28, 28, 1])) * k" I' D$ g3 O4 @4 I- \5 ~
model.add(layers.LeakyReLU()) $ K6 l: c3 u. J7 Z
model.add(layers.Dropout(0.3))
. g( Z/ ?6 p1 X. k/ a % Y9 _1 S5 m. h) n9 F
model.add(layers.Conv2D(128, (5, 5), strides=(2, 2), padding='same')) . h* L) w5 ?- g% E3 J- ]
model.add(layers.LeakyReLU())
/ [5 ?" t1 u7 ~8 r, C4 J model.add(layers.Dropout(0.3))
4 U- I* H, v6 m% ^/ Q9 A, K2 [ 4 Z8 r2 B# k0 n/ T
model.add(layers.Flatten())
* X2 f' `0 c8 S7 t3 z. b! k model.add(layers.Dense(1)) - w# J' j- F6 h+ L: I+ C
8 O/ c% \! W2 U! ^+ Y7 l return model , s9 {( y. n/ H7 F9 K
; z9 z$ c. q( {7 Y* r+ X
# 使用(尚未训练的)判别器来对图片的真伪进行判断。模型将被训练为为真实图片输出正值,为伪造图片输出负值。 6 b0 t0 R' ^2 ~* V. u. v0 S
discriminator = make_discriminator_model() " }( O8 u' {2 A( P& H) T! H1 n
decision = discriminator(generated_image)
! n3 Y# N5 z+ }8 V8 j print (decision) 6 r9 k5 H+ F- [. I
& G; T9 C6 | p2 g: H1 _
# 首先定义一个辅助函数,用于计算交叉熵损失的,这个两个模型通用。 ! v+ L3 c/ a, r- N# O
cross_entropy = tf.keras.losses.BinaryCrossentropy(from_logits=True) 8 l/ p# T5 e9 O D1 k' N
* T3 o+ ]$ M: ?& S, G # 生成器的损失和优化器 " e7 v0 e. M, l# H
def generator_loss(fake_output): 1 N4 F2 S' b- J5 d9 r/ Q
return cross_entropy(tf.ones_like(fake_output), fake_output)
& p. e; l0 R J) p generator_optimizer = tf.keras.optimizers.Adam(1e-4) . ^* }6 A# k; X, n$ G
; g ^% P I Y0 o2 F9 v: j$ F
# 判别器的损失和优化器 / U! [0 q& O* ~5 I) t
def discriminator_loss(real_output, fake_output):
7 f4 i1 ^' e/ F6 }& c6 d5 o( J real_loss = cross_entropy(tf.ones_like(real_output), real_output)
) H; H* s6 B* u6 @" } fake_loss = cross_entropy(tf.zeros_like(fake_output), fake_output) 3 D! H3 ^, i5 r) y
total_loss = real_loss + fake_loss
1 X& X! h1 Z7 N9 F; | return total_loss ' M% j9 l9 a( x; Z
discriminator_optimizer = tf.keras.optimizers.Adam(1e-4)
/ v1 X2 a6 V5 O$ a8 L9 V1 r
8 k6 K/ X4 `4 b$ q3 \/ M6 T # 保存检查点
1 z I9 ?3 L6 d$ S) h. P; e checkpoint_dir = './training_checkpoints' 7 B6 D C2 ^) ?7 y, B9 u
checkpoint_prefix = os.path.join(checkpoint_dir, "ckpt")
/ W" x$ F: w1 D" i0 ~ checkpoint = tf.train.Checkpoint(generator_optimizer=generator_optimizer,
8 {$ U* B; |/ s/ n' j/ u+ \8 @ discriminator_optimizer=discriminator_optimizer,
% c# j' f+ H/ c+ k generator=generator, 1 _3 E) o ?- b2 \1 L- l: E
discriminator=discriminator) 0 }+ C$ I# ^. X) r
& E% g. X; a: h: w
# 定义训练过程
) E+ ~' T! u/ D# n: H. K7 q EPOCHS = 50
/ w" B! L# [! Z& L3 z& u; h- o noise_dim = 100 , l2 c& b/ w6 x9 O
num_examples_to_generate = 16 - ^: F' ^: j" w3 O; C' B. W% T# i
+ Z7 l: u) Y2 g7 d; J* ]) y
# 我们将重复使用该种子(因此在动画 GIF 中更容易可视化进度)
# g- g8 X) w, m8 g seed = tf.random.normal([num_examples_to_generate, noise_dim]) # ?: i: D9 U8 Y9 L
. Q! M. r+ a% o) `! e0 ~$ g
# 注意 `tf.function` 的使用
) L9 }1 M- } h* Q+ P # 该注解使函数被“编译” ) ~( e. q/ n+ L8 ^5 }
@tf.function 6 b8 D0 e" ?' b( S- {; p8 J+ I
def train_step(images): " A- q6 O4 d, {# Q5 S/ {
noise = tf.random.normal([BATCH_SIZE, noise_dim]) : ~+ R( {! ]3 u
/ {: W2 _+ R3 j) ]* P i
with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape: % m6 R, S+ ^$ Z
generated_images = generator(noise, training=True) - d2 h4 T c! Z: _4 J0 J
D. _; ~9 H9 s3 I/ N
real_output = discriminator(images, training=True)
8 D( G+ W7 z" @) I0 v& f4 p fake_output = discriminator(generated_images, training=True) 5 b# V$ I) {, G+ C7 M; b, x; B
$ ]& x+ z" s: W# D3 X+ `# j/ g6 j
gen_loss = generator_loss(fake_output)
- ^9 n" T% e& R2 p2 z8 J disc_loss = discriminator_loss(real_output, fake_output) % i. G. D1 q- F e" Y/ K
* P/ Q4 ^8 A0 ]5 n. e gradients_of_generator = gen_tape.gradient(gen_loss, generator.trainable_variables) $ u# A9 x) Q7 g
gradients_of_discriminator = disc_tape.gradient(disc_loss, discriminator.trainable_variables)
, d v5 l+ R. O( z
! h; O% C6 `: `+ Q+ G6 W generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables))
$ Q( N5 A% t" F5 N" p, D2 w& G discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables))
9 Q3 ]/ |! }: S: x; z3 L8 E
/ T# n5 L9 G, h$ m def train(dataset, epochs):
8 ~4 n3 u! D* K% L" f5 F for epoch in range(epochs):
' v: w% Y8 p6 s0 g6 N! q start = time.time()
* X9 r9 Y* f# |3 S: A# v2 C" o : A; X, A" d2 H( _
for image_batch in dataset: - i9 l% X A/ W( K% Y# R
train_step(image_batch) 8 a2 F& i% [7 @$ Y4 z$ y- m
6 L" K+ R' e# N9 h # 继续进行时为 GIF 生成图像 # G, U9 Y" }/ r: V( Y j
display.clear_output(wait=True) Q) M8 l4 W5 s" d9 f- p
generate_and_save_images(generator, ' C6 i. M! O, U5 F% O
epoch + 1, % H6 U. n0 j0 }" s+ z
seed)
& {& c! e1 k9 p: t& x2 M $ ~6 A% J- O. T3 n- P& G
# 每 15 个 epoch 保存一次模型
3 ~+ o( G. F& s4 ` if (epoch + 1) % 15 == 0: 8 E K I5 [( m
checkpoint.save(file_prefix = checkpoint_prefix) / I% A, L$ w& L4 i
1 w* t6 l6 [" v
print ('Time for epoch {} is {} sec'.format(epoch + 1, time.time()-start))
6 ~! R) Y9 |& R) P4 B) i( u ' |, N0 e" V7 D7 j7 A% B
# 最后一个 epoch 结束后生成图片 $ J$ U" c- Y3 Z
display.clear_output(wait=True)
$ D; G' w$ @) Y6 N& N generate_and_save_images(generator,
7 X0 a" n; B" B5 |1 v/ o epochs, & q5 b8 c9 |9 a9 ^2 L8 h
seed)
& F% s& a$ j& k
& \3 @; J$ w8 U, @% m( H5 } # 生成与保存图片
. O$ i' a: }2 q) F* j3 x def generate_and_save_images(model, epoch, test_input): 6 E7 ^+ H) e# s; Y4 [' L
# 注意 training` 设定为 False . w3 P; H2 w, h% S& W! X) H
# 因此,所有层都在推理模式下运行(batchnorm)。 8 L9 ` \/ e! t" o8 y
predictions = model(test_input, training=False)
( n% {0 p6 A& i8 J3 j7 w) V) v 4 a- m6 Y, C' n: q3 {/ k
fig = plt.figure(figsize=(4,4))
* Q5 |1 s& l% o0 |, X . @' C1 p4 `6 r$ {6 p2 V7 v9 a
for i in range(predictions.shape[0]):
: y# Y- V8 f! t9 i H$ i plt.subplot(4, 4, i+1)
- W" X. a% u7 m5 p4 V2 G plt.imshow(predictions[i, :, :, 0] * 127.5 + 127.5, cmap='gray')
+ C* Q* v: G0 N [8 c8 t) p plt.axis('off') - Z; ]8 S1 {" T& h
; v; Q+ ?# b' G( Z3 E
plt.savefig('image_at_epoch_{:04d}.png'.format(epoch)) / {% l& Y8 m/ l3 E
plt.show()
: @# g9 J& a+ d1 d+ g. ]" y4 i
7 e0 w: A" q! M. V. {9 B # 训练模型
3 X! W* K2 | [0 n+ _' y! N) E train(train_dataset, EPOCHS) # N) t. \: z: o# j+ }5 K& F
9 r) M/ y8 P3 R2 p, @0 G, r
# 恢复最新的检查点 # x- X1 p) l9 |: S
checkpoint.restore(tf.train.latest_checkpoint(checkpoint_dir))
% R; J. z9 e: y9 h( O2 ~ , Q2 z* f2 \# y8 q' E9 m
# 评估模型
0 m# Q6 S% o" u; j # 使用 epoch 数生成单张图片
0 ~$ y6 U( n5 T1 z* z& i def display_image(epoch_no):
" n* s3 j5 ~( G& x- [, Z return PIL.Image.open('image_at_epoch_{:04d}.png'.format(epoch_no))
/ t7 y& L6 t1 d, s$ ^, W
7 L& ~5 C- P: M display_image(EPOCHS) 8 X% j( F5 o8 w% c/ ^- }0 T2 Y
: Z2 h# O7 C ^& v) U anim_file = 'dcgan.gif' , G. x1 q% H J
% \% [ R6 C, v2 @- z
with imageio.get_writer(anim_file, mode='I') as writer:
3 B. M3 H( s& g& v7 j/ L" a filenames = glob.glob('image*.png')
f+ f- K v+ D( H. x filenames = sorted(filenames)
" J5 j7 u9 R! ?) L# ~ last = -1
$ ^+ I# r* T7 p' U7 U- t for i,filename in enumerate(filenames): , H/ t, Y/ V L1 w: s- i
frame = 2*(i**0.5) , X# T1 A8 S3 |& z7 y
if round(frame) > round(last):
* m! @) f6 \# y* _! m# M last = frame 9 r+ Y" Z, p) s) d
else:
7 Z2 O: l* Y6 v( K" c continue . `$ I0 Q9 `- j0 j7 v
image = imageio.imread(filename)
: V4 b$ C( h, d( h2 I writer.append_data(image) + L8 V2 g% I- v; U2 F" U* C
image = imageio.imread(filename)
2 u2 d* b# t, M writer.append_data(image) K4 }9 z& Z8 |3 J8 F @" H0 \) B
! O: n3 x) g5 ~5 \
import IPython
1 P+ B5 U) D# B) _9 ` if IPython.version_info > (6,2,0,''):
" @7 t8 k3 @5 R y, n a display.Image(filename=anim_file) ' |0 q3 ]* g! F
参考:https://www.tensorflow.org/tutorials/generative/dcgan
% D: w' ?4 C1 [9 e1 z ————————————————
4 N" u7 J0 z: m5 {( z; |. _' F 版权声明:本文为CSDN博主「一颗小树x」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。
6 Q3 G* |8 [- a7 c5 @ 原文链接:https://blog.csdn.net/qq_41204464/article/details/118279111 1 ]0 ]! u) B0 D- y% i. G# a
7 ?4 i) J7 ]# M# n/ e9 x6 {* r
4 m% O3 `' d" T
zan