在上一篇《TensorFlow入门之MNIST样例代码分析》中,我们讲解了如果来用一个三层全连接网络实现手写数字识别。但是在实际运用中我们需要更有效率,更加灵活的代码。在TensorFlow实战这本书中给出了更好的实现,他将程序分为三个模块,分别是前向传播过程模块,训练模块和验证检测模块。并且在这个版本中添加了模型持久化功能,我们可以将模型保存下来,方便之后的模型检验,并且我们可以一边训练新的模型,一边来检验模型,代码更加的灵活高效。

前向传播模块

首先将前向传播过程抽象出来,作为一个可以作为训练测试共享的模块,取名为mnist_inference.py,将这个过程抽象出来的好处是,一是可以保证在训练或者测试的过程中前向传播的一致性,提高代码的复用性。还有一点是我们可以更好地将其与滑动平均模型与模型持久化功能结合,更加灵活的来检验新的模型。mnist_inference.py代码如下:

  1. # -*- coding: utf-8 -*-
  2. import tensorflow as tf
  3. # 定义神经网络结构相关的参数
  4. INPUT_NODE = 784
  5. OUTPUT_NODE = 10
  6. LAYER1_NODE = 500
  7. # 通过tf.get_variable函数来获取变量。在训练神经网络时会创建这些变量;在测试时会通
  8. # 过保存的模型加载这些变量的取值。而且更加方便的是,因为可以在变量加载时将滑动平均变
  9. # 量重命名,所以可以直接通过相同的名字在训练时使用变量自身,而在测试时使用变量的滑动
  10. # 平均值。在这个函数中也会将变量的正则化损失加入到损失集合。
  11. def get_weight_variable(shape, regularizer):
  12. weights = tf.get_variable(
  13. "weights", shape,
  14. initializer=tf.truncated_normal_initializer(stddev=0.1)
  15. )
  16. # 当给出了正则化生成函数时,将当前变量的正则化损失加入名字为losses的集合。在这里
  17. # 使用了add_to_collection函数将一个张量加入一个集合,而这个集合的名称为losses。
  18. # 这是自定义的集合,不在TensorFlow自动管理的集合列表中。
  19. if regularizer != None:
  20. tf.add_to_collection('losses', regularizer(weights))
  21. return weights
  22. # 定义神经网络的前向传播过程
  23. def inference(input_tensor, regularizer):
  24. # 声明第一层神经网络的变量并完成前向传播过程。
  25. with tf.variable_scope('layer1'):
  26. # 这里通过tf.get_variable或者tf.Variable没有本质区别,因为在训练或者测试
  27. # 中没有在同一个程序中多次调用这个函数。如果在同一个程序中多次调用,在第一次
  28. # 调用之后需要将reuse参数设置为True。
  29. weights = get_weight_variable(
  30. [INPUT_NODE, LAYER1_NODE], regularizer
  31. )
  32. biases = tf.get_variable(
  33. "biases", [LAYER1_NODE],
  34. initializer=tf.constant_initializer(0.0)
  35. )
  36. layer1 = tf.nn.relu(tf.matmul(input_tensor, weights)+biases)
  37. # 类似的声明第二层神经网络的变量并完成前向传播过程。
  38. with tf.variable_scope('layer2'):
  39. weights = get_weight_variable(
  40. [LAYER1_NODE, OUTPUT_NODE], regularizer
  41. )
  42. biases = tf.get_variable(
  43. "biases", [OUTPUT_NODE],
  44. initializer=tf.constant_initializer(0.0)
  45. )
  46. layer2 = tf.matmul(layer1, weights) + biases
  47. # 返回最后前向传播的结果
  48. return layer2

训练模块

将训练模型的模块提取出来,训练模块命名为mnist_train.py,在下面的代码中每过1000个step我们就保存一次模型。代码如下:

  1. # -*- coding: utf-8 -*-
  2. import os
  3. import tensorflow as tf
  4. from tensorflow.examples.tutorials.mnist import input_data
  5. # 加载mnist_inference.py中定义的常量和前向传播的函数。
  6. import mnist_inference
  7. # 配置神经网络的参数。
  8. BATCH_SIZE = 100
  9. LEARNING_RATE_BASE = 0.8
  10. LEARNING_RATE_DECAY = 0.99
  11. REGULARIZATION_RATE = 0.0001
  12. TRAINING_STEPS = 30000
  13. MOVING_AVERAGE_DECAY = 0.99
  14. # 模型保存的路径和文件名
  15. MODEL_SAVE_PATH = "./model/"
  16. MODEL_NAME = "model.ckpt"
  17. def train(mnist):
  18. # 定义输入输出placeholder。
  19. x = tf.placeholder(tf.float32, [None, mnist_inference.INPUT_NODE], name='x-input')
  20. y_ = tf.placeholder(tf.float32, [None, mnist_inference.OUTPUT_NODE], name='y-input')
  21. regularizer = tf.contrib.layers.l2_regularizer(REGULARIZATION_RATE)
  22. # 直接使用mnist_inference.py中定义的前向传播过程
  23. y = mnist_inference.inference(x, regularizer)
  24. global_step = tf.Variable(0, trainable=False)
  25. # 定义损失函数、学习率、滑动平均操作以及训练过程
  26. variable_averages = tf.train.ExponentialMovingAverage(
  27. MOVING_AVERAGE_DECAY, global_step
  28. )
  29. variable_averages_op = variable_averages.apply(
  30. tf.trainable_variables()
  31. )
  32. cross_entropy = tf.nn.sparse_softmax_cross_entropy_with_logits(
  33. logits=y, labels=tf.argmax(y_, 1)
  34. )
  35. cross_entropy_mean = tf.reduce_mean(cross_entropy)
  36. loss = cross_entropy_mean + tf.add_n(tf.get_collection('losses'))
  37. learning_rate = tf.train.exponential_decay(
  38. LEARNING_RATE_BASE,
  39. global_step,
  40. mnist.train.num_examples / BATCH_SIZE,
  41. LEARNING_RATE_DECAY
  42. )
  43. train_step = tf.train.GradientDescentOptimizer(learning_rate)\
  44. .minimize(loss, global_step=global_step)
  45. with tf.control_dependencies([train_step, variable_averages_op]):
  46. train_op = tf.no_op(name='train')
  47. # 初始化TensorFlow持久化类
  48. saver = tf.train.Saver()
  49. with tf.Session() as sess:
  50. tf.global_variables_initializer().run()
  51. # 在训练过程中不再测试模型在验证数据上的表现,验证和测试的过程将会有一个独
  52. # 立的程序来完成。
  53. for i in range(TRAINING_STEPS):
  54. xs, ys = mnist.train.next_batch(BATCH_SIZE)
  55. _, loss_value, step = sess.run([train_op, loss, global_step],
  56. feed_dict={x: xs, y_: ys})
  57. # 每1000轮保存一次模型
  58. if i % 1000 == 0:
  59. # 输出当前的训练情况。这里只输出了模型在当前训练batch上的损失
  60. # 函数大小。通过损失函数的大小可以大概了解训练的情况。在验证数
  61. # 据集上正确率的信息会有一个单独的程序来生成
  62. print("After %d training step(s), loss on training "
  63. "batch is %g." % (step, loss_value))
  64. # 保存当前的模型。注意这里给出了global_step参数,这样可以让每个
  65. # 被保存的模型的文件名末尾加上训练的轮数,比如“model.ckpt-1000”,
  66. # 表示训练1000轮之后得到的模型。
  67. saver.save(
  68. sess, os.path.join(MODEL_SAVE_PATH, MODEL_NAME),
  69. global_step=global_step
  70. )
  71. def main(argv=None):
  72. mnist = input_data.read_data_sets("./data", one_hot=True)
  73. train(mnist)
  74. if __name__ == "__main__":
  75. tf.app.run()

验证与测试模块

验证模块与测试模块可以对保存好的训练模型进行验证与测试,在下面的代码中我们选择每过10秒钟验证一个最新的模型。这样做的好处是可以将训练与验证或者测试分割开来,同时进行。该模块命名为mnist_eval.py

  1. # -*- coding: utf-8 -*-
  2. import time
  3. import tensorflow as tf
  4. from tensorflow.examples.tutorials.mnist import input_data
  5. # 加载mnist_inference.py 和mnist_train.py中定义的常量和函数。
  6. import mnist_inference
  7. import mnist_train
  8. # 每10秒加载一次最新的模型,并且在测试数据上测试最新模型的正确率
  9. EVAL_INTERVAL_SECS = 10
  10. def evaluate(mnist):
  11. with tf.Graph().as_default() as g:
  12. # 定义输入输出的格式。
  13. x = tf.placeholder(
  14. tf.float32, [None, mnist_inference.INPUT_NODE], name='x-input'
  15. )
  16. y_ = tf.placeholder(
  17. tf.float32, [None, mnist_inference.OUTPUT_NODE], name='y-input'
  18. )
  19. validate_feed = {x: mnist.validation.images,
  20. y_: mnist.validation.labels}
  21. # 直接通过调用封装好的函数来计算前向传播的结果。因为测试时不关注ze正则化损失的值
  22. # 所以这里用于计算正则化损失的函数被设置为None。
  23. y = mnist_inference.inference(x, None)
  24. # 使用前向传播的结果计算正确率。如果需要对未知的样例进行分类,那么使用
  25. # tf.argmax(y,1)就可以得到输入样例的预测类别了。
  26. correct_prediction = tf.equal(tf.argmax(y, 1), tf.argmax(y_, 1))
  27. accuracy = tf.reduce_mean(tf.cast(correct_prediction, tf.float32))
  28. # 通过变量重命名的方式来加载模型,这样在前向传播的过程中就不需要调用求滑动平均
  29. # 的函数来获取平均值了。这样就可以完全共用mnist_inference.py中定义的
  30. # 前向传播过程。
  31. variable_averages = tf.train.ExponentialMovingAverage(
  32. mnist_train.MOVING_AVERAGE_DECAY
  33. )
  34. variables_to_restore = variable_averages.variables_to_restore()
  35. saver = tf.train.Saver(variables_to_restore)
  36. # 每隔EVAL_INTERVAL_SECS秒调用一次计算正确率的过程以检验训练过程中正确率的
  37. # 变化。
  38. while True:
  39. with tf.Session() as sess:
  40. # tf.train.get_checkpoint_state函数会通过checkpoint文件自动
  41. # 找到目录中最新模型的文件名。
  42. ckpt = tf.train.get_checkpoint_state(
  43. mnist_train.MODEL_SAVE_PATH
  44. )
  45. if ckpt and ckpt.model_checkpoint_path:
  46. # 加载模型。
  47. saver.restore(sess, ckpt.model_checkpoint_path)
  48. # 通过文件名得到模型保存时迭代的轮数。
  49. global_step = ckpt.model_checkpoint_path\
  50. .split('/')[-1].split('-')[-1]
  51. accuracy_score = sess.run(accuracy,
  52. feed_dict=validate_feed)
  53. print("After %s training step(s), validation "
  54. "accuracy = %g" % (global_step, accuracy_score))
  55. else:
  56. print("No checkpoint file found")
  57. return
  58. time.sleep(EVAL_INTERVAL_SECS)
  59. def main(argv=None):
  60. mnist = input_data.read_data_sets("./data", one_hot=True)
  61. evaluate(mnist)
  62. if __name__ == "__main__":
  63. tf.app.run()

总结

这个样例是一个非常好的可以用来理解TensorFlow的程序,特别是TensorFlow的计算图的理解,还有模型持久化与恢复,变量的管理,滑动平均模型的实现等等。还有这种灵活的模块分块的思想也值得学习。

TensorFlow入门之MNIST最佳实践-深度学习的更多相关文章

  1. TensorFlow入门之MNIST最佳实践

    在上一篇<TensorFlow入门之MNIST样例代码分析>中,我们讲解了如果来用一个三层全连接网络实现手写数字识别.但是在实际运用中我们需要更有效率,更加灵活的代码.在TensorFlo ...

  2. RESTful接口设计原则/最佳实践(学习笔记)

    RESTful接口设计原则/最佳实践(学习笔记) 原文地址:http://www.vinaysahni.com/best-practices-for-a-pragmatic-restful-api 1 ...

  3. 分享《机器学习实战基于Scikit-Learn和TensorFlow》中英文PDF源代码+《深度学习之TensorFlow入门原理与进阶实战》PDF+源代码

    下载:https://pan.baidu.com/s/1qKaDd9PSUUGbBQNB3tkDzw <机器学习实战:基于Scikit-Learn和TensorFlow>高清中文版PDF+ ...

  4. 吴裕雄 python 神经网络——TensorFlow训练神经网络:MNIST最佳实践

    import os import tensorflow as tf from tensorflow.examples.tutorials.mnist import input_data INPUT_N ...

  5. 分别基于TensorFlow、PyTorch、Keras的深度学习动手练习项目

    ×下面资源个人全都跑了一遍,不会出现仅是字符而无法运行的状况,运行环境: Geoffrey Hinton在多次访谈中讲到深度学习研究人员不要仅仅只停留在理论上,要多编程.个人在学习中也体会到单单的看理 ...

  6. TensorFlow入门之MNIST样例代码分析

    这几天想系统的学习一下TensorFlow,为之后的工作打下一些基础.看了下<TensorFlow:实战Google深度学习框架>这本书,目前个人觉得这本书还是对初学者挺友好的,作者站在初 ...

  7. 【一统江湖的大前端(9)】TensorFlow.js 开箱即用的深度学习工具

    示例代码托管在:http://www.github.com/dashnowords/blogs 博客园地址:<大史住在大前端>原创博文目录 目录 一. 上手TensorFlow.js 二. ...

  8. NLP入门(五)用深度学习实现命名实体识别(NER)

    前言   在文章:NLP入门(四)命名实体识别(NER)中,笔者介绍了两个实现命名实体识别的工具--NLTK和Stanford NLP.在本文中,我们将会学习到如何使用深度学习工具来自己一步步地实现N ...

  9. 从TensorFlow到PyTorch:九大深度学习框架哪款最适合你?

    开源的深度学习神经网络正步入成熟,而现在有许多框架具备为个性化方案提供先进的机器学习和人工智能的能力.那么如何决定哪个开源框架最适合你呢?本文试图通过对比深度学习各大框架的优缺点,从而为各位读者提供一 ...

随机推荐

  1. Notes 20180307 : 运算符

    我们前边曾说过程序=数据结构+算法,数据结构讲的是数据在内存中的存储形式,这个我会作为2018的一个重点来研究,不过在这里不做赘述,前半年的工作以JavaSE为主.算法则是我们在数据结构的基础上对其的 ...

  2. HTML表格属性及简单实例

    这里主要总结记录下表格的一些属性和简单的样式,方便以后不时之需. 1.<table> 用来定义HTML的表格,具有本地属性 border 表示边框,border属性的值必须为1或空字符串( ...

  3. 关于端口冲突的解决方式Error: listen EACCES 0.0.0.80

    笔者昨天下午临走前安装了vs 2017想要运行一下项目的NET后端来让本机的前端直接对接后端,但是没注意到运行vs后IIS直接占用了本机的80端口.第二天跑nodeJS的时候直接Error: list ...

  4. postman中 form-data、x-www-form-urlencoded、raw、binary的区别【转】

    链接:https://blog.csdn.net/wangjun5159/article/details/47781443 1.form-data: 就是http请求中的multipart/form- ...

  5. Pycharm常用的快捷键

    常用快捷键: Ctrl + D              复制选定的区域或行 Ctrl + Y               删除选定的行 Ctrl + Alt + L         代码格式化 Ct ...

  6. MongoDB 4.0.6 Manual

    General mongod options: -v [ --verbose ] [=arg(=v)] be more verbose (include multiple times for more ...

  7. 基于OMAPL138的字符驱动_GPIO驱动AD9833(三)之中断申请IRQ

    基于OMAPL138的字符驱动_GPIO驱动AD9833(三)之中断申请IRQ 0. 导语 学习进入到了下一个阶段,还是以AD9833为例,这次学习是向设备申请中断,实现触发,在未来很多场景,比如做用 ...

  8. Python学习手册之__main__ 模块,常用第三方模块和打包发布

    在上一篇文章中,我们介绍了 Python 的 元组拆包.三元运算符和对 Python 的 else 语句进行了深入讲解,现在我们介绍 Python 的 __main__ 模块.常用第三方模块和打包发布 ...

  9. Java 8 中有趣的操作 Stream

    Stream 不是java io中的stream 对象创建 我们没有必要使用一个迭代来创建对象,直接使用流就可以 String[] strs = {"haha","hoh ...

  10. HashMap底层实现原理及扩容机制

    HashMap的数据结构:数组+链表+红黑树:Java7中的HashMap只由数组+链表构成:Java8引入了红黑树,提高了HashMap的性能:借鉴一张图来说明,原文:https://www.jia ...