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