tensorflow训练中出现nan问题的解决
作者:你不来我不老 发布时间:2023-02-10 09:34:09
标签:tensorflow,nan
深度学习中对于网络的训练是参数更新的过程,需要注意一种情况就是输入数据未做归一化时,如果前向传播结果已经是[0,0,0,1,0,0,0,0]这种形式,而真实结果是[1,0,0,0,0,0,0,0,0],此时由于得出的结论不惧有概率性,而是错误的估计值,此时反向传播会使得权重和偏置值变的无穷大,导致数据溢出,也就出现了nan的问题。
解决办法:
1、对输入数据进行归一化处理,如将输入的图片数据除以255将其转化成0-1之间的数据;
2、对于层数较多的情况,各层都做batch_nomorlization;
3、对设置Weights权重使用tf.truncated_normal(0, 0.01, [3,3,1,64])生成,同时值的均值为0,方差要小一些;
4、激活函数可以使用tanh;
5、减小学习率lr。
实例:
import tensorflow as tf
from tensorflow.examples.tutorials.mnist import input_data
mnist = input_data.read_data_sets('data',one_hot = True)
def add_layer(input_data,in_size, out_size,activation_function=None):
Weights = tf.Variable(tf.random_normal([in_size,out_size]))
Biases = tf.Variable(tf.zeros([1, out_size])+0.1)
Wx_plus_b = tf.add(tf.matmul(input_data, Weights), Biases)
if activation_function==None:
outputs = Wx_plus_b
else:
outputs = activation_function(Wx_plus_b)
#return outputs#, Weights
return {'outdata':outputs, 'w':Weights}
def get_accuracy(t_y):
# global l1
# accu = tf.reduce_mean(tf.cast(tf.equal(tf.argmax(l1['outdata'],1),tf.argmax(t_y,1)), dtype = tf.float32))
global prediction
accu = tf.reduce_mean(tf.cast(tf.equal(tf.argmax(prediction['outdata'],1),tf.argmax(t_y,1)), dtype = tf.float32))
return accu
X = tf.placeholder(tf.float32, [None, 784])
Y = tf.placeholder(tf.float32, [None, 10])
#l1 = add_layer(X, 784, 10, tf.nn.softmax)
#cross_entropy = tf.reduce_mean(-tf.reduce_sum(Y*tf.log(l1['outdata']), reduction_indices= [1]))
#l1 = add_layer(X, 784, 1024, tf.nn.relu)
l1 = add_layer(X, 784, 1024, None)
prediction = add_layer(l1['outdata'], 1024, 10, tf.nn.softmax)
cross_entropy = tf.reduce_mean(-tf.reduce_sum(Y*tf.log(prediction['outdata']), reduction_indices= [1]))
optimizer = tf.train.GradientDescentOptimizer(0.000001)
train = optimizer.minimize(cross_entropy)
newW = tf.Variable(tf.random_normal([1024,10]))
newOut = tf.matmul(l1['outdata'],newW)
newSoftMax = tf.nn.softmax(newOut)
init = tf.global_variables_initializer()
with tf.Session() as sess:
sess.run(init)
#print(sess.run(l1_Weights))
for i in range(2):
X_train, y_train = mnist.train.next_batch(1)
X_train = X_train/255 #需要进行归一化处理
#print(sess.run(l1['w'],feed_dict={X:X_train}))
#print(sess.run(prediction['w'],feed_dict={X:X_train, Y:y_train}))
#print(sess.run(l1['outdata'],feed_dict={X:X_train, Y:y_train}).shape)
print(sess.run(prediction['outdata'],feed_dict={X:X_train, Y:y_train}))
print(sess.run(newOut, feed_dict={X:X_train}))
print(sess.run(newSoftMax, feed_dict={X:X_train}))
print(y_train)
#print(sess.run(l1['outdata'], feed_dict={X:X_train}))
sess.run(train, feed_dict={X:X_train, Y:y_train})
if i%100 == 0:
#print(sess.run(cross_entropy, feed_dict={X:X_train, Y:y_train}))
accuracy = get_accuracy(mnist.test.labels)
print(sess.run(accuracy,feed_dict={X:mnist.test.images}))
#if i%100==0:
#print(sess.run(prediction, feed_dict={X:X_train}))
#print(sess.run(cross_entropy, feed_dict={X:X_train,Y:y_train}))
来源:http://blog.csdn.net/fireflychh/article/details/73691373
0
投稿
猜你喜欢
- Python报错:对象不存在此属性保错代码:我就搞不懂了,怎么会没有此属性② 原因:Python报错位置不对③总结下:在给对象属性赋值的时候
- 在运行复杂的Python程序时,执行时间会很长,这时也许想提高程序的执行效率。但该怎么做呢?首先,要有个工具能够检测代码中的瓶颈,例如,找到
- MySQL4.1以前版本服务器只能使用单一字符集,从MySQL4.1版本开始,不仅服务器能够使用多种字符集,而且在服务器、数据库、数据表、数
- response.getWriter().write() 功能:向前台页面显示一段信息。当在普通的url方式中,会生成一个新的页面来显示内容
- 声明,本文中所称CSS雪碧即为CSS Sprites,这个词组一直没有一个固定或者约定俗成的中文翻译,一些人开始称之为CSS雪碧,我们且当作
- WinHttp; // Microsoft WinHTTP Services, version 5.1Alias HTTPREQUEST_P
- (一)原理 小偷程序实际上是通过了XML中的XMLHTTP组件调用其它网站上的网页。比如新闻小偷程序,
- 函数javascript函数相信大家都写过不少了,所以我们这里只是简单介绍一下.创建函数:function f(x) {........}v
- 1、互动流通的活跃度是社区网站的关键,产品设计者大都需要在此猛下药。facebook有利用率最高的minifeed,myspace有“好友的
- 严格来说,Having并不需要一个子表,但没有子表的Having并没有实际意义。如果你只需要一个表,那么你可以用Where子句达到一切目的。
- 本文实例讲述了JavaScript点击按钮后弹出透明浮动层的方法。分享给大家供大家参考。具体分析如下:这里实现点击后页面变灰色,并用JS弹出
- 最近关于HTML5吵得火热,很多人认为HTML5出现会秒杀Flash,以至于在各大web前端开 * 坛吵得不可开交。论坛里三言两语说的不够 尽
- 验证关键词是否为sql保留字的在线工具:<html> <head><t
- 创建列表sample_list = ['a',1,('a','b')]Python 列表操作
- Python pywifi ERROR Open handle failed这个问题的网上的资料很少,可能是因为简单吧。这里记录下解决办法。
- 如下所示:import serialimport sysimport osimport timeimport redef wait_for_
- 1. 单行导入与多行导入在 Go 语言中,一个包可包含多个 .go 文件(这些文件必须得在同一级文件夹中),只要这些 .go 文件的头部都使
- 假设要生成一千万个随机数,常规的做法如下:var numbers = [];for (var&nbs
- 有些框架本身就支持多配置文件,例如Ruby On Rails,nodejs下的expressjs。python下的Flask虽然本身支持配置
- 定义计算N的阶乘的函数1)使用循环计算阶乘def frac(n): r = 1 if n<=1: