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