基于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)) # 绘制下一层

绘图结果,上面为原数据,下面为编码又解码之后的数据

经验分享 程序员 微信小程序 职场和发展