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