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