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