QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 5748|回复: 0
打印 上一主题 下一主题

深度卷积生成对抗网络DCGAN——生成手写数字图片

[复制链接]
字体大小: 正常 放大
杨利霞        

5273

主题

82

听众

17万

积分

  • TA的每日心情
    开心
    2021-8-11 17:59
  • 签到天数: 17 天

    [LV.4]偶尔看看III

    网络挑战赛参赛者

    网络挑战赛参赛者

    自我介绍
    本人女,毕业于内蒙古科技大学,担任文职专业,毕业专业英语。

    群组2018美赛大象算法课程

    群组2018美赛护航培训课程

    群组2019年 数学中国站长建

    群组2019年数据分析师课程

    群组2018年大象老师国赛优

    跳转到指定楼层
    1#
    发表于 2021-6-28 11:54 |只看该作者 |倒序浏览
    |招呼Ta 关注Ta

    ) ]3 @: R# N( `) h# o3 l( x5 r0 ^: e深度卷积生成对抗网络DCGAN——生成手写数字图片
    6 j/ w5 Y+ J2 z, {5 ]前言5 P" f! z9 r' J7 k& F$ {
    本文使用深度卷积生成对抗网络(DCGAN)生成手写数字图片,代码使用Keras API与tf.GradientTape 编写的,其中tf.GradientTrape是训练模型时用到的。+ }* V, x% d3 P( C- W! w7 w

    6 o  I5 L# ?/ {* c6 d; t6 E5 q
    ; d. |0 Z* i# b. C# m
    本文用到imageio 库来生成gif图片,如果没有安装的,需要安装下:
    4 B* ?6 g. b. R! a. K% \2 T0 g  f  a4 Q: A& o8 y% K$ G

    " ?0 j: l( }) d# 用于生成 GIF 图片, c# J" d. E9 {) y: o! z
    pip install -q imageio
    3 c3 `8 K7 g2 Q* `目录! f& \  ~! Y( G# t
    % x3 |# m5 t' D# P; Q: ]: q

    : N4 K0 N/ ^* ^& a( `前言5 B- k9 L: P  H9 A' U

    & l  i( }( ]* }# S

    / y( Z; p% |' `# G8 q一、什么是生成对抗网络?$ g) j; b8 w' Y; N9 r3 \

    ( U# {  D2 }& |1 ^" P
    . H: ~( X0 o" c3 A
    二、加载数据集; j9 r* ?: d9 S5 G) G% P

    1 I% l7 h0 O% S0 c/ G

    0 B* w/ T: f9 d- E) ]" y7 s9 v三、创建模型
    ( o# T6 B9 ?( j6 V+ n, Y1 V5 I$ A# X" ^9 f4 P% p
    4 {7 x6 d# b; U4 Z3 o' H
    3.1 生成器. }1 e9 e0 Y- _5 q
    7 P* K& `8 A' @, l8 V* W& J+ I

    , [; S6 ~' A: c0 P. S3.1 判别器
    / \, n0 r+ y+ n. F7 n, a
    ( G+ s  k' }7 ^  R
    9 {, r; U2 R& o8 l( e  u, _, m
    四、定义损失函数和优化器
    3 R+ \( x+ E- V% o5 k8 b
    1 y& ?" B6 |, B$ O% R

    ) J: f" \# @& {( V, s  g8 A4.1 生成器的损失和优化器7 n4 P. x- D3 N7 P
    / J) O  @. \! s% N
      S4 H5 x7 }; m3 R* E! ]. q
    4.2 判别器的损失和优化器7 ?2 N2 S6 p7 _  m

    7 _$ w# @7 q5 W% O4 O2 o( B/ @# I. ?

    8 P/ l) K' y+ s. N7 H( L# q3 l五、训练模型
    % Y  }8 f: B0 V" w  ^' d, R" l1 s7 B0 ~" E) G- x3 t. o
    0 D  L6 s  R5 q! ^% ?9 l* b
    5.1 保存检查点
    9 I" c  b& a# z+ s% O! ?; K
    , L! {3 X- G! p, e) G+ y

    ; O8 Q+ ~* ~0 p$ @, o5.2 定义训练过程  T* E1 n) r! e2 M3 o- q
    . p, R$ s: G$ i( e

      [/ |% B7 Y9 R5 y2 v5.3 训练模型
    0 b% g; ]7 J3 j! B1 r* R9 M: |% K; ~  |" \/ M  X- }" d
    : ?4 W. v# ~# U' m6 n% W, @
    六、评估模型
    9 Q$ g8 g/ u3 D0 J: z. E; c. ?# X/ K
    # g. g3 h( m6 D

    " U$ e/ T/ ^/ B一、什么是生成对抗网络?
    / Q0 Q1 Y; N! u# }5 K8 z生成对抗网络(GAN),包含生成器和判别器,两个模型通过对抗过程同时训练。
    6 f# G% u9 N% g2 [6 t- i
    " D$ l, r1 w* D& n( B" ^3 `0 k
    : z* h! B# O, D) d
    生成器,可以理解为“艺术家、创造者”,它学习创造看起来真实的图像。/ D7 F; M, ]/ Z% l* J; b; R

    ( F: C8 s- ]# t, o) J

    ! Z. Y1 n( n" y, K  _8 s判别器,可以理解为“艺术评论家、审核者”,它学习区分真假图像。* H% V& m0 h+ A4 G  e: r7 e

    " M7 P3 }6 h* m. R

    9 x# W& o  M+ S. p( q5 c训练过程中,生成器在生成逼真图像方便逐渐变强,而判别器在辨别这些图像的能力上逐渐变强。; @: ]! g. _4 Q! s6 j5 j* o3 \
    & ^; L, E0 B% C8 e0 o
    , N) n2 e( H' I& f9 P" M7 ?+ s+ q
    当判别器不能再区分真实图片和伪造图片时,训练过程达到平衡。
    5 \8 A* u3 o, m* H7 f& \" Q% z1 w0 v: \7 L  G5 R

      T& F6 I* y8 c1 i4 G: _本文,在MNIST数据集上演示了该过程。随着训练的进行,生成器所生成的一系列图片,越来越像真实的手写数字。; P& Q( A$ F6 E' V' b, X+ {+ c% n
    + Q0 H5 W4 Y2 w5 r( Z. P  I/ S
    ( X: U* w2 G* t6 z' X1 O
    二、加载数据集
    0 f* P* N5 A! b1 I7 Z7 w使用MNIST数据,来训练生成器和判别器。生成器将生成类似于MNIST数据集的手写数字。
    2 m- [! u; `$ H* y/ D* i* F1 S6 @; r3 [" o

    4 Y, _- C$ i/ s7 N6 z5 l+ p2 ~(train_images, train_labels), (_, _) = tf.keras.datasets.mnist.load_data()( L* A1 _  h  S3 ?6 R4 t# V8 }7 P7 V

    $ [' L" ]9 Q, D" H: ntrain_images = train_images.reshape(train_images.shape[0], 28, 28, 1).astype('float32')+ {( b# w9 T- q+ C
    train_images = (train_images - 127.5) / 127.5 # 将图片标准化到 [-1, 1] 区间内
    5 ^+ Y; v2 n% Y  s- m+ M( p
    % B0 G) m. u/ G+ f, n" w5 Z& y1 lBUFFER_SIZE = 60000: Y; H7 N  n" X. t
    BATCH_SIZE = 2561 \4 R# E, q9 {. z) H. K# ~) r% K

    1 k8 A. R  p/ s; B7 z" K# 批量化和打乱数据
    # i& x& u' W+ p% N2 S9 Ztrain_dataset = tf.data.Dataset.from_tensor_slices(train_images).shuffle(BUFFER_SIZE).batch(BATCH_SIZE)0 F6 O' V4 n. v% W
    三、创建模型7 o/ p& R: C/ [8 Y' ^" A  X
    主要创建两个模型,一个是生成器,另一个是判别器。
    ; ~  }6 R9 F+ h% _1 \( d) {: z. a3 B2 O* D3 M) c' g+ R

    " O% L6 @% x: _9 [3.1 生成器) M" p! @1 a  E. _$ p5 t* K
    生成器使用 tf.keras.layers.Conv2DTranspose 层,来从随机噪声中产生图片。3 ]$ m" Y9 }8 i
    $ q# J9 a7 m/ F3 g" E( h  O
    1 _! ?! Y9 F2 _$ X2 }
    然后把从随机噪声中产生图片,作为输入数据,输入到Dense层,开始。
    & V" i1 X$ A& n; d) H
    " U) E( ?1 F9 O* K, R6 p  s& E
    5 [% `2 W$ X8 s
    后面,经过多次上采样,达到所预期 28x28x1 的图片尺寸。/ C4 c. R2 a  Z6 ?6 F( i: d
    0 }. Z9 m; D/ ?4 z6 V8 r
    3 U: @0 }, B' W
    def make_generator_model():
    5 S. X% o) _( _+ k0 V4 ?3 }* m2 k    model = tf.keras.Sequential()
    # ?2 l1 F$ c$ T    model.add(layers.Dense(7*7*256, use_bias=False, input_shape=(100,)))+ A& l$ U4 e' |0 G9 I
        model.add(layers.BatchNormalization()), |. U- ~, W8 r: p+ g6 v5 F
        model.add(layers.LeakyReLU())
    ( U7 X2 u  D- \% F4 t
    # ?  N0 U' z3 ]1 c    model.add(layers.Reshape((7, 7, 256)))9 w9 c- h1 h# L
        assert model.output_shape == (None, 7, 7, 256) # 注意:batch size 没有限制, ^7 J; ?  H5 d, g) T* K

    $ \6 n+ O! v# v4 G) a    model.add(layers.Conv2DTranspose(128, (5, 5), strides=(1, 1), padding='same', use_bias=False))6 N/ j8 t; A: M; w
        assert model.output_shape == (None, 7, 7, 128)/ g+ a; ^6 M+ P1 `
        model.add(layers.BatchNormalization())4 K, |/ B; h, a. i
        model.add(layers.LeakyReLU())
      D7 S8 z; n) p4 [3 m; ^/ r . v& N2 j; s2 _/ S' R( ?* C( c# ~
        model.add(layers.Conv2DTranspose(64, (5, 5), strides=(2, 2), padding='same', use_bias=False))' r' W0 W# z% P- l; F' w
        assert model.output_shape == (None, 14, 14, 64)+ o  A: g3 e2 v! A) F
        model.add(layers.BatchNormalization())3 [; g5 k, {  p% Y( J* @) r; M  r
        model.add(layers.LeakyReLU())$ D8 v4 D9 E, \
    9 f5 f" Z- v- i! x, d5 D
        model.add(layers.Conv2DTranspose(1, (5, 5), strides=(2, 2), padding='same', use_bias=False, activation='tanh'))( [- q$ F- f9 O: {6 _+ W
        assert model.output_shape == (None, 28, 28, 1)1 P: J9 B* Z3 y7 ]+ K
    , L* F4 q% h; k( U$ S" v) Z7 e2 U
        return model. k' D4 L8 L9 \& M
    用tf.keras.utils.plot_model( ),看一下模型结构
    5 O) ^) t6 o7 U$ \9 c$ a& U9 {" E4 `8 T
    ; Y0 U# _- O; L6 W6 S- u" q
    " C/ ~$ T/ b  _2 j

    ' Y  B% I+ B3 `& Z9 K/ z( I
    ( L- X# {* U; N# n
    用summary(),看一下模型结构和参数( A4 j& I0 ^+ y# Y# {1 Q
    3 N/ V9 \+ q3 O* V+ u
    / Z: z8 P" X$ }& C

    & s/ ?+ v8 X0 \$ c' i1 ~1 t
    * `1 W" l, a# `# A1 {2 g

    ; y" B  T; O: M5 ~! p. w% o+ M

      X1 k% S7 P4 |, ]8 X; X5 \使用尚未训练的生成器,创建一张图片,这时的图片是随机噪声中产生。
    ) q* G2 ^# x1 r& M# l
    9 i  v8 O; a6 K

    # j. }" A: C9 Z* e, bgenerator = make_generator_model()6 k3 K* G+ s3 `" l
    ' l) u( E0 t2 f* }- N$ \7 J
    noise = tf.random.normal([1, 100])% O) F1 n! w6 H$ {, |( X' e) M
    generated_image = generator(noise, training=False)6 w6 R( h6 ^8 S) c' [

    ' H( s: z9 k8 J3 h' g9 @plt.imshow(generated_image[0, :, :, 0], cmap='gray'): c6 U: ?0 F  V1 @: F, }( a! \

      x9 z! r0 k* d' l+ D, _9 v. i
    8 _7 v7 j, L. H, f( ]
    9 @4 P' C; j6 a7 A  P
    0 U! l7 a* u  I/ |$ t1 ~
    3.1 判别器! A, c1 D. S8 L5 m2 l' S5 H
    判别器是基于 CNN卷积神经网络 的图片分类器。
    * @' m* n, i1 y. R4 s% h' Y5 W7 w4 n! P/ C5 m; G
    : R  w0 d& n# K# B0 N
    def make_discriminator_model():; q. p2 }0 v) `- F
        model = tf.keras.Sequential(). _! X/ d1 j% ~3 Y1 B, T# ~
        model.add(layers.Conv2D(64, (5, 5), strides=(2, 2), padding='same',' }* _# o9 Q' Y# p* y+ B3 ^- O
                                         input_shape=[28, 28, 1]))
    7 A; A2 q2 G, M! j& f    model.add(layers.LeakyReLU())
    9 V  f& z; O0 m) v$ R& H2 j    model.add(layers.Dropout(0.3)). P! X" ?' i* }1 A" T! S/ n

    - y# p6 h6 i( l8 }  p8 ?7 w    model.add(layers.Conv2D(128, (5, 5), strides=(2, 2), padding='same')); I6 j; K3 l, \, J
        model.add(layers.LeakyReLU())
    ; W8 n" z! @" u" g) Q    model.add(layers.Dropout(0.3)); q& q3 t' f0 B6 ?% a- O  W6 k

    " j/ ~" S1 }% v. D* G    model.add(layers.Flatten())" Q. R, B9 V1 I5 Z* q
        model.add(layers.Dense(1))
    7 Z; P5 Z2 a7 i. R
    " D/ o# U3 c# E. ?3 ^+ A    return model
    2 G: X1 |0 R5 g用tf.keras.utils.plot_model( ),看一下模型结构- K4 i* ^* n. a7 ^

    3 `+ w! t9 W  \) B$ q
    * K" ?4 Y' q( U. X9 T  D" _

    2 ~) z) Z5 r3 H
    3 A$ G  ^* J5 ~: z3 Q8 B0 v
    7 p9 z  d0 d) ]2 v7 |) q
    5 e, [- m, `5 E  v1 k
    用summary(),看一下模型结构和参数
    7 {- ?3 b+ B9 N4 D% N
    + F$ R+ C& Z: {& P" c% b' _7 @1 J; W

    ; J7 T5 j7 ]+ t" `! M
    $ S, t2 V# A  e% \4 z  \
    1 ]$ T6 \9 C) h3 H. p: K
    / C; L6 T8 T1 U5 J7 k
    4 ]0 o0 R1 Q" f7 a( f
    四、定义损失函数和优化器  Q7 o8 ?# i  V  l7 ~
    由于有两个模型,一个是生成器,另一个是判别器;所以要分别为两个模型定义损失函数和优化器。
    % u  Q* y( N3 ~& ^& j) V" Y/ A$ C; c3 F

    . r, K2 m) Q" l1 R& J+ ^  F9 L首先定义一个辅助函数,用于计算交叉熵损失的,这个两个模型通用。# Y$ J) L7 r2 Q3 P! e  c% y

    " N7 g& j8 b0 C( r! v
    ) v8 F3 j( k5 Z4 _* R0 g
    # 该方法返回计算交叉熵损失的辅助函数
    ! G, A' L7 J3 j) N" |& K2 I! Ycross_entropy = tf.keras.losses.BinaryCrossentropy(from_logits=True)
    8 e1 S/ a0 G4 C  u4.1 生成器的损失和优化器
    ) W! t) f3 L% z" n: ~0 j0 B5 R) t' e1)生成器损失
    7 P8 j$ K( a- ^* T: _5 C+ ?4 o
    ( l- S5 I3 `& \8 A" d" L) w
    ( S) P0 _& y3 @
    生成器损失,是量化其欺骗判别器的能力;如果生成器表现良好,判别器将会把伪造图片判断为真实图片(或1)。
      G2 l, q! r6 G+ g2 R5 ^# a7 L0 K5 @8 y! V4 ~9 q5 Y3 w* O; A5 W* t

    . |" e% W, S# f9 x1 [, t这里我们将把判别器在生成图片上的判断结果,与一个值全为1的数组进行对比。
    ! R3 O/ H3 W% S: P; _! a4 V; a
    ( `4 d0 ^! h/ O' A0 ?" m2 I% d

    " Y; y- R- U& E6 Udef generator_loss(fake_output):
    - _; i  Q' e" `* C7 I    return cross_entropy(tf.ones_like(fake_output), fake_output)- {0 v' p, j5 R
    2)生成器优化器
    $ i- i: R+ V+ {! W7 |0 }7 Z$ b& C
    - Q! o! J0 R) L9 l4 u) S

    ) \% Y% a( G1 s/ g: r+ vgenerator_optimizer = tf.keras.optimizers.Adam(1e-4)! Y" z, R. w7 `' i' r( x) x- L3 d# E9 w
    4.2 判别器的损失和优化器8 T6 v- g( g; H; _- W/ ]) o
    1)判别器损失$ J: R, c  Y7 L5 a; g9 W

      L+ z. K, Z, Y& H( @0 T
    ) x1 K8 `5 W' S8 B
    判别器损失,是量化判断真伪图片的能力。它将判别器对真实图片的预测值,与全值为1的数组进行对比;将判别器对伪造(生成的)图片的预测值,与全值为0的数组进行对比。: I$ ~" l/ t* k) F3 s* n  D* Q/ i0 U' S2 J
    6 t; ?1 C! O, D: _
    ; Y( {9 N. D' c! g- q* t; f- F# F
    def discriminator_loss(real_output, fake_output):
    1 ?* f: X/ d0 @/ u5 s5 \3 [    real_loss = cross_entropy(tf.ones_like(real_output), real_output)  Z2 a' Z5 F, c" P: Y4 P: }, R6 G
        fake_loss = cross_entropy(tf.zeros_like(fake_output), fake_output)
    # |4 @6 `& i; P$ J" t& v" ]    total_loss = real_loss + fake_loss4 C" Q% O* e% ^' `4 K& t& r' K! p/ P
        return total_loss
    , V7 d( f1 [% C1 ?' I2)判别器优化器4 g- m* @0 S/ j: w: ^
    & T- ^9 A; f7 W3 ]! Q  W
    / `/ a  E3 ~, i( q
    discriminator_optimizer = tf.keras.optimizers.Adam(1e-4)
    0 [. [% j. D1 \- G3 j五、训练模型4 v5 v% A- a: f
    5.1 保存检查点  M' [2 a$ @. G2 a! A: k
    保存检查点,能帮助保存和恢复模型,在长时间训练任务被中断的情况下比较有帮助。
    / j' T+ M5 N; P# e9 o( y0 x
      I6 P" s: _, S) i; v( _8 D
    + K6 o1 X2 P4 ^, F
    checkpoint_dir = './training_checkpoints'
    7 h  |' V  h6 {. Z+ L- ~5 ^checkpoint_prefix = os.path.join(checkpoint_dir, "ckpt")
    + K3 I. \7 E/ ycheckpoint = tf.train.Checkpoint(generator_optimizer=generator_optimizer,
    4 q+ y2 v0 n. m8 T& h8 e                                 discriminator_optimizer=discriminator_optimizer,  d  S  @1 `2 o+ i. O- f
                                     generator=generator,
    0 X4 d6 S1 Q. C* s2 [+ c) R1 @- V                                 discriminator=discriminator)
    , K) d3 M# q5 ?# P3 f5.2 定义训练过程& B, m4 m1 m  V( H3 o
    EPOCHS = 50
    - P: P/ V! z1 z/ `2 V5 E" snoise_dim = 100
    . B9 o9 q4 ^1 inum_examples_to_generate = 16
    7 `  {: o' }' J. O/ s
    $ S4 F. ~3 T- T' i/ h# ^" d, O+ W + ?( R2 }* S1 m. T% X! S0 Q
    # 我们将重复使用该种子(因此在动画 GIF 中更容易可视化进度)  ?, e/ X2 W. `
    seed = tf.random.normal([num_examples_to_generate, noise_dim])
    ' k4 x) i& d. P; E7 x. D训练过程中,在生成器接收到一个“随机噪声中产生的图片”作为输入开始。' H7 R% ^8 j. @& M. s

    , t2 s1 o3 M7 ~9 T* |- H- z
    2 _' B4 {: R5 u9 P
    判别器随后被用于区分真实图片(训练集的)和伪造图片(生成器生成的)。
    6 ?; T$ Z( ^4 {8 p$ F4 J; _9 V, ^% L- @1 O4 M

    2 G& F( b( k- g9 ]" z" ^% y两个模型都计算损失函数,并且分别计算梯度用于更新生成器与判别器。  g0 ^2 z4 I# D3 D" x3 s

    5 _- c4 [+ \: b9 |9 @/ y+ a& s

    ; t) P& B: {  r5 g5 |# 注意 `tf.function` 的使用; C! ]4 b9 B9 Y! J% F0 K
    # 该注解使函数被“编译”
    6 z- T7 N& d  d. G5 I) Z@tf.function. G9 _8 D6 ]7 `4 H
    def train_step(images):7 `9 D: F" l" ~- T' i
        noise = tf.random.normal([BATCH_SIZE, noise_dim])
    6 @2 J3 s: x; r. E+ H: U + z: a$ m5 r- K; h
        with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape:
    . p0 V" l  N+ E. z& D2 q7 |5 W* N; j( t      generated_images = generator(noise, training=True)9 W8 d6 p9 v2 N9 L
    % T3 b* `5 n" y% o1 b0 g5 V/ P4 v+ P
          real_output = discriminator(images, training=True)
    ' L/ E! f& U( O/ g      fake_output = discriminator(generated_images, training=True)
    & _& Y5 B  c0 v+ Q
    & Z3 A4 A: ~3 ?- a- ^# }      gen_loss = generator_loss(fake_output)
    8 o0 d0 J9 L+ j      disc_loss = discriminator_loss(real_output, fake_output)
    9 J0 }1 k$ e0 J% L 6 e. [5 ~9 Y6 R( u/ x  A, H
        gradients_of_generator = gen_tape.gradient(gen_loss, generator.trainable_variables)) k8 h7 D  A' {2 b- R
        gradients_of_discriminator = disc_tape.gradient(disc_loss, discriminator.trainable_variables)
    ' Q, H! r5 Y3 @9 t 8 I% C, C6 r5 v. q
        generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables)): i7 g: \8 A/ J6 M# G5 |
        discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables))8 A$ M7 `  i4 V% |5 }; R* N8 `
    3 v6 c0 {. q8 H- F6 W
    def train(dataset, epochs):  q; V+ Z  i$ a2 v. @* P) R& f
      for epoch in range(epochs):( `0 a# w5 j( ?+ H
        start = time.time()6 |$ |: o0 ~& E
    . `' j# I! t/ @5 \3 g, M+ d  _
        for image_batch in dataset:5 P. n2 W* f' q4 h% _; I
          train_step(image_batch)
    ( a: h3 o, P4 U ' T: p2 S& f. O) S1 u% V
        # 继续进行时为 GIF 生成图像% M: i# u) B; {, t; }9 |* N
        display.clear_output(wait=True)
    $ X; y) g7 X. t3 p7 K2 k    generate_and_save_images(generator,
    + b* Z8 n3 U1 k: Y& l2 T: O' n                             epoch + 1,
    & U% ~8 _& \1 @. V                             seed)# ^, \' Q/ h" q# {9 H7 a! c% `4 q

    # j+ T* p8 z: s' `% M# ~) r    # 每 15 个 epoch 保存一次模型4 d+ D5 B- ~$ b7 y8 K! W
        if (epoch + 1) % 15 == 0:: |6 Q& `( s/ K
          checkpoint.save(file_prefix = checkpoint_prefix)
    4 F/ u! r2 V: W , h3 \/ _6 t( d7 ~# z
        print ('Time for epoch {} is {} sec'.format(epoch + 1, time.time()-start))" E, `  L, v3 c
    : W( G' R9 h1 H2 \: U3 a
      # 最后一个 epoch 结束后生成图片
    ; B9 b+ l, m; m: x! b9 `. _  display.clear_output(wait=True)# @7 T. h* m4 q$ B; U+ e
      generate_and_save_images(generator,- \! u! F7 f! P  W" A$ _8 s
                               epochs,: t5 k7 u+ g0 t- K
                               seed)
    $ a# g5 _- f+ K$ u & Z0 m# P: @0 B, q& s3 Q
    # 生成与保存图片5 Z+ c( a/ t9 o" b6 K
    def generate_and_save_images(model, epoch, test_input):9 V$ [9 t8 s. s) f6 t
      # 注意 training` 设定为 False
    ( Z0 ?5 O: L2 a7 {+ P% k  # 因此,所有层都在推理模式下运行(batchnorm)。2 T9 ]# H. S6 F7 o' u
      predictions = model(test_input, training=False); \4 k2 Q1 {2 ]) U* Q1 E6 q+ W
    ; h# n; y$ w: r6 W3 Y- c
      fig = plt.figure(figsize=(4,4))) P7 |  J; ?: ?9 @; [- |& ?
    " Q6 W( |8 k' c
      for i in range(predictions.shape[0]):
    1 l* C6 K, Z$ ?2 U3 w      plt.subplot(4, 4, i+1)
    ; B6 {# [+ }( Z6 t7 j      plt.imshow(predictions[i, :, :, 0] * 127.5 + 127.5, cmap='gray')
    + e7 M/ y& u/ M3 U9 F      plt.axis('off')
      l: F7 J0 a' _1 U  x1 e
    # [( p0 k# i  C1 l0 K3 V  plt.savefig('image_at_epoch_{:04d}.png'.format(epoch))
    , v# W6 `/ e' e3 l6 h7 o1 C5 |  plt.show()
    $ A% f( Z% O* x! j$ v' Q5.3 训练模型
    ) W6 F# D1 w# W! F8 s7 Y调用上面定义的train()函数,来同时训练生成器和判别器。
    7 I, j9 r* ~, I1 C" E
    3 s9 |" s, S" L. V- G6 h

    * f3 O4 ?, w# U4 q; b& b) e注意,训练GAN可能比较难的;生成器和判别器不能互相压制对方,需要两种达到平衡,它们用相似的学习率训练。8 A) M2 [  H' L1 \, _. ]  `  l3 n* N

    ) b9 w$ Z/ @5 m3 y6 c4 M. D
    ! M! b- n. h5 i* `+ M# g
    %%time
    4 e3 n4 p. }% d- R1 t+ v: A! K9 J# itrain(train_dataset, EPOCHS)( X+ B; p8 `4 M" J0 J) M3 N& N- t
    在刚开始训练时,生成的图片看起来很像随机噪声,随着训练过程的进行,生成的数字越来越真实。训练大约50轮后,生成器生成的图片看起来很像MNIST数字了。* r: s9 y7 q8 P

    + r# E! @! \' |. j" `2 O' Y) L

    $ ^* y; ~& C" G; g+ P% U训练了15轮的效果:" C  m; d% s7 j! y6 @9 h  ^: E8 n
    5 f" x8 ?# x0 Z6 S: T
    6 D( n; J) k4 t5 b0 U8 ^4 g

    ! b$ r1 m) e" r6 I/ e' v
    8 e. M% S, w9 n
    + d# \  S2 B* J6 m& J6 w; @% m4 L2 A& T
    ( E+ N+ r% `; }1 n2 k5 S
    训练了30轮的效果:
    / G3 x" d+ B- z5 W7 |( ^$ X- n
    $ `, K" `) b( P  W+ _7 I

    ( Y: O0 B1 ?7 U7 U& h8 k2 G. o8 Z# `. {! Q! f
    % n7 H$ [7 {% Q- k" r6 l# e8 U" L

    * ^" B/ g1 ~* ^* P  R0 T. r

    9 l" V# L# h& @$ m# d' l/ A1 G: T1 K训练过程:
      K" m  p' u- a6 D' ~1 J' H8 T3 P
    $ O# `  v& Q7 U6 T! [

    ! U) ~, B/ q+ X( G% C
    ) h9 f% Y9 w/ N' E9 |2 W
    , B& ?! c/ t% D0 }

    # _$ h% w# Q/ r* G

    , ?5 `8 y. |$ m恢复最新的检查点
    3 r9 O5 P8 w- z+ ~6 o
    - l# |  |  d7 M( h9 W

    : L- d7 z$ m! R' W$ X% Q/ ucheckpoint.restore(tf.train.latest_checkpoint(checkpoint_dir))
    + d8 A) t4 e4 @六、评估模型
    8 N7 ]4 [6 C0 U: T5 M这里通过直接查看生成的图片,来看模型的效果。使用训练过程中生成的图片,通过imageio生成动态gif。
    ( z1 ~" u7 n# V" t& R
    0 [; |) `0 d& F9 W& ?& M) p- W

    # C3 f5 U9 Q  ^; ]# r# 使用 epoch 数生成单张图片* i2 v' P; o; ~8 j& F6 m1 d& N( k
    def display_image(epoch_no):0 p' {5 @6 F0 h- e9 A& x( a
      return PIL.Image.open('image_at_epoch_{:04d}.png'.format(epoch_no))
    0 ?+ j& H( Z  Z3 X / i$ C& A" d# \& n3 m
    display_image(EPOCHS)' l9 G7 }# `, C9 i! Y2 t
    anim_file = 'dcgan.gif'
    2 ?% Y  }* j/ t/ z& Q2 n + O; R8 T# s/ E7 [/ J$ `2 D1 L. z
    with imageio.get_writer(anim_file, mode='I') as writer:
    $ p; F7 r1 ^- Z1 `. t# g9 Z  filenames = glob.glob('image*.png')7 [8 _& C, b8 l
      filenames = sorted(filenames)$ X# @( U  O1 S( Y  Y. T6 \
      last = -1
    % ~& Q  s- b) z. B0 u  h6 i  for i,filename in enumerate(filenames):
    7 r( r5 C: R" }( `4 V    frame = 2*(i**0.5)2 g" O% R+ y; \9 `% o4 E8 |' l2 n
        if round(frame) > round(last):
      B5 j9 C4 @& d/ D0 c0 J      last = frame7 O6 w- K- t% O* C( N
        else:
    7 c" g8 o/ s5 X- o/ P* ~7 O$ l      continue
    8 r% k0 q+ _% Q. q% ]    image = imageio.imread(filename)
    8 x8 N( i( Y& d# A# M! }    writer.append_data(image)' \6 p3 M" T- g5 Q, N
      image = imageio.imread(filename)8 m# C* q$ ~8 x+ o# ?( F% b% v+ y
      writer.append_data(image)
    : @  O4 u8 }( O; ^5 o; x6 L0 S
    5 l' `+ U/ ?  e- `import IPython9 k7 `+ z5 i# }- m9 J; X2 ^2 ~
    if IPython.version_info > (6,2,0,''):8 Z- h4 L/ P& h& ^. u: K- n
      display.Image(filename=anim_file)3 B' U9 B4 b/ |5 e( W6 \/ A& M
    # b7 r( g* j' W* T1 [; V

    1 t. C6 b3 `5 z4 u9 F' j8 N! O' k! s0 ?4 Z: m8 V( M7 s6 X" z; N7 L& d

    . P1 [; a. J, p) t3 [完整代码:
    ! U9 I( r; l, M( j+ |# [9 C! F: I1 f; I8 ?/ C  Q1 d- m
    5 ]7 v, N3 E6 L
    import tensorflow as tf
    1 `/ s% L2 Z' B: I6 s1 Timport glob
    - ~/ p" H% c3 _. w1 _5 R3 Uimport imageio& |: ?: u/ h5 W  p
    import matplotlib.pyplot as plt( \+ V5 t7 @, |
    import numpy as np
    2 {( C4 i% F5 B) N3 V6 L9 j) nimport os$ v8 d1 r; [# |2 s
    import PIL7 Z9 a8 V  m1 @0 }
    from tensorflow.keras import layers& y/ I! y4 Y. @/ |& q; ~" s% m
    import time
    8 Z7 w0 f& S0 d+ c" V
    1 o- Z. d% L% M1 z# g3 y, i: ufrom IPython import display
    7 \& ], n) R9 h! ~4 I1 \9 I
    % R- N0 g3 h# a5 Q(train_images, train_labels), (_, _) = tf.keras.datasets.mnist.load_data()( T+ ~$ ?# Z/ S) t' y
    ) y' m/ f, f2 ], b% ]! ?
    train_images = train_images.reshape(train_images.shape[0], 28, 28, 1).astype('float32')
    , [  }! i- V, }& o! B2 mtrain_images = (train_images - 127.5) / 127.5 # 将图片标准化到 [-1, 1] 区间内
    0 V5 I' ^1 R$ G' J/ z' v7 V. ~3 V 8 W& M: E% v" l% t
    BUFFER_SIZE = 60000& F# `6 A  M" F# a( S
    BATCH_SIZE = 256
    4 a" ?) E$ c! Q, X2 o
    , F, [, B& D1 x6 Y( P. F' u# 批量化和打乱数据: l5 }. L1 p: M2 B4 k
    train_dataset = tf.data.Dataset.from_tensor_slices(train_images).shuffle(BUFFER_SIZE).batch(BATCH_SIZE)
    ! \/ K5 u' U- z2 I2 S" a6 U: b . N4 n; m/ H) Q# p7 f! i8 L
    # 创建模型--生成器
    ! U; j5 |# s; Z. f7 m# q3 `8 q2 odef make_generator_model():
    - Z* \) J4 E- m1 z, U. D    model = tf.keras.Sequential()+ ^+ u/ m, q, L0 M( M# {8 {
        model.add(layers.Dense(7*7*256, use_bias=False, input_shape=(100,)))1 j( v+ m' l  B4 J
        model.add(layers.BatchNormalization())" N" ~3 r5 @$ V1 v1 F
        model.add(layers.LeakyReLU())8 c2 p" x1 K) Y( o: v% \; H' E9 w+ H

    0 ~2 g: R, Y8 o) G7 N4 R/ F- T: @    model.add(layers.Reshape((7, 7, 256)))
    7 Q' s' T8 v* M- Q" }8 m    assert model.output_shape == (None, 7, 7, 256) # 注意:batch size 没有限制1 g! \' P- N$ q/ V1 L
    / N. E* I4 h3 s6 N" J
        model.add(layers.Conv2DTranspose(128, (5, 5), strides=(1, 1), padding='same', use_bias=False))
    / E1 ?. y3 w1 X7 p' e    assert model.output_shape == (None, 7, 7, 128)9 H; ?3 s7 W" t8 @) L1 g, c6 ^0 P# E
        model.add(layers.BatchNormalization())7 C! m6 E7 C; s6 E8 ?% L
        model.add(layers.LeakyReLU())% U3 C% Z; N0 |3 h9 d

    ( p, S  p. h; G6 ]    model.add(layers.Conv2DTranspose(64, (5, 5), strides=(2, 2), padding='same', use_bias=False))
    0 S: S2 [' m: W  |' A1 q# N    assert model.output_shape == (None, 14, 14, 64)
    6 t, g- A5 m" p# Y/ a/ i+ z    model.add(layers.BatchNormalization())
    " b) P7 F7 B5 K3 @; v5 H( l    model.add(layers.LeakyReLU()). a. l. A& n3 L: `  Q
    # D* z* i0 y( G. N8 J8 ^
        model.add(layers.Conv2DTranspose(1, (5, 5), strides=(2, 2), padding='same', use_bias=False, activation='tanh'))
    4 e! y% G( a, s' u    assert model.output_shape == (None, 28, 28, 1). W5 e0 @9 j3 L3 v5 Q

    4 j7 C" H. W6 h. G2 M9 t$ o    return model
    ; B8 C( l/ t* b) i! A. Y
    , h9 I1 q  J) H; Y) t) t/ i" Y3 j# 使用尚未训练的生成器,创建一张图片,这时的图片是随机噪声中产生。+ w/ f. j  }8 {8 L  i  G" d
    generator = make_generator_model()
    5 p: \  @$ S2 V" a! Z8 Y
    $ Z* M5 Y0 S9 P  [3 v9 anoise = tf.random.normal([1, 100])
    6 [# d5 V- y% V* K; m: ^$ pgenerated_image = generator(noise, training=False)
    1 h; h) }6 O/ |* j- X* H
    8 B5 r; Y  L. nplt.imshow(generated_image[0, :, :, 0], cmap='gray')( E3 K& h: L+ X
    tf.keras.utils.plot_model(generator)' T# D) m/ h3 Z- W, C; _" T

    " X) i0 p4 y) i% B5 h2 M# 判别器1 _4 B& T5 L! f; U) b. q+ X( E
    def make_discriminator_model():
    , k8 `9 P! m1 F    model = tf.keras.Sequential()9 S* f' a7 n3 ]" R- K+ K0 N+ [
        model.add(layers.Conv2D(64, (5, 5), strides=(2, 2), padding='same',
    8 l5 u8 {) z2 A% d3 a                                     input_shape=[28, 28, 1]))
      Q5 \! A2 d& l  }    model.add(layers.LeakyReLU())
    % w& Q+ e0 ?6 m( @/ w" P    model.add(layers.Dropout(0.3))( S! |* y5 N3 t0 x. p% G

    $ X/ @1 _: D: i2 p) Q# F    model.add(layers.Conv2D(128, (5, 5), strides=(2, 2), padding='same'))
    5 K3 d' \; N" q* ?# {: R    model.add(layers.LeakyReLU())
    2 J0 X( t- y3 o6 b$ E    model.add(layers.Dropout(0.3))
    ; G3 y9 s# p& ]) s/ A: K# B8 r2 f9 P # K2 h" n2 A7 Y$ G) Q' ?' m
        model.add(layers.Flatten())
      H0 Y8 n8 z3 E1 i# ]    model.add(layers.Dense(1))" a! h" B+ p0 `& D
    - ~2 j& E2 U, P0 f* C4 G, }
        return model2 o+ ~- L- l4 \+ `9 S
    ) Q' A! v- A9 o- `( W, x
    # 使用(尚未训练的)判别器来对图片的真伪进行判断。模型将被训练为为真实图片输出正值,为伪造图片输出负值。
    & Q( q7 {( r: G! v* Z9 a6 Cdiscriminator = make_discriminator_model()) u) ]' N7 t  ~- z3 C6 k
    decision = discriminator(generated_image)
    ( S# O- a- |/ B# l  xprint (decision)3 k& e  w, D3 k6 Q8 ?& |
    ) x% m# Z6 [2 H! o
    # 首先定义一个辅助函数,用于计算交叉熵损失的,这个两个模型通用。
    : N1 W  J) U5 f* w( Z* wcross_entropy = tf.keras.losses.BinaryCrossentropy(from_logits=True)) T+ j) P+ \2 i1 o) N3 U5 z
    " X- |* v/ N/ V6 J
    # 生成器的损失和优化器5 x% ?# K9 ~4 I) m* E% V
    def generator_loss(fake_output):9 h6 _$ M( `1 f. p' G
        return cross_entropy(tf.ones_like(fake_output), fake_output)
      ]6 j/ [- L8 V2 r3 v2 Wgenerator_optimizer = tf.keras.optimizers.Adam(1e-4)
    0 e+ R  H% S  w4 s
      i; T1 e7 y7 {# 判别器的损失和优化器
    + Q/ |1 T7 Y8 t% F: _2 c# u5 tdef discriminator_loss(real_output, fake_output):
    " p2 ^' R4 b* z7 l6 X    real_loss = cross_entropy(tf.ones_like(real_output), real_output)" ^) p) L' u& ?4 j: H3 W  _, ~9 z
        fake_loss = cross_entropy(tf.zeros_like(fake_output), fake_output)
    0 S6 ], v' R6 I    total_loss = real_loss + fake_loss
    5 A- b- |( s& ?- c3 I: Y! l    return total_loss" o  P4 T' m$ f- S, w) {; w
    discriminator_optimizer = tf.keras.optimizers.Adam(1e-4)0 p8 Y5 [  M- e1 f3 H. f" Q$ H

    ) P% r$ o2 v4 L9 r: y: z% Y# 保存检查点2 T. |' ~0 |' g& k3 g
    checkpoint_dir = './training_checkpoints'- R9 A& {3 p. T+ N( L# l) g3 v
    checkpoint_prefix = os.path.join(checkpoint_dir, "ckpt")
    + ^* n* n. ^4 X4 [0 echeckpoint = tf.train.Checkpoint(generator_optimizer=generator_optimizer,
    # p8 f0 F8 t% g2 c. P5 ~9 H                                 discriminator_optimizer=discriminator_optimizer," K+ r0 O) H" Z; ?, D9 y
                                     generator=generator,- t" G8 O, K  I  X0 w4 |
                                     discriminator=discriminator)
    # b' O9 W; f7 ] 9 Y& n3 q6 \( l/ J
    # 定义训练过程
    ( S: X$ \: q( j. J6 R0 {5 P9 Y. L& LEPOCHS = 50
    * z& \7 d* X: [, e- O, Inoise_dim = 100
    3 N% P. c* Z4 q+ dnum_examples_to_generate = 163 B5 i1 F- L5 A" j9 L! p0 C3 i) x

    ) L& A; ^# m# W# 我们将重复使用该种子(因此在动画 GIF 中更容易可视化进度)
    5 P0 \  _8 x% j$ ~5 R, I9 {# O/ s+ {seed = tf.random.normal([num_examples_to_generate, noise_dim])
    5 Q$ @& Y2 A( P) _/ x 3 P4 L& ~; o3 A  o/ g
    # 注意 `tf.function` 的使用, F0 E/ \" T$ I' t$ W- v" j
    # 该注解使函数被“编译”, {  }+ F1 l4 S: y+ _3 j9 i
    @tf.function
    ' ?' t* Q1 u) D+ r6 H( \! Vdef train_step(images):7 K0 W. i# V+ t" x; A7 A. h
        noise = tf.random.normal([BATCH_SIZE, noise_dim])0 R& T/ t) Z6 X
    7 e/ H( a* [8 {6 U1 e
        with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape:
    4 U4 g9 Z0 v) l* u& P$ C3 H      generated_images = generator(noise, training=True)
    6 |1 L$ i7 J, m# n& [7 k * }7 r& n4 ^2 m- ?% N2 f9 ]5 \
          real_output = discriminator(images, training=True)1 v; H4 h1 k3 j' z6 {  Q; y9 c" Y
          fake_output = discriminator(generated_images, training=True)
    % d1 o" b- A) O- S
      E! @% [! O/ u      gen_loss = generator_loss(fake_output)
    0 ?8 I7 I% W5 n0 L1 E3 c; Y      disc_loss = discriminator_loss(real_output, fake_output)3 p) H/ F6 K/ R4 R
    3 d1 K  L6 q- p/ g/ Y5 A
        gradients_of_generator = gen_tape.gradient(gen_loss, generator.trainable_variables)
    # ?! B% I# j$ P$ m7 f    gradients_of_discriminator = disc_tape.gradient(disc_loss, discriminator.trainable_variables)
    3 [- A0 K  |" H) A) J
    4 ]! @9 D  [" v: d' B    generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables))
    ; M+ _% f2 P6 |. ~5 @5 T    discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables))6 q$ @3 i( s" S- ?
    ( g  Y! w# `3 m+ y2 w
    def train(dataset, epochs):$ `' B9 x0 r0 n: |: F- W
      for epoch in range(epochs):
    ; L, U! }: b" Q, d) \, v% O) j    start = time.time()* o" q' W2 f. I! Q8 R) l! {1 {
    % k) b6 o. f6 A, r  m4 z
        for image_batch in dataset:  U8 W5 {3 `8 C4 V2 ~
          train_step(image_batch)
    . O( E5 ^  m9 G
    9 C+ l6 B+ R  b+ `0 _% z    # 继续进行时为 GIF 生成图像
    4 @. g' n# i2 Z* u0 m* k    display.clear_output(wait=True)9 ~9 i8 E" ?( t+ p& E. F$ J
        generate_and_save_images(generator,) ?2 ~6 H8 ^( p) u  J
                                 epoch + 1,
    6 q( \/ p% S0 N8 W+ U4 L                             seed)" M* \4 }5 t3 y, u# P

    ! s$ U  @! K, S    # 每 15 个 epoch 保存一次模型: j, l0 R; U1 f. K) i4 ]+ Q
        if (epoch + 1) % 15 == 0:
    % z! C9 q; O+ P) j& l$ r4 Q      checkpoint.save(file_prefix = checkpoint_prefix)6 |' r) {/ n/ J. R" d5 |3 v6 n, s

    9 T( Z- n" @3 k3 g; s( V    print ('Time for epoch {} is {} sec'.format(epoch + 1, time.time()-start))* {8 Q' @8 f% N2 m  ?+ N% b
    + z" s. x2 b) s0 Q/ Y
      # 最后一个 epoch 结束后生成图片
    , f' {8 m- l$ i* N' Q, F* P7 ~  display.clear_output(wait=True)5 E+ X: F' D. N# f" s1 V+ u7 \7 l
      generate_and_save_images(generator,- R7 X4 i5 I8 U) T. m# j  V( `
                               epochs,
    ' Q8 J) N/ }. X; W$ `( w) Q                           seed)
      a3 O, t; c+ C& ~& ? : c. R# m1 X( G9 {6 G
    # 生成与保存图片
    " ?: a1 m* V( Odef generate_and_save_images(model, epoch, test_input):# g  H) w! X; A* \+ ]
      # 注意 training` 设定为 False1 I4 m; N  n% v5 i& x. [% W% i
      # 因此,所有层都在推理模式下运行(batchnorm)。
    : v1 A. g9 y4 w- g" M6 b1 w# n+ Q  predictions = model(test_input, training=False)
    7 z, Y' t; U; a. m( E. Z . g9 G: e9 y0 E2 ^8 U; e& N- F
      fig = plt.figure(figsize=(4,4))$ D8 n# }% F: y3 z: N! w( x
    # I" s2 K% B" B4 C
      for i in range(predictions.shape[0]):" u9 I* f" N- m, n5 H9 k! j
          plt.subplot(4, 4, i+1)
    9 O* x, W* o$ [  _+ V      plt.imshow(predictions[i, :, :, 0] * 127.5 + 127.5, cmap='gray')! \# e/ _9 h+ D$ Y8 L
          plt.axis('off')5 F9 ~+ a' \9 k8 u

    ! T% l" R/ N2 t5 V6 O  plt.savefig('image_at_epoch_{:04d}.png'.format(epoch))
    ' c) I. }& s6 H3 d& Z# p! O  plt.show()
      B3 v8 W, `3 r: j3 @8 s% { ( L4 ^% o& D% e0 X2 v; q  n8 r
    # 训练模型' O- `, s4 \# g2 r$ T( T3 X# b
    train(train_dataset, EPOCHS)1 c: @; ?" }. s3 g% ~, T; e
    - H9 C- F  h2 Q' S  T2 d3 U
    # 恢复最新的检查点
    7 c: [. @" {; u$ u6 kcheckpoint.restore(tf.train.latest_checkpoint(checkpoint_dir)). M; i- {8 y$ {5 c% w4 p& S
    2 Q% Y# \  w9 ?  Z& j9 B
    # 评估模型
    7 v2 K/ t! u8 V9 a4 C' F# 使用 epoch 数生成单张图片. L5 i# R4 u6 k1 \$ @
    def display_image(epoch_no):
    " s. }2 g! {0 k4 d- G  return PIL.Image.open('image_at_epoch_{:04d}.png'.format(epoch_no))' m6 _+ z8 v$ @# P3 R" I3 d

      L/ |" o% k) g# t8 z: {' Qdisplay_image(EPOCHS)4 x: b0 F6 H, x6 c4 d
    $ Q9 Z5 W4 n; C( @; @
    anim_file = 'dcgan.gif'
    . `! q) U3 m& }+ H7 I2 c6 h
    1 ^0 ~* T" J- Dwith imageio.get_writer(anim_file, mode='I') as writer:& c% W! X7 F; A9 `/ {
      filenames = glob.glob('image*.png')
    : g$ B0 k# {( @( P  filenames = sorted(filenames)
    ) r& N' t3 e- [# `  last = -1' x4 k& A' S5 F4 R* V% U
      for i,filename in enumerate(filenames):
    : M6 @* ?$ d: x1 g4 [    frame = 2*(i**0.5)! r- q4 R# Y+ n" A4 L7 n% ~+ D8 f
        if round(frame) > round(last):
    - M5 x6 m! A* h! D. U/ Q$ {1 u; v% z, l      last = frame' A: f1 F4 V9 b9 a, A8 s* @  _
        else:# o- ^$ y7 p/ d& [) j) F6 I8 R9 _
          continue
    8 o# Z5 R* r4 X6 }8 _  P    image = imageio.imread(filename)
    1 |: }, H0 U  H- u/ u) Z    writer.append_data(image)) A1 S9 U8 ~( \* d0 c, j+ \; M
      image = imageio.imread(filename)
    9 N, h0 {( }0 I! Y* I% W  writer.append_data(image)9 `, W& e! n4 R0 t2 F& K; P% z

    % _5 t; x& Q  ?9 _9 `import IPython/ j, f# l5 t* s: q! l! p  \5 V, j
    if IPython.version_info > (6,2,0,''):
    + q7 `- F8 h/ Z- R  display.Image(filename=anim_file)
    % w! }9 w$ D1 X2 e2 \+ f- V  [参考:https://www.tensorflow.org/tutorials/generative/dcgan
    ; [4 A, ]+ k5 A( Q; @————————————————; T+ h5 W1 E1 y0 U! [) T' g
    版权声明:本文为CSDN博主「一颗小树x」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。) }; l: x  B0 w& |. E
    原文链接:https://blog.csdn.net/qq_41204464/article/details/1182791115 {* r8 Y9 J/ Q7 W

    6 a2 j) L: |3 U7 K( r$ M6 U/ y# Y6 D. r( d' E
    zan
    转播转播0 分享淘帖0 分享分享0 收藏收藏0 支持支持0 反对反对0 微信微信
    您需要登录后才可以回帖 登录 | 注册地址

    qq
    收缩
    • 电话咨询

    • 04714969085
    fastpost

    关于我们| 联系我们| 诚征英才| 对外合作| 产品服务| QQ

    手机版|Archiver| |繁體中文 手机客户端  

    蒙公网安备 15010502000194号

    Powered by Discuz! X2.5   © 2001-2013 数学建模网-数学中国 ( 蒙ICP备14002410号-3 蒙BBS备-0002号 )     论坛法律顾问:王兆丰

    GMT+8, 2026-7-28 21:54 , Processed in 0.388157 second(s), 51 queries .

    回顶部