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