tensorflow如何让损失函数只对一个变量求导(tensorflow:损失函数和优化器)

本文目录
tensorflow:损失函数和优化器
1. 损失函数是在graph中定义的经过operation的tensor。
2.损失函数最终要带入到优化器的minimize方法中做参数。minimize方法内部包含了compute_gradients和apply_gradients方法。
3. 优化器的minimize方法返回的是operation,一般命名为train_step。
4.session的run方法的参数,如果是operation,则返回值为None;如果是tensor,则返回值是ndarray。因此sess.run(train_step,feed_dict)无返回结果,只起训练作用。
5.Variable定义时必须给出初始值。Variable是变量,其值保存在session中,session的global_variable_initializer实际上是初始值的保存。
TensorFlow入门
1.可扩展参数: 从 kwargs 字典里获取,可限制key的取值
vars : name scope 下的变量(字典)
placeholders : 外部变量占位,一般是特征和标签(字典)
layers : 神经网络的 layer (列表)
activations : 每个 layer 的输出结果(列表)
inputs : 输入
output : 输出
loss : 损失
accuracy : 准确率
optimizer :优化器
opt_op :最优化的 op 操作
_build 是私有 build 方法,在继承Model的具体实现时对 layers 进行 append 操作,下面介绍 build 方法:
调用 _build ,所有变量设置共享空间( self.name )
构建模型序列:给输入,通过layer()返回输出,又将这个输出再次作为输入到下一个 layer() 中,循环这一过程;最终,取最后一层 layer 的结果作为 output
保存 name scope下 的变量到 self.vars
模型效果度量: _loss 方法, _accuracy 方法
常用的结构化数据文件格式有csv、txt 、libsvm,本篇文章主要说明结构化数据(csv/txt)如何在TF框架进行高效、灵活的读取,避免一些不合理的方式,迈出算法开发标准化、工程化的第一步。
最常见的读取数据的方式是利用pandas包将csv、txt读取为DataFrame,一次全部放入内存。这是一种非常低效的方式,应该尽量避免这种读取方法。其他第三方封装python读取方式的包(TFLearn等)也不建议使用,推荐使用TF框架的OP操作进行数据读取。
高效的 TensorFlow 读取方式是将数据读取转换成 OP,通过 session run 的方式拉去数据。读取线程源源不断地将文件系统中的文件读入到一个内存的队列中,而负责计算的是另一个线程,计算需要数据时,直接从内存队列中取就可以了,这样就可以解决GPU因为IO而空闲的问题。同时,不会一次性的preload到内存,再大的数据量也不会超出内存的限制。
梯度下降 :梯度下降是一个在机器学习中用于寻找较佳结果(曲线的最小值)的迭代优化算法。梯度的含义是斜率或者斜坡的倾斜度。下降的含义是代价函数的下降。
算法是迭代的,意思是需要多次使用算法获取结果,以得到最优化结果。梯度下降的迭代性质能使欠拟合演变成获得对数据的较佳拟合。
梯度下降中有一个称为学习率的参量。如上图左所示,刚开始学习率较大,因此下降步长更大。随着点的下降,学习率变得越来越小,从而下降步长也变小。同时,代价函数也在减小,或者说代价在减小,有时候也称为损失函数或者损失,两者是一样的。(损失/代价的减小是一个概念)。
只有在数据很庞大的时候(在机器学习中,数据一般情况下都会很大),我们才需要使用epochs,batch size,iteration这些术语,在这种情况下,一次性将数据输入计算机是不可能的。因此,为了解决这个问题,我们需要把数据分成小块,一块一块的传递给计算机,在每一步的末端更新神经网络的权重,拟合给定的数据。
(1)batchsize:批大小。在深度学习中,一般采用SGD训练,即每次训练在训练集中取batchsize个样本训练;
(2)iteration:1个iteration等于使用batchsize个样本训练一次;
(3)epoch:1个epoch等于使用训练集中的全部样本训练一次;
如何给某一层添加L2正则化
tf.nn.l2_loss()和tf.contrib.layers.l2_regularizer(),使用示例如下:
import tensorflow as tf
weights = tf.constant(, dtype=tf.float32)
sess = tf.InteractiveSession()
# 计算的是所有元素的平方和再除以2
print(tf.nn.l2_loss(weights).eval()())
# 等价于
print(tf.contrib.layers.l2_regularizer(1.)(weights).eval())
# output: 45.5
接下来将介绍两种方法将l2正则化项添加到损失函数后:
一、遍历trainable variables添加L2正则化项:
1.遍历可训练参数,将每个参数传入tf.nn.l2_loss()进行计算并相加起来;
2.乘以weight_decay并与base_loss相加。
weight_decay = 0.001
base_loss = tf.reduce_mean(tf.nn.sparse_softmax_cross_entropy_with_logits(labels=labels, logits=logits))
l2_loss = weight_decay * tf.add_n()
loss = base_loss + l2_loss
注意:该过程对每个trainable variable都进行了l2正则化,包括权值w和偏置b。有种说法是如果对偏执b进行l2正则化将会导致欠拟合,一般只需要对权值w进行正则化,所以来看第二种方法。
二、在构造网络层时传入l2正则化函数:
如下所示,在构造网络层时,将’kernel_initializer’参数设为l2正则化函数,则tensorflow会将该权重变量(卷积核)的l2正则化项加入到集合 tf.GraphKeys.REGULARIZATOIN_LOSSES里。
x = tf.layers.conv2d(x, 512, (3, 3),
padding=’same’,
activation=tf.nn.relu,
kernel_initializer=tf.truncated_normal_initializer(stddev=0.01),
kernel_regularizer=tf.contrib.layers.l2_regularizer(0.001)
在计算loss时使用tf.get_collection()来获取tf.GraphKeys.REGULARIZATOIN_LOSSES集合,然后相加即可:
base_loss = tf.reduce_mean(tf.nn.sparse_softmax_cross_entropy_with_logits(labels=labels, logits=logits))
l2_loss = tf.get_collection(tf.GraphKeys.REGULARIZATION_LOSSES)
loss = tf.add_n( + l2_loss, name="loss")
欢迎补充指正。

更多文章:
apache不能在本地计算机启动(关于“Windows不能在本地计算机启动Apache2.并参考特定服务错误代码1“问题解决)
2026年1月4日 19:00
deepin开机的四个选项(win10和deepin双系统怎样设置启动项)
2026年3月2日 15:45
deficiency词源(minus的详细意思minus的详细意思是什么)
2026年9月23日 15:30
mysql workbench怎么运行sql文件(如何使用MySQL Workbench导入.sql文件)
2025年5月23日 03:00
webpack缺点(如何理解webpack文档中对AMD缺点的描述)
2026年9月10日 12:30
discuz采集插件(火车头采集的数据怎么发布在discuz网站)
2025年11月7日 09:30
accommodation theory名词解释(带sion后缀的单词~~快快快!)
2026年7月27日 18:15
htmlcssjs软件下载手机(前端div+css怎么放到手机里面查看效果)
2026年2月9日 05:00
centos7怎么安装yum(CentOS7 配置 yum 源和 epel 源)
2025年11月6日 00:30
面向对象程序设计是java吗(Java是一种面向对象的编程语言吗)
2025年7月24日 13:45
linux配置与管理web服务器(高分请教如何搭建Linux下的web服务器)
2025年12月1日 02:00
vfp中的常用函数(求Visual Foxpro常用数值函数)
2026年4月7日 23:15
淘宝智能版导航代码(求淘宝店铺导航条代码,半透明状态的如下图)
2026年1月29日 13:00
shady是什么意思(Eminem为什么又叫slim shady,什么意思)
2025年9月24日 14:15












