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