AutoencodersTF
在训练CNN时,其中一个问题是我们需要大量的标注数据。以图像分类为例,我们需要将图像分成不同的类别,这通常需要手动完成。
然而,我们可能希望使用原始(未标注)数据来训练CNN特征提取器,这种方法被称为自监督学习。在这种情况下,我们将使用训练图像作为网络的输入和输出。自动编码器的核心思想是,我们会有一个编码器网络,将输入图像转换为某种潜在空间(通常是一个较小尺寸的向量),然后通过解码器网络,其目标是重建原始图像。
由于我们训练自动编码器的目的是尽可能捕捉原始图像中的信息以实现准确的重建,网络会尝试找到输入图像的最佳嵌入来捕捉其含义。

图片来源于 Keras 博客
以下大部分示例都受 这篇文章 的启发。
让我们为MNIST创建最简单的自动编码器:
import tensorflow as tf
from tensorflow.keras.datasets import mnist
import numpy as np
import matplotlib.pyplot as plt
(x_train, y_trainclass), (x_test, y_testclass) = mnist.load_data()def plotn(n,x):
fig,ax = plt.subplots(1,n)
for i,z in enumerate(x[0:n]):
ax[i].imshow(z.reshape(28,28) if z.size==28*28 else z.reshape(14,14) if z.size==14*14 else z)
plt.show()
plotn(5,x_train)from tensorflow.keras.layers import Input, Dense, Conv2D, MaxPooling2D, UpSampling2D, Lambda
from tensorflow.keras.models import Model
from tensorflow.keras.losses import binary_crossentropy,mse
input_img = Input(shape=(28, 28, 1))
x = Conv2D(16, (3, 3), activation='relu', padding='same')(input_img)
x = MaxPooling2D((2, 2), padding='same')(x)
x = Conv2D(8, (3, 3), activation='relu', padding='same')(x)
x = MaxPooling2D((2, 2), padding='same')(x)
x = Conv2D(8, (3, 3), activation='relu', padding='same')(x)
encoded = MaxPooling2D((2, 2), padding='same')(x)
encoder = Model(input_img,encoded)
input_rep = Input(shape=(4,4,8))
x = Conv2D(8, (3, 3), activation='relu', padding='same')(input_rep)
x = UpSampling2D((2, 2))(x)
x = Conv2D(8, (3, 3), activation='relu', padding='same')(x)
x = UpSampling2D((2, 2))(x)
x = Conv2D(16, (3, 3), activation='relu')(x)
x = UpSampling2D((2, 2))(x)
decoded = Conv2D(1, (3, 3), activation='sigmoid', padding='same')(x)
decoder = Model(input_rep,decoded)
autoencoder = Model(input_img, decoder(encoder(input_img)))
autoencoder.compile(optimizer='adam', loss='binary_crossentropy')x_train = x_train.astype('float32') / 255.
x_test = x_test.astype('float32') / 255.
x_train = np.reshape(x_train, (len(x_train), 28, 28, 1))
x_test = np.reshape(x_test, (len(x_test), 28, 28, 1))autoencoder.fit(x_train, x_train,
epochs=25,
batch_size=128,
shuffle=True,
validation_data=(x_test, x_test))Train on 60000 samples, validate on 10000 samples
Epoch 1/25
59648/60000 [============================>.] - ETA: 0s - loss: 0.2134/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/engine/training.py:2325: UserWarning: `Model.state_updates` will be removed in a future version. This property should not be used in TensorFlow 2.0, as `updates` are applied automatically.
warnings.warn('`Model.state_updates` will be removed in a future version. '
60000/60000 [==============================] - 6s 99us/sample - loss: 0.2130 - val_loss: 0.1454
Epoch 2/25
60000/60000 [==============================] - 5s 86us/sample - loss: 0.1353 - val_loss: 0.1258
Epoch 3/25
60000/60000 [==============================] - 5s 87us/sample - loss: 0.1225 - val_loss: 0.1177
Epoch 4/25
60000/60000 [==============================] - 5s 85us/sample - loss: 0.1163 - val_loss: 0.1126
Epoch 5/25
60000/60000 [==============================] - 5s 87us/sample - loss: 0.1120 - val_loss: 0.1091
Epoch 6/25
60000/60000 [==============================] - 5s 86us/sample - loss: 0.1093 - val_loss: 0.1070
Epoch 7/25
60000/60000 [==============================] - 5s 87us/sample - loss: 0.1072 - val_loss: 0.1055
Epoch 8/25
60000/60000 [==============================] - 5s 87us/sample - loss: 0.1057 - val_loss: 0.1041
Epoch 9/25
60000/60000 [==============================] - 5s 85us/sample - loss: 0.1045 - val_loss: 0.1028
Epoch 10/25
60000/60000 [==============================] - 5s 84us/sample - loss: 0.1035 - val_loss: 0.1022
Epoch 11/25
60000/60000 [==============================] - 5s 84us/sample - loss: 0.1026 - val_loss: 0.1011
Epoch 12/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.1018 - val_loss: 0.1003
Epoch 13/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.1012 - val_loss: 0.0996
Epoch 14/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.1005 - val_loss: 0.0991
Epoch 15/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.1000 - val_loss: 0.0988
Epoch 16/25
60000/60000 [==============================] - 5s 82us/sample - loss: 0.0995 - val_loss: 0.0981
Epoch 17/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.0990 - val_loss: 0.0976
Epoch 18/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.0986 - val_loss: 0.0974
Epoch 19/25
60000/60000 [==============================] - 5s 84us/sample - loss: 0.0982 - val_loss: 0.0969
Epoch 20/25
60000/60000 [==============================] - 5s 85us/sample - loss: 0.0978 - val_loss: 0.0970
Epoch 21/25
60000/60000 [==============================] - 5s 84us/sample - loss: 0.0975 - val_loss: 0.0962
Epoch 22/25
60000/60000 [==============================] - 5s 84us/sample - loss: 0.0971 - val_loss: 0.0960
Epoch 23/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.0968 - val_loss: 0.0958
Epoch 24/25
60000/60000 [==============================] - 5s 84us/sample - loss: 0.0966 - val_loss: 0.0953
Epoch 25/25
60000/60000 [==============================] - 5s 83us/sample - loss: 0.0963 - val_loss: 0.0953
<tensorflow.python.keras.callbacks.History at 0x7f3fa179b690>y_test = autoencoder.predict(x_test[0:5])
plotn(5,x_test)
plotn(5,y_test)/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/engine/training.py:2325: UserWarning: `Model.state_updates` will be removed in a future version. This property should not be used in TensorFlow 2.0, as `updates` are applied automatically.
warnings.warn('`Model.state_updates` will be removed in a future version. '
encoder = Model(input_img, encoded)
encoded_imgs = encoder.predict(x_test[0:5])/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/engine/training.py:2325: UserWarning: `Model.state_updates` will be removed in a future version. This property should not be used in TensorFlow 2.0, as `updates` are applied automatically.
warnings.warn('`Model.state_updates` will be removed in a future version. '
plotn(5,encoded_imgs.reshape(5,-1,8))print(encoded_imgs.max(),encoded_imgs.min())
res = decoder.predict(7*np.random.rand(7,4,4,8))
plotn(7,res)6.3110805 0.0
任务 1:尝试使用非常小的潜在向量大小(例如 2)训练自动编码器,并绘制与不同数字对应的点。提示:在卷积部分之后使用全连接的密集层,将向量大小减少到所需值。
任务 2:从不同的数字开始,获取它们的潜在空间表示,观察在潜在空间中添加一些噪声对生成数字的影响。
去噪#
自编码器可以被有效地用于从图像中去除噪声。为了训练一个去噪器,我们将从无噪声的图像开始,并向它们添加人工噪声。然后,我们将带有噪声的图像作为输入,无噪声的图像作为输出,输入到自编码器中。
让我们看看这在 MNIST 数据集上的效果如何:
def noisify(data):
return np.clip(data+np.random.normal(loc=0.5,scale=0.5,size=data.shape),0.,1.)
x_train_noise = noisify(x_train)
x_test_noise = noisify(x_test)
plotn(5,x_train_noise)autoencoder.fit(x_train_noise, x_train,
epochs=25,
batch_size=128,
shuffle=True,
validation_data=(x_test_noise, x_test))Train on 60000 samples, validate on 10000 samples
Epoch 1/25
60000/60000 [==============================] - 6s 101us/sample - loss: 0.1576 - val_loss: 0.1566
Epoch 2/25
60000/60000 [==============================] - 6s 95us/sample - loss: 0.1564 - val_loss: 0.1553
Epoch 3/25
60000/60000 [==============================] - 6s 94us/sample - loss: 0.1555 - val_loss: 0.1539
Epoch 4/25
60000/60000 [==============================] - 6s 95us/sample - loss: 0.1545 - val_loss: 0.1530
Epoch 5/25
60000/60000 [==============================] - 6s 95us/sample - loss: 0.1538 - val_loss: 0.1517
Epoch 6/25
60000/60000 [==============================] - 6s 93us/sample - loss: 0.1528 - val_loss: 0.1506
Epoch 7/25
60000/60000 [==============================] - 6s 93us/sample - loss: 0.1521 - val_loss: 0.1499
Epoch 8/25
60000/60000 [==============================] - 5s 92us/sample - loss: 0.1514 - val_loss: 0.1495
Epoch 9/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1508 - val_loss: 0.1487
Epoch 10/25
60000/60000 [==============================] - 6s 93us/sample - loss: 0.1500 - val_loss: 0.1483
Epoch 11/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1495 - val_loss: 0.1484
Epoch 12/25
60000/60000 [==============================] - 6s 94us/sample - loss: 0.1487 - val_loss: 0.1468
Epoch 13/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1482 - val_loss: 0.1467
Epoch 14/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1476 - val_loss: 0.1459
Epoch 15/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1469 - val_loss: 0.1450
Epoch 16/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1463 - val_loss: 0.1442
Epoch 17/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1457 - val_loss: 0.1441
Epoch 18/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1451 - val_loss: 0.1429
Epoch 19/25
60000/60000 [==============================] - 6s 92us/sample - loss: 0.1445 - val_loss: 0.1425
Epoch 20/25
60000/60000 [==============================] - 6s 94us/sample - loss: 0.1440 - val_loss: 0.1418
Epoch 21/25
60000/60000 [==============================] - 6s 93us/sample - loss: 0.1435 - val_loss: 0.1423
Epoch 22/25
60000/60000 [==============================] - 6s 93us/sample - loss: 0.1430 - val_loss: 0.1409
Epoch 23/25
60000/60000 [==============================] - 6s 94us/sample - loss: 0.1426 - val_loss: 0.1405
Epoch 24/25
60000/60000 [==============================] - 6s 93us/sample - loss: 0.1422 - val_loss: 0.1409
Epoch 25/25
60000/60000 [==============================] - 6s 93us/sample - loss: 0.1418 - val_loss: 0.1398
<tensorflow.python.keras.callbacks.History at 0x7f3fa612c4d0>y_test = autoencoder.predict(x_test_noise[0:5])
plotn(5,x_test_noise)
plotn(5,y_test)练习: 查看在MNIST数字上训练的去噪器如何处理不同的图像。作为一个例子,你可以使用Fashion MNIST数据集,它具有相同的图像大小。请注意,去噪器仅在与其训练时相同类型的图像上效果良好(即输入数据的概率分布相同)。
超分辨率#
与去噪器类似,我们可以训练自动编码器来提高图像的分辨率。为了训练超分辨率网络,我们将从高分辨率图像开始,并自动将其缩小以生成网络输入。然后,我们将小图像作为输入,高分辨率图像作为输出,输入到自动编码器中。
让我们将 MNIST 缩小到 14x14:
x_train_lr = tf.keras.layers.AveragePooling2D()(x_train).numpy()
x_test_lr = tf.keras.layers.AveragePooling2D()(x_test).numpy()
plotn(5,x_train_lr)from tensorflow.keras.layers import Input, Dense, Conv2D, MaxPooling2D, UpSampling2D, Lambda
from tensorflow.keras.models import Model
from tensorflow.keras.losses import binary_crossentropy,mse
input_img = Input(shape=(14, 14, 1))
x = Conv2D(16, (3, 3), activation='relu', padding='same')(input_img)
x = MaxPooling2D((2, 2), padding='same')(x)
x = Conv2D(8, (3, 3), activation='relu', padding='same')(x)
encoded = MaxPooling2D((2, 2), padding='same')(x)
encoder = Model(input_img,encoded)
input_rep = Input(shape=(4,4,8))
x = Conv2D(8, (3, 3), activation='relu', padding='same')(input_rep)
x = UpSampling2D((2, 2))(x)
x = Conv2D(8, (3, 3), activation='relu', padding='same')(x)
x = UpSampling2D((2, 2))(x)
x = Conv2D(16, (3, 3), activation='relu')(x)
x = UpSampling2D((2, 2))(x)
decoded = Conv2D(1, (3, 3), activation='sigmoid', padding='same')(x)
decoder = Model(input_rep,decoded)
autoencoder = Model(input_img, decoder(encoder(input_img)))
autoencoder.compile(optimizer='adam', loss='binary_crossentropy')autoencoder.fit(x_train_lr, x_train,
epochs=25,
batch_size=128,
shuffle=True,
validation_data=(x_test_lr, x_test))Epoch 1/25
469/469 [==============================] - 6s 10ms/step - loss: 0.3413 - val_loss: 0.1519
Epoch 2/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1457 - val_loss: 0.1292
Epoch 3/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1273 - val_loss: 0.1202
Epoch 4/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1189 - val_loss: 0.1142
Epoch 5/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1148 - val_loss: 0.1107
Epoch 6/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1115 - val_loss: 0.1083
Epoch 7/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1093 - val_loss: 0.1063
Epoch 8/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1071 - val_loss: 0.1046
Epoch 9/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1060 - val_loss: 0.1037
Epoch 10/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1048 - val_loss: 0.1026
Epoch 11/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1039 - val_loss: 0.1019
Epoch 12/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1030 - val_loss: 0.1012
Epoch 13/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1024 - val_loss: 0.1004
Epoch 14/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1017 - val_loss: 0.0999
Epoch 15/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1010 - val_loss: 0.0993
Epoch 16/25
469/469 [==============================] - 4s 9ms/step - loss: 0.1005 - val_loss: 0.0989
Epoch 17/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0999 - val_loss: 0.0983
Epoch 18/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0995 - val_loss: 0.0982
Epoch 19/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0990 - val_loss: 0.0975
Epoch 20/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0987 - val_loss: 0.0971
Epoch 21/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0981 - val_loss: 0.0971
Epoch 22/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0979 - val_loss: 0.0965
Epoch 23/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0977 - val_loss: 0.0959
Epoch 24/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0972 - val_loss: 0.0957
Epoch 25/25
469/469 [==============================] - 4s 9ms/step - loss: 0.0972 - val_loss: 0.0955
<tensorflow.python.keras.callbacks.History at 0x7f66790ada90>y_test_lr = autoencoder.predict(x_test_lr[0:5])
plotn(5,x_test_lr)
plotn(5,y_test_lr)练习: 尝试在 CIFAR-10 上训练超分辨率网络进行2倍和4倍上采样。使用噪声作为4倍上采样模型的输入,并观察结果。
变分自编码器 (VAE)#
传统的自编码器通过某种方式降低输入数据的维度,从而提取输入图像的重要特征。然而,潜在向量通常没有太多意义。换句话说,以 MNIST 数据集为例,确定哪些数字对应于不同的潜在向量并不是一件容易的事,因为相近的潜在向量不一定对应相同的数字。
另一方面,为了训练生成模型,理解潜在空间会更有帮助。这一想法引出了变分自编码器 (VAE)。
VAE 是一种自编码器,它学习预测潜在参数的统计分布,即所谓的潜在分布。例如,我们可以假设潜在向量服从分布 $N(\mathrm{z_mean},e^{\mathrm{z_log_sigma}})$,其中 $\mathrm{z_mean}, \mathrm{z_log_sigma} \in\mathbb{R}^d$。VAE 的编码器学习预测这些参数,然后解码器从该分布中随机取一个向量来重建对象。
总结如下:
- 从输入向量中,我们预测
z_mean和z_log_sigma(不是直接预测标准差,而是预测它的对数) - 我们从分布 $N(\mathrm{z_mean},e^{\mathrm{z_log_sigma}})$ 中采样一个向量
sample - 解码器尝试使用
sample作为输入向量来解码原始图像
intermediate_dim = 512
latent_dim = 2
batch_size = 128
tf.compat.v1.disable_eager_execution()
inputs = Input(shape=(784,))
h = Dense(intermediate_dim, activation='relu')(inputs)
z_mean = Dense(latent_dim)(h)
z_log_sigma = Dense(latent_dim)(h)@tf.function
def sampling(args):
z_mean, z_log_sigma = args
bs = tf.shape(z_mean)[0]
epsilon = tf.random.normal(shape=(bs, latent_dim))
return z_mean + tf.exp(z_log_sigma) * epsilon
z = Lambda(sampling)([z_mean, z_log_sigma])encoder = Model(inputs, [z_mean, z_log_sigma, z])
latent_inputs = Input(shape=(latent_dim,))
x = Dense(intermediate_dim, activation='relu')(latent_inputs)
outputs = Dense(784, activation='sigmoid')(x)
decoder = Model(latent_inputs, outputs)
outputs = decoder(encoder(inputs)[2])
vae = Model(inputs, outputs)变分自编码器使用由两部分组成的复杂损失函数:
- 重建损失 是一种损失函数,用于衡量重建图像与目标图像的接近程度(可以是均方误差 MSE)。它与普通自编码器中的损失函数相同。
- KL 损失,确保潜在变量分布接近正态分布。它基于 Kullback-Leibler 散度 的概念——一种用于估计两个统计分布相似程度的度量方法。
@tf.function
def vae_loss(x1,x2):
reconstruction_loss = mse(x1,x2)*784
tmp = 1 + z_log_sigma - tf.square(z_mean) - tf.exp(z_log_sigma)
kl_loss = -0.5*tf.reduce_sum(tmp, axis=-1)
return tf.convert_to_tensor(tf.reduce_mean(reconstruction_loss + kl_loss))
vae.compile(optimizer='rmsprop', loss=vae_loss)x_train_flat = x_train.reshape((len(x_train), np.prod(x_train.shape[1:])))
x_test_flat = x_test.reshape((len(x_test), np.prod(x_test.shape[1:])))
vae.fit(x_train_flat, x_train_flat,
shuffle=True,
epochs=25,
batch_size=batch_size,
validation_data=(x_test_flat, x_test_flat))Train on 60000 samples, validate on 10000 samples
Epoch 1/25
59520/60000 [============================>.] - ETA: 0s - loss: 48.6396/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/engine/training.py:2325: UserWarning: `Model.state_updates` will be removed in a future version. This property should not be used in TensorFlow 2.0, as `updates` are applied automatically.
warnings.warn('`Model.state_updates` will be removed in a future version. '
60000/60000 [==============================] - 4s 64us/sample - loss: 48.5874 - val_loss: 41.8877
Epoch 2/25
60000/60000 [==============================] - 3s 57us/sample - loss: 41.1296 - val_loss: 40.2556
Epoch 3/25
60000/60000 [==============================] - 3s 56us/sample - loss: 40.0063 - val_loss: 39.3692
Epoch 4/25
60000/60000 [==============================] - 3s 56us/sample - loss: 39.2531 - val_loss: 38.7666
Epoch 5/25
60000/60000 [==============================] - 3s 57us/sample - loss: 38.7147 - val_loss: 38.6124
Epoch 6/25
60000/60000 [==============================] - 3s 57us/sample - loss: 38.2962 - val_loss: 38.1867
Epoch 7/25
60000/60000 [==============================] - 3s 56us/sample - loss: 37.9756 - val_loss: 37.9831
Epoch 8/25
60000/60000 [==============================] - 3s 57us/sample - loss: 37.6933 - val_loss: 37.5475
Epoch 9/25
60000/60000 [==============================] - 3s 57us/sample - loss: 37.4323 - val_loss: 37.2913
Epoch 10/25
60000/60000 [==============================] - 3s 56us/sample - loss: 37.2133 - val_loss: 37.1992
Epoch 11/25
60000/60000 [==============================] - 3s 57us/sample - loss: 36.9966 - val_loss: 36.9521
Epoch 12/25
60000/60000 [==============================] - 3s 57us/sample - loss: 36.8204 - val_loss: 36.8431
Epoch 13/25
60000/60000 [==============================] - 3s 57us/sample - loss: 36.6490 - val_loss: 36.6979
Epoch 14/25
60000/60000 [==============================] - 3s 57us/sample - loss: 36.5023 - val_loss: 36.6661
Epoch 15/25
60000/60000 [==============================] - 3s 57us/sample - loss: 36.3456 - val_loss: 36.4957
Epoch 16/25
60000/60000 [==============================] - 3s 56us/sample - loss: 36.2266 - val_loss: 36.6669
Epoch 17/25
60000/60000 [==============================] - 3s 57us/sample - loss: 36.1045 - val_loss: 36.4855
Epoch 18/25
60000/60000 [==============================] - 3s 57us/sample - loss: 35.9922 - val_loss: 36.4150
Epoch 19/25
60000/60000 [==============================] - 3s 56us/sample - loss: 35.8968 - val_loss: 36.1196
Epoch 20/25
60000/60000 [==============================] - 3s 57us/sample - loss: 35.7991 - val_loss: 36.0708
Epoch 21/25
60000/60000 [==============================] - 3s 56us/sample - loss: 35.7129 - val_loss: 36.1686
Epoch 22/25
60000/60000 [==============================] - 3s 57us/sample - loss: 35.6214 - val_loss: 36.1080
Epoch 23/25
60000/60000 [==============================] - 3s 57us/sample - loss: 35.5357 - val_loss: 36.2309
Epoch 24/25
60000/60000 [==============================] - 3s 56us/sample - loss: 35.4528 - val_loss: 36.1416
Epoch 25/25
60000/60000 [==============================] - 3s 56us/sample - loss: 35.3650 - val_loss: 35.7258
<tensorflow.python.keras.callbacks.History at 0x7f0e00233890>y_test = vae.predict(x_test_flat[0:5])
plotn(5,x_test_flat)
plotn(5,y_test)/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/engine/training.py:2325: UserWarning: `Model.state_updates` will be removed in a future version. This property should not be used in TensorFlow 2.0, as `updates` are applied automatically.
warnings.warn('`Model.state_updates` will be removed in a future version. '
x_test_encoded = encoder.predict(x_test_flat)[0]
plt.figure(figsize=(6, 6))
plt.scatter(x_test_encoded[:, 0], x_test_encoded[:, 1], c=y_testclass)
plt.colorbar()
plt.show()/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/engine/training.py:2325: UserWarning: `Model.state_updates` will be removed in a future version. This property should not be used in TensorFlow 2.0, as `updates` are applied automatically.
warnings.warn('`Model.state_updates` will be removed in a future version. '
def plotsample(n):
dx = np.linspace(-1,1,n)
dy = np.linspace(-1,1,n)
fig,ax = plt.subplots(n,n)
for i,xi in enumerate(dx):
for j,xj in enumerate(dy):
res = decoder.predict(np.array([xi,xj]).reshape(-1,2))[0]
ax[i,j].imshow(res.reshape(28,28))
ax[i,j].axis('off')
plt.show()
plotsample(10)/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/engine/training.py:2325: UserWarning: `Model.state_updates` will be removed in a future version. This property should not be used in TensorFlow 2.0, as `updates` are applied automatically.
warnings.warn('`Model.state_updates` will be removed in a future version. '
任务:在我们的样本中,我们已经训练了全连接VAE。现在从上面的传统自动编码器中获取CNN,并创建基于CNN的VAE。
额外材料#
免责声明:
本文档使用AI翻译服务Co-op Translator进行翻译。尽管我们努力确保翻译的准确性,但请注意,自动翻译可能包含错误或不准确之处。原始语言的文档应被视为权威来源。对于关键信息,建议使用专业人工翻译。我们不对因使用此翻译而产生的任何误解或误读承担责任。