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