基于Tensorflow2的基本自编码器实现(MNIST)
基于Tensorflow2的基本自编码器实现(MNIST)
关于自编码器的知识这里暂不做过多介绍,我们直接在手写数字集MNIST上进行演示效果。
1. 导包
import tensorflow as tf # 2.0 import matplotlib.pyplot as plt
2. 数据准备
# 加载数据 (x_train, _), (x_test, _) = tf.keras.datasets.mnist.load_data() # 不需要标签,因此使用占位符 # x_train.shape 为(60000, 28, 28) # x_test.shape 为(10000, 28, 28) x_train = x_train.reshape([x_train.shape[0],28*28]) # 改变shape x_test = x_test.reshape([x_test.shape[0],28*28]) # x_train.shape 为(60000, 784) # x_test.shape 为(10000, 784) # 归一化为0-1之间的数据(愿MNIST数据为0-255之间) x_train = tf.cast(x_train,tf.float32)/255 x_test = tf.cast(x_test,tf.float32)/255
3. 模型创建
将原784维的数据压缩至32维,再还原为784维的数据
# 设置参数 input_size = 784 hidden_size = 32 output_size = 784 # 使用函数式方式创建Model input = tf.keras.layers.Input(shape = (input_size,)) # encode en = tf.keras.layers.Dense(hidden_size,activation=relu)(input) # decode de = tf.keras.layers.Dense(output_size,activation=sigmoid)(en) # 创建模型,指定输入与输出 model = tf.keras.Model(inputs = input,outputs = de)
使用print(model.summary())查看模型,可以看到每层的shape大小及参数的数量
4. 模型编译与训练
model.compile(optimizer=tf.optimizers.Adam(),loss =tf.losses.mse,metrics=[acc])
model.fit(x_train,x_train,
epochs=50,
batch_size=255,
shuffle=True,
validation_data=(x_test,x_test)) # 目标数据也是目标数据也是x_test
5. 从模型中获取编码器与解码器
获取编码器encode,输入为784维,输出为32维
input_en = tf.keras.layers.Input(shape = (input_size,)) en = model.layers[1](input_en) # 利用了模型中第一个Dense层训练好的参数 encode = tf.keras.Model(inputs = input_en ,outputs = en)
获取解码器decode,输入为32维,输出为784维
input_de = tf.keras.layers.Input(shape = (hidden_size,)) output_de = model.layers[-1](input_de) # -1 调用之前训练好模型的最后一层 ,利用了模型中第二个Dense层训练好的参数 decode = tf.keras.Model(inputs = input_de ,outputs = output_de)
6. 使用测试集进行测试
x_test = x_test.numpy()
print(x_test.shape)# (10000, 784)
encode_test = encode.predict(x_test)# 压缩 -> (10000, 32)
decode_test = decode.predict(encode_test) # 解码 -> (10000, 784)
# 绘图
n = 15
plt.figure(figsize=(20,4)) # 宽20 高4的 画布
for i in range(1,n):
ax = plt.subplot(2,n,i) # 几行几列的第几个
plt.imshow(x_test[i].reshape(28,28)) # 绘制上一层
ax = plt.subplot(2,n,i+n) # 几行几列的第几个
plt.imshow(decode_test[i].reshape(28,28)) # 绘制下一层
绘图结果,上面为原数据,下面为编码又解码之后的数据
