深度学习神经网络 MNIST手写数据辨识 1 前向传播和反向传播

本文主要是介绍深度学习神经网络 MNIST手写数据辨识 1 前向传播和反向传播,希望对大家解决编程问题提供一定的参考价值,需要的开发者们随着小编来一起学习吧!

首先是前向传播的程序。为了更清晰我们分段讲解。

第一部分导入模块,并设置输入节点为28*28,输出节点为10(0到9共10个数字),第一层的节点为500(随便设的)

import tensorflow as tf
INPUT_NODE = 784
OUTPUT_NODE = 10
LAYER1_NODE = 500

然后是生成单个层次网络的结构,判断损失函数是否加入正则

#定义神经网络的输入,参数和输出,定义前向传播过程
def get_weight(shape,regularizer):w = tf.Variable(tf.random_normal(shape,stddev=0.1),dtype=tf.float32) #生成随机参数if regularizer != None:tf.add_to_collection('losses',tf.contrib.layers.l2_regularizer(regularizer)(w))return w

同时设置偏置项,偏置项不需要正则化。

def get_bias(shape):b = tf.Variable(tf.constant(0.01,shape=shape))return b

在总的前向传播网络中设置两层网络:

def forward(x,regularizer):w1 = get_weight([INPUT_NODE,LAYER1_NODE],regularizer)b1 = get_bias([LAYER1_NODE])y1 = tf.nn.relu(tf.matmul(x,w1)+b1)w2 = get_weight([LAYER1_NODE, OUTPUT_NODE], regularizer)b2 = get_bias([OUTPUT_NODE])y = tf.matmul(y1, w2) + b2return y

然后反向传播。这里实现了一种机制:每次训练前,先查看一下已有的模型,

首先仍然是加载模型和设置初始常量:正则系数为0.0001,不算很大。然后滑动平均值衰减设为0.99.

import tensorflow as tf
from tensorflow.examples.tutorials.mnist import input_data
import mnist_forward2
import osBATCH_SIZE = 200
LEARNING_RATE_BASE = 0.1
LEARNING_RATE_DECAY = 0.99
REGULARIZER = 0.0001STEPS = 50000MOVING_AVERAGE_DECAY = 0.99MODEL_SAVE_PATH="./model/" #模型保存路径
MODEL_NAME="mnist_model" #模型保存文件名

然后是反向传播函数  def backward(mnist) :

输入数据和输出占位就先不说了,这里提一下损失函数:

采用最后输出为softmax的网络激活函数,并把损失函数定义为交叉熵

    #定义损失函数ce = tf.nn.sparse_softmax_cross_entropy_with_logits(logits=y,labels=tf.argmax(y_,1))cem = tf.reduce_mean(ce)loss = cem + tf.add_n(tf.get_collection('losses'))

学习率的设置方法和以前一样,然后定义反向传播方法,并设置和启用滑动平均值。

之后我们使用保存模型的函数:

    saver = tf.train.Saver()

在会话中我们先查看模型目录下有没有训练好的模型和参数,如果有,就恢复:

    with tf.Session() as sess:ckpt = tf.train.get_checkpoint_state(MODEL_SAVE_PATH)if ckpt and ckpt.model_checkpoint_path:  # 先判断是否有模型saver.restore(sess, ckpt.model_checkpoint_path)  # 恢复模型到当前会话#可以观察到当前的会话已经包含当前的正确globalstep了currentstep = ckpt.model_checkpoint_path.split('/')[-1].split('-')[-1]print(currentstep)

值得注意的是,我们之前在当前的模型里使用了滑动平均值,这里恢复的时候恢复了滑动平均后的数据,然后继续根据global_step来计算新的滑动平均值。而且,因为在模型中我们嵌入了global_step,所以恢复的时候,global_step也被恢复了。

然后开始训练。

        for i in range(STEPS):xs,ys = mnist.train.next_batch(BATCH_SIZE)_,loss_value,step = sess.run([train_op,loss,global_step],feed_dict={x:xs,y_:ys})if i % 1000 == 0:print("After " + str(i) + " steps, loss is: " + str(loss_value))saver.save(sess,os.path.join(MODEL_SAVE_PATH,MODEL_NAME),global_step=global_step)

设置自动执行的函数main() :

def main():mnist = input_data.read_data_sets("./data/",one_hot=True)backward(mnist)if __name__ == '__main__':main()

现在前向传播和后向传播都已经设置好了。大家多运行几次,就会发现每次都是从上一次训练好的模型中开始然后继续训练的。

这篇关于深度学习神经网络 MNIST手写数据辨识 1 前向传播和反向传播的文章就介绍到这儿,希望我们推荐的文章对编程师们有所帮助!



http://www.chinasem.cn/article/883656

相关文章

Python获取中国节假日数据记录入JSON文件

《Python获取中国节假日数据记录入JSON文件》项目系统内置的日历应用为了提升用户体验,特别设置了在调休日期显示“休”的UI图标功能,那么问题是这些调休数据从哪里来呢?我尝试一种更为智能的方法:P... 目录节假日数据获取存入jsON文件节假日数据读取封装完整代码项目系统内置的日历应用为了提升用户体验,

SpringCloud动态配置注解@RefreshScope与@Component的深度解析

《SpringCloud动态配置注解@RefreshScope与@Component的深度解析》在现代微服务架构中,动态配置管理是一个关键需求,本文将为大家介绍SpringCloud中相关的注解@Re... 目录引言1. @RefreshScope 的作用与原理1.1 什么是 @RefreshScope1.

Java利用JSONPath操作JSON数据的技术指南

《Java利用JSONPath操作JSON数据的技术指南》JSONPath是一种强大的工具,用于查询和操作JSON数据,类似于SQL的语法,它为处理复杂的JSON数据结构提供了简单且高效... 目录1、简述2、什么是 jsONPath?3、Java 示例3.1 基本查询3.2 过滤查询3.3 递归搜索3.4

MySQL大表数据的分区与分库分表的实现

《MySQL大表数据的分区与分库分表的实现》数据库的分区和分库分表是两种常用的技术方案,本文主要介绍了MySQL大表数据的分区与分库分表的实现,文中通过示例代码介绍的非常详细,对大家的学习或者工作具有... 目录1. mysql大表数据的分区1.1 什么是分区?1.2 分区的类型1.3 分区的优点1.4 分

Mysql删除几亿条数据表中的部分数据的方法实现

《Mysql删除几亿条数据表中的部分数据的方法实现》在MySQL中删除一个大表中的数据时,需要特别注意操作的性能和对系统的影响,本文主要介绍了Mysql删除几亿条数据表中的部分数据的方法实现,具有一定... 目录1、需求2、方案1. 使用 DELETE 语句分批删除2. 使用 INPLACE ALTER T

Python 中的异步与同步深度解析(实践记录)

《Python中的异步与同步深度解析(实践记录)》在Python编程世界里,异步和同步的概念是理解程序执行流程和性能优化的关键,这篇文章将带你深入了解它们的差异,以及阻塞和非阻塞的特性,同时通过实际... 目录python中的异步与同步:深度解析与实践异步与同步的定义异步同步阻塞与非阻塞的概念阻塞非阻塞同步

Python Dash框架在数据可视化仪表板中的应用与实践记录

《PythonDash框架在数据可视化仪表板中的应用与实践记录》Python的PlotlyDash库提供了一种简便且强大的方式来构建和展示互动式数据仪表板,本篇文章将深入探讨如何使用Dash设计一... 目录python Dash框架在数据可视化仪表板中的应用与实践1. 什么是Plotly Dash?1.1

Redis 中的热点键和数据倾斜示例详解

《Redis中的热点键和数据倾斜示例详解》热点键是指在Redis中被频繁访问的特定键,这些键由于其高访问频率,可能导致Redis服务器的性能问题,尤其是在高并发场景下,本文给大家介绍Redis中的热... 目录Redis 中的热点键和数据倾斜热点键(Hot Key)定义特点应对策略示例数据倾斜(Data S

Python实现将MySQL中所有表的数据都导出为CSV文件并压缩

《Python实现将MySQL中所有表的数据都导出为CSV文件并压缩》这篇文章主要为大家详细介绍了如何使用Python将MySQL数据库中所有表的数据都导出为CSV文件到一个目录,并压缩为zip文件到... python将mysql数据库中所有表的数据都导出为CSV文件到一个目录,并压缩为zip文件到另一个

使用PyTorch实现手写数字识别功能

《使用PyTorch实现手写数字识别功能》在人工智能的世界里,计算机视觉是最具魅力的领域之一,通过PyTorch这一强大的深度学习框架,我们将在经典的MNIST数据集上,见证一个神经网络从零开始学会识... 目录当计算机学会“看”数字搭建开发环境MNIST数据集解析1. 认识手写数字数据库2. 数据预处理的