TensorFlow MNIST手写数据集的实现方法


Posted in Python onFebruary 05, 2020

MNIST数据集介绍

MNIST数据集中包含了各种各样的手写数字图片,数据集的官网是:http://yann.lecun.com/exdb/mnist/index.html,我们可以从这里下载数据集。使用如下的代码对数据集进行加载:

from tensorflow.examples.tutorials.mnist import input_data
mnist = input_data.read_data_sets('MNIST_data', one_hot=True)

运行上述代码会自动下载数据集并将文件解压在MNIST_data文件夹下面。代码中的one_hot=True,表示将样本的标签转化为one_hot编码。

MNIST数据集中的图片是28*28的,每张图被转化为一个行向量,长度是28*28=784,每一个值代表一个像素点。数据集中共有60000张手写数据图片,其中55000张训练数据,5000张测试数据。

在MNIST中,mnist.train.images是一个形状为[55000, 784]的张量,其中的第一个维度是用来索引图片,第二个维度图片中的像素。MNIST数据集包含有三部分,训练数据集,验证数据集,测试数据集(mnist.validation)。

标签是介于0-9之间的数字,用于描述图片中的数字,转化为one-hot向量即表示的数字对应的下标为1,其余的值为0。标签的训练数据是[55000,10]的数字矩阵。

下面定义了一个简单的网络对数据集进行训练,代码如下:

import tensorflow as tf
import numpy as np
from tensorflow.examples.tutorials.mnist import input_data
import matplotlib.pyplot as plt
mnist = input_data.read_data_sets('MNIST_data', one_hot=True)
tf.reset_default_graph()
x = tf.placeholder(tf.float32, [None, 784])
y = tf.placeholder(tf.float32, [None, 10])
w = tf.Variable(tf.random_normal([784, 10]))
b = tf.Variable(tf.zeros([10]))
pred = tf.matmul(x, w) + b
pred = tf.nn.softmax(pred)
cost = tf.reduce_mean(-tf.reduce_sum(y * tf.log(pred), reduction_indices=1))
learning_rate = 0.01
optimizer = tf.train.GradientDescentOptimizer(learning_rate).minimize(cost)
training_epochs = 25
batch_size = 100
display_step = 1
save_path = 'model/'
saver = tf.train.Saver()
with tf.Session() as sess:
  sess.run(tf.global_variables_initializer())
  for epoch in range(training_epochs):
    avg_cost = 0
    total_batch = int(mnist.train.num_examples/batch_size)
    for i in range(total_batch):
      batch_xs, batch_ys = mnist.train.next_batch(batch_size)
      _, c = sess.run([optimizer, cost], feed_dict={x:batch_xs, y:batch_ys})
      avg_cost += c / total_batch
    if (epoch + 1) % display_step == 0:
      print('epoch= ', epoch+1, ' cost= ', avg_cost)
  print('finished')
  correct_prediction = tf.equal(tf.argmax(pred, 1), tf.argmax(y, 1))
  accuracy = tf.reduce_mean(tf.cast(correct_prediction, tf.float32))
  print('accuracy: ', accuracy.eval({x:mnist.test.images, y:mnist.test.labels}))
  save = saver.save(sess, save_path=save_path+'mnist.cpkt')
print(" starting 2nd session ...... ")
with tf.Session() as sess:
  sess.run(tf.global_variables_initializer())
  saver.restore(sess, save_path=save_path+'mnist.cpkt')
  correct_prediction = tf.equal(tf.argmax(pred, 1), tf.argmax(y, 1))
  accuracy = tf.reduce_mean(tf.cast(correct_prediction, tf.float32))
  print('accuracy: ', accuracy.eval({x: mnist.test.images, y: mnist.test.labels}))
  output = tf.argmax(pred, 1)
  batch_xs, batch_ys = mnist.test.next_batch(2)
  outputval= sess.run([output], feed_dict={x:batch_xs, y:batch_ys})
  print(outputval)
  im = batch_xs[0]
  im = im.reshape(-1, 28)
  plt.imshow(im, cmap='gray')
  plt.show()
  im = batch_xs[1]
  im = im.reshape(-1, 28)
  plt.imshow(im, cmap='gray')
  plt.show()

总结

以上所述是小编给大家介绍的TensorFlow MNIST手写数据集的实现方法,希望对大家有所帮助!

Python 相关文章推荐
python实现的文件同步服务器实例
Jun 02 Python
浅谈python内置变量-reversed(seq)
Jun 21 Python
Python操作MongoDB数据库的方法示例
Jan 04 Python
TensorFlow 实战之实现卷积神经网络的实例讲解
Feb 26 Python
Python使用sorted对字典的key或value排序
Nov 15 Python
python实现ip地址查询经纬度定位详解
Aug 30 Python
使用python实现画AR模型时序图
Nov 20 Python
Python+OpenCV 实现图片无损旋转90°且无黑边
Dec 12 Python
python如何获取apk的packagename和activity
Jan 10 Python
Python SMTP发送电子邮件的示例
Sep 23 Python
python 中关于pycharm选择运行环境的问题
Oct 31 Python
python中pyqtgraph知识点总结
Jan 26 Python
tensorflow之并行读入数据详解
Feb 05 #Python
tensorflow mnist 数据加载实现并画图效果
Feb 05 #Python
tensorflow 自定义损失函数示例代码
Feb 05 #Python
利用Tensorflow的队列多线程读取数据方式
Feb 05 #Python
Tensorflow 多线程与多进程数据加载实例
Feb 05 #Python
TensorFlow自定义损失函数来预测商品销售量
Feb 05 #Python
解决Tensorflow 内存泄露问题
Feb 05 #Python
You might like
最贵的咖啡是怎么产生的,它的风味怎么样?
2021/03/04 新手入门
php正则取img标记中任意属性(正则替换去掉或改变图片img标记中的任意属性)
2013/08/13 PHP
typecho插件编写教程(五):核心代码
2015/05/28 PHP
Javascript 类、命名空间、代码组织代码
2011/07/31 Javascript
使用js正则控制input标签只允许输入的值
2013/07/29 Javascript
JS实现可展开折叠层的鼠标拖曳效果
2015/10/09 Javascript
JS实现的跨浏览器解析XML文件实例
2016/06/21 Javascript
JavaScript——DOM操作——Window.document对象详解
2016/07/14 Javascript
微信小程序 scroll-view组件实现列表页实例代码
2016/12/14 Javascript
js实现上下左右弹框划出效果
2017/03/08 Javascript
angular过滤器实现排序功能
2017/06/27 Javascript
二维码图片生成器QRCode.js简单介绍
2017/08/18 Javascript
浅析JavaScript中的特殊数据类型
2017/12/15 Javascript
koa2服务端使用jwt进行鉴权及路由权限分发的流程分析
2019/07/22 Javascript
微信小程序实现自定义底部导航
2020/11/18 Javascript
Python中用于计算对数的log()方法
2015/05/15 Python
用Python写飞机大战游戏之pygame入门(4):获取鼠标的位置及运动
2015/11/05 Python
Python读取指定目录下指定后缀文件并保存为docx
2017/04/23 Python
python笔记:mysql、redis操作方法
2017/06/28 Python
Python 串口读写的实现方法
2019/06/12 Python
Django在pycharm下修改默认启动端口的方法
2019/07/26 Python
关于pytorch多GPU训练实例与性能对比分析
2019/08/19 Python
python实现KNN近邻算法
2020/12/30 Python
在HTML5中如何使用CSS建立不可选的文字
2014/10/17 HTML / CSS
CSS Grid布局教程之网格单元格布局
2014/12/30 HTML / CSS
一波HTML5 Canvas基础绘图实例代码集合
2016/02/28 HTML / CSS
编写函数,将一个3*3矩阵转置
2013/10/09 面试题
中学教师自我鉴定
2014/02/07 职场文书
理工学院学生自我鉴定
2014/02/23 职场文书
市场总经理岗位职责
2014/04/11 职场文书
公关活动策划方案
2014/05/25 职场文书
销售类求职信
2014/06/13 职场文书
本溪水洞导游词
2015/02/11 职场文书
2019大学生预备党员转正思想汇报
2019/06/21 职场文书
Python基础知识之变量的详解
2021/04/14 Python
MySQL 四种连接和多表查询详解
2021/07/16 MySQL