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