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