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