深度学习神经网络 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

相关文章

SpringBoot分段处理List集合多线程批量插入数据方式

《SpringBoot分段处理List集合多线程批量插入数据方式》文章介绍如何处理大数据量List批量插入数据库的优化方案:通过拆分List并分配独立线程处理,结合Spring线程池与异步方法提升效率... 目录项目场景解决方案1.实体类2.Mapper3.spring容器注入线程池bejsan对象4.创建

PHP轻松处理千万行数据的方法详解

《PHP轻松处理千万行数据的方法详解》说到处理大数据集,PHP通常不是第一个想到的语言,但如果你曾经需要处理数百万行数据而不让服务器崩溃或内存耗尽,你就会知道PHP用对了工具有多强大,下面小编就... 目录问题的本质php 中的数据流处理:为什么必不可少生成器:内存高效的迭代方式流量控制:避免系统过载一次性

C#实现千万数据秒级导入的代码

《C#实现千万数据秒级导入的代码》在实际开发中excel导入很常见,现代社会中很容易遇到大数据处理业务,所以本文我就给大家分享一下千万数据秒级导入怎么实现,文中有详细的代码示例供大家参考,需要的朋友可... 目录前言一、数据存储二、处理逻辑优化前代码处理逻辑优化后的代码总结前言在实际开发中excel导入很

深度解析Python中递归下降解析器的原理与实现

《深度解析Python中递归下降解析器的原理与实现》在编译器设计、配置文件处理和数据转换领域,递归下降解析器是最常用且最直观的解析技术,本文将详细介绍递归下降解析器的原理与实现,感兴趣的小伙伴可以跟随... 目录引言:解析器的核心价值一、递归下降解析器基础1.1 核心概念解析1.2 基本架构二、简单算术表达

MyBatis-plus处理存储json数据过程

《MyBatis-plus处理存储json数据过程》文章介绍MyBatis-Plus3.4.21处理对象与集合的差异:对象可用内置Handler配合autoResultMap,集合需自定义处理器继承F... 目录1、如果是对象2、如果需要转换的是List集合总结对象和集合分两种情况处理,目前我用的MP的版本

深度解析Java @Serial 注解及常见错误案例

《深度解析Java@Serial注解及常见错误案例》Java14引入@Serial注解,用于编译时校验序列化成员,替代传统方式解决运行时错误,适用于Serializable类的方法/字段,需注意签... 目录Java @Serial 注解深度解析1. 注解本质2. 核心作用(1) 主要用途(2) 适用位置3

Java MCP 的鉴权深度解析

《JavaMCP的鉴权深度解析》文章介绍JavaMCP鉴权的实现方式,指出客户端可通过queryString、header或env传递鉴权信息,服务器端支持工具单独鉴权、过滤器集中鉴权及启动时鉴权... 目录一、MCP Client 侧(负责传递,比较简单)(1)常见的 mcpServers json 配置

GSON框架下将百度天气JSON数据转JavaBean

《GSON框架下将百度天气JSON数据转JavaBean》这篇文章主要为大家详细介绍了如何在GSON框架下实现将百度天气JSON数据转JavaBean,文中的示例代码讲解详细,感兴趣的小伙伴可以了解下... 目录前言一、百度天气jsON1、请求参数2、返回参数3、属性映射二、GSON属性映射实战1、类对象映

Maven中生命周期深度解析与实战指南

《Maven中生命周期深度解析与实战指南》这篇文章主要为大家详细介绍了Maven生命周期实战指南,包含核心概念、阶段详解、SpringBoot特化场景及企业级实践建议,希望对大家有一定的帮助... 目录一、Maven 生命周期哲学二、default生命周期核心阶段详解(高频使用)三、clean生命周期核心阶

C# LiteDB处理时间序列数据的高性能解决方案

《C#LiteDB处理时间序列数据的高性能解决方案》LiteDB作为.NET生态下的轻量级嵌入式NoSQL数据库,一直是时间序列处理的优选方案,本文将为大家大家简单介绍一下LiteDB处理时间序列数... 目录为什么选择LiteDB处理时间序列数据第一章:LiteDB时间序列数据模型设计1.1 核心设计原则