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