Tensorflow2.0笔记 - AutoEncoder做FashionMnist数据集训练

本笔记记录自编码器做FashionMnist数据集训练,关于autoencoder的原理,请自行百度。

复制代码
import os
import time
import tensorflow as tf
from tensorflow import keras
from tensorflow.keras import datasets, layers, optimizers, Sequential, metrics, Input,losses
from PIL import Image
from matplotlib import pyplot as plt
import numpy as np
from tensorflow.keras.models import Model

os.environ['TF_CPP_MIN_LOG_LEVEL']='2'
#tf.random.set_seed(12345)
tf.__version__

#加载fashion mnist数据集
(x_train, _), (x_test, _) = datasets.fashion_mnist.load_data()
#图片像素数据范围限值到[0,1]
x_train = x_train.astype('float32') / 255.
x_test = x_test.astype('float32') / 255.

print (x_train.shape)
print (x_test.shape)

h_dim = 64 
class Autoencoder(Model):
  def __init__(self, h_dim):
    super(Autoencoder, self).__init__()
    self.h_dim = h_dim   
    #encoder层,[b, 28, 28] => [b, 784] => [b, h_dim]
    self.encoder = tf.keras.Sequential([
      layers.Flatten(),
      layers.Dense(256, activation='relu'),
      layers.Dense(h_dim, activation='relu'),
    ])
    #decoder层,[b, h_dim] => [b,784] => [b, 28, 28]
    self.decoder = tf.keras.Sequential([
      layers.Dense(784, activation='sigmoid'),
      #恢复成28x28的图片
      layers.Reshape((28, 28))
    ])

  def call(self, x):
    encoded = self.encoder(x)
    decoded = self.decoder(encoded)
    return decoded

model = Autoencoder(h_dim)

model.compile(optimizer='adam', loss=losses.MeanSquaredError())
model.fit(x_train, x_train,
                epochs=10,
                shuffle=True,
                validation_data=(x_test, x_test))


encoded_imgs = model.encoder(x_test).numpy()
decoded_imgs = model.decoder(encoded_imgs).numpy()
n = 10
plt.figure(figsize=(20, 4))
for i in range(n):
  #绘制原始图像
  ax = plt.subplot(2, n, i + 1)
  plt.imshow(x_test[i])
  plt.title("original")
  plt.gray()
  ax.get_xaxis().set_visible(False)
  ax.get_yaxis().set_visible(False)

  #绘制重建的图像
  ax = plt.subplot(2, n, i + 1 + n)
  plt.imshow(decoded_imgs[i])
  plt.title("reconstructed")
  plt.gray()
  ax.get_xaxis().set_visible(False)
  ax.get_yaxis().set_visible(False)
plt.show()

运行结果:

相关推荐
数智化码农几秒前
半导体制造人力资源数字化选型:适配双用工体系,兼顾合规与人效
大数据·运维·python·制造
东离与糖宝4 分钟前
SSE流式输出详解:大模型打字机效果底层原理
人工智能
李兆龙的博客5 分钟前
从一到无穷大 #91:从 Habitat 看存储平台的整合与分工
数据库·人工智能·架构
statistican_ABin7 分钟前
中国婚姻与生育观念变迁分析报告
python
阿文和她的Key11 分钟前
OpenAI 关 Pro 入口事件复盘:企业 AI 架构的稳定性问题,不只是故障应急
人工智能·架构
sdzhyt14 分钟前
从“台账分散”到“智能调度”,AI如何走进应急避难场所管理一线?
人工智能
. . . . .15 分钟前
comfyUI原理
人工智能·算法·机器学习
清水白石00816 分钟前
别被 `Process.start()` 骗了:Python multiprocessing 从面试题到线上生产的进程启动机制实战
开发语言·python
在所不辞兄17 分钟前
【零基础学智能仿真-16】循环神经网络与LSTM/GRU——学习力学响应的历史记忆
人工智能·rnn·神经网络·gru·lstm·工程技术·工程仿真
愚公搬代码21 分钟前
【愚公系列】《造浪者:AI创业实战地图》002-AI创业的六个本质差异
人工智能