Tensorflow2实现三层神经网络的前向传输

2024-04-27 12:32

本文主要是介绍Tensorflow2实现三层神经网络的前向传输,希望对大家解决编程问题提供一定的参考价值,需要的开发者们随着小编来一起学习吧!

使用Tensorflow2自己实现三层神经网络的前向传输

导入所需要的包

import tensorflow as tf

导入数据集,本次采用的tensorflow提供的经典是mnist手写数据集

# x:[60k,28,28],
# y:[60k]
(x, y),_ = tf.keras.datasets.mnist.load_data()
# x:[0-255]->[0,1]    y:[0-9]
# 将x,y 转换为Tensor,并且将x归一化
x = tf.convert_to_tensor(x, dtype=tf.float32)/255.
y = tf.convert_to_tensor(y, dtype=tf.int32)
print(x.shape,y.shape,x.dtype,y.dtype) 

输出x,y的shape为下图,x表示60000张28*28 的图片,y对应60000个标签,范围为【0-9】输出x,y的形状
设置batch为128,即一次训练128条数据。

# 设置batch为128
train_db = tf.data.Dataset.from_tensor_slices((x,y)).batch(128)
train_iter = iter(train_db)
sample = next(train_iter)
# 一个batch的形状
print('batch:',sample[0].shape,sample[1].shape)

以下为训练所需参数和过程,本次设计为三层神经网络。输入层为28*28的图片,节点为784,第二层为256个节点,第三层为128个节点,输出层为10个节点,注释中,b为训练数据的个数(维数)。

# 创建权值
# 降维过程 [b,784]->[b,256]->[b,128]->[b,10]
# [dim_in, dim_out],[dim_out]
# 随机生成一个权重矩阵,并且初始化每一层的偏置
# 由于下文中的梯度下降法,tape默认只会跟踪tf.Variable类型的信息,所以进行转换。
w1 = tf.Variable(tf.random.truncated_normal([784,256],stddev=0.1))
b1 =  tf.Variable(tf.zeros([256]))
w2 =  tf.Variable(tf.random.truncated_normal([256,128],stddev=0.1))
b2 =  tf.Variable(tf.zeros([128]))
w3 =  tf.Variable(tf.random.truncated_normal([128,10],stddev=0.1))
b3 =  tf.Variable(tf.zeros([10]))
lr = 1e-3  #0.001   10的-3次方

训练过程如下代码,设置epoch为10:

for epoch in range(10):# enumerate处理后可以返回当前步骤的step,便于打印当前信息print('epoch',epoch)for step,(x,y) in enumerate(train_db):#x :[128,28,28]#y :[128]x = tf.reshape(x,[-1,28*28])with tf.GradientTape() as tape:  # x :[128,28*28]# h1 = x@w1+b1# [b,784]@[784*256]+[256]->[b,256]+[256]->[b,256]+[b,256]h1 = x@w1 +tf.broadcast_to(b1,[x.shape[0],256])h1 = tf.nn.relu(h1)h2 =  h1@w2 + b2h2 =  tf.nn.relu(h2)out =  h2@w3 + b3# compute loss 计算误差# out:[b,10]y_onehot = tf.one_hot(y,depth=10)# mse = mean(sum(y-out)^2)loss = tf.square(y_onehot-out)# mean: scalarloss = tf.reduce_mean(loss)# compute gradientsgrads = tape.gradient(loss,[w1,b1,w2,b2,w3,b3])# w1 = w1 - lr * w1_gradw1.assign_sub(lr * grads[0])  # 保持w1原地更新,保持引用不变,类型不变b1.assign_sub(lr * grads[1])w2.assign_sub(lr * grads[2])b2.assign_sub(lr * grads[3])w3.assign_sub(lr * grads[4])b3.assign_sub(lr * grads[5])if step % 100 == 0:print(step,'  loss:',float(loss))

运行结果如下图:
0-5
6-10

这篇关于Tensorflow2实现三层神经网络的前向传输的文章就介绍到这儿,希望我们推荐的文章对编程师们有所帮助!



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

相关文章

Linux下删除乱码文件和目录的实现方式

《Linux下删除乱码文件和目录的实现方式》:本文主要介绍Linux下删除乱码文件和目录的实现方式,具有很好的参考价值,希望对大家有所帮助,如有错误或未考虑完全的地方,望不吝赐教... 目录linux下删除乱码文件和目录方法1方法2总结Linux下删除乱码文件和目录方法1使用ls -i命令找到文件或目录

SpringBoot+EasyExcel实现自定义复杂样式导入导出

《SpringBoot+EasyExcel实现自定义复杂样式导入导出》这篇文章主要为大家详细介绍了SpringBoot如何结果EasyExcel实现自定义复杂样式导入导出功能,文中的示例代码讲解详细,... 目录安装处理自定义导出复杂场景1、列不固定,动态列2、动态下拉3、自定义锁定行/列,添加密码4、合并

mybatis执行insert返回id实现详解

《mybatis执行insert返回id实现详解》MyBatis插入操作默认返回受影响行数,需通过useGeneratedKeys+keyProperty或selectKey获取主键ID,确保主键为自... 目录 两种方式获取自增 ID:1. ​​useGeneratedKeys+keyProperty(推

Spring Boot集成Druid实现数据源管理与监控的详细步骤

《SpringBoot集成Druid实现数据源管理与监控的详细步骤》本文介绍如何在SpringBoot项目中集成Druid数据库连接池,包括环境搭建、Maven依赖配置、SpringBoot配置文件... 目录1. 引言1.1 环境准备1.2 Druid介绍2. 配置Druid连接池3. 查看Druid监控

Linux在线解压jar包的实现方式

《Linux在线解压jar包的实现方式》:本文主要介绍Linux在线解压jar包的实现方式,具有很好的参考价值,希望对大家有所帮助,如有错误或未考虑完全的地方,望不吝赐教... 目录linux在线解压jar包解压 jar包的步骤总结Linux在线解压jar包在 Centos 中解压 jar 包可以使用 u

c++ 类成员变量默认初始值的实现

《c++类成员变量默认初始值的实现》本文主要介绍了c++类成员变量默认初始值,文中通过示例代码介绍的非常详细,对大家的学习或者工作具有一定的参考学习价值,需要的朋友们下面随着小编来一起学习学习吧... 目录C++类成员变量初始化c++类的变量的初始化在C++中,如果使用类成员变量时未给定其初始值,那么它将被

Qt使用QSqlDatabase连接MySQL实现增删改查功能

《Qt使用QSqlDatabase连接MySQL实现增删改查功能》这篇文章主要为大家详细介绍了Qt如何使用QSqlDatabase连接MySQL实现增删改查功能,文中的示例代码讲解详细,感兴趣的小伙伴... 目录一、创建数据表二、连接mysql数据库三、封装成一个完整的轻量级 ORM 风格类3.1 表结构

基于Python实现一个图片拆分工具

《基于Python实现一个图片拆分工具》这篇文章主要为大家详细介绍了如何基于Python实现一个图片拆分工具,可以根据需要的行数和列数进行拆分,感兴趣的小伙伴可以跟随小编一起学习一下... 简单介绍先自己选择输入的图片,默认是输出到项目文件夹中,可以自己选择其他的文件夹,选择需要拆分的行数和列数,可以通过

Python中将嵌套列表扁平化的多种实现方法

《Python中将嵌套列表扁平化的多种实现方法》在Python编程中,我们常常会遇到需要将嵌套列表(即列表中包含列表)转换为一个一维的扁平列表的需求,本文将给大家介绍了多种实现这一目标的方法,需要的朋... 目录python中将嵌套列表扁平化的方法技术背景实现步骤1. 使用嵌套列表推导式2. 使用itert

Python使用pip工具实现包自动更新的多种方法

《Python使用pip工具实现包自动更新的多种方法》本文深入探讨了使用Python的pip工具实现包自动更新的各种方法和技术,我们将从基础概念开始,逐步介绍手动更新方法、自动化脚本编写、结合CI/C... 目录1. 背景介绍1.1 目的和范围1.2 预期读者1.3 文档结构概述1.4 术语表1.4.1 核