tensorflow 自定义损失函数示例代码


Posted in Python onFebruary 05, 2020

这个自定义损失函数的背景:(一般回归用的损失函数是MSE, 但要看实际遇到的情况而有所改变)

我们现在想要做一个回归,来预估某个商品的销量,现在我们知道,一件商品的成本是1元,售价是10元。

如果我们用均方差来算的话,如果预估多一个,则损失一块钱,预估少一个,则损失9元钱(少赚的)。

显然,我宁愿预估多了,也不想预估少了。

所以,我们就自己定义一个损失函数,用来分段地看,当yhat 比 y大时怎么样,当yhat比y小时怎么样。

(yhat沿用吴恩达课堂中的叫法)

import tensorflow as tf
from numpy.random import RandomState
batch_size = 8
# 两个输入节点
x = tf.placeholder(tf.float32, shape=(None, 2), name="x-input")
# 回归问题一般只有一个输出节点
y_ = tf.placeholder(tf.float32, shape=(None, 1), name="y-input")
# 定义了一个单层的神经网络前向传播的过程,这里就是简单加权和
w1 = tf.Variable(tf.random_normal([2, 1], stddev=1, seed=1))
y = tf.matmul(x, w1)
# 定义预测多了和预测少了的成本
loss_less = 10
loss_more = 1
#在windows下,下面用这个where替代,因为调用tf.select会报错
loss = tf.reduce_sum(tf.where(tf.greater(y, y_), (y - y_)*loss_more, (y_-y)*loss_less))
train_step = tf.train.AdamOptimizer(0.001).minimize(loss)
#通过随机数生成一个模拟数据集
rdm = RandomState(1)
dataset_size = 128
X = rdm.rand(dataset_size, 2)
"""
设置回归的正确值为两个输入的和加上一个随机量,之所以要加上一个随机量是
为了加入不可预测的噪音,否则不同损失函数的意义就不大了,因为不同损失函数
都会在能完全预测正确的时候最低。一般来说,噪音为一个均值为0的小量,所以
这里的噪音设置为-0.05, 0.05的随机数。
"""
Y = [[x1 + x2 + rdm.rand()/10.0-0.05] for (x1, x2) in X]
with tf.Session() as sess:
 init = tf.global_variables_initializer()
 sess.run(init)
 steps = 5000
 for i in range(steps):
  start = (i * batch_size) % dataset_size
  end = min(start + batch_size, dataset_size)
  sess.run(train_step, feed_dict={x:X[start:end], y_:Y[start:end]})
 print(sess.run(w1))

[[ 1.01934695]
[ 1.04280889]

最终结果如上面所示。

因为我们当初生成训练数据的时候,y是x1 + x2,所以回归结果应该是1,1才对。
但是,由于我们加了自己定义的损失函数,所以,倾向于预估多一点。

如果,我们将loss_less和loss_more对调,我们看一下结果:

[[ 0.95525807]
[ 0.9813394 ]]

通过这个例子,我们可以看出,对于相同的神经网络,不同的损失函数会对训练出来的模型产生重要的影响。

引用:以上实例为《Tensorflow实战 Google深度学习框架》中提供。

总结

以上所述是小编给大家介绍的tensorflow 自定义损失函数示例,希望对大家有所帮助!

Python 相关文章推荐
Python中os.path用法分析
Jan 15 Python
使用Python对Csv文件操作实例代码
May 12 Python
Windows下Anaconda的安装和简单使用方法
Jan 04 Python
python-opencv颜色提取分割方法
Dec 08 Python
python实现多线程端口扫描
Aug 31 Python
Python序列化pickle模块使用详解
Mar 05 Python
python名片管理系统开发
Jun 18 Python
Python自动登录QQ的实现示例
Aug 28 Python
Python urllib3软件包的使用说明
Nov 18 Python
pytorch显存一直变大的解决方案
Apr 08 Python
使用pandas生成/读取csv文件的方法实例
Jul 09 Python
pycharm无法安装cv2模块问题
May 20 Python
利用Tensorflow的队列多线程读取数据方式
Feb 05 #Python
Tensorflow 多线程与多进程数据加载实例
Feb 05 #Python
TensorFlow自定义损失函数来预测商品销售量
Feb 05 #Python
解决Tensorflow 内存泄露问题
Feb 05 #Python
TensorFlow实现指数衰减学习率的方法
Feb 05 #Python
关于Tensorflow使用CPU报错的解决方式
Feb 05 #Python
解决Tensorflow sess.run导致的内存溢出问题
Feb 05 #Python
You might like
PHP 程序授权验证开发思路
2009/07/09 PHP
php中session与cookie的比较
2015/01/27 PHP
php求今天、昨天、明天时间戳的简单实现方法
2016/07/28 PHP
PHP互换两个变量值的方法(不用第三变量)
2016/11/14 PHP
IE之动态添加DOM节点触发window.resize事件
2010/07/27 Javascript
jQuery修改class属性和CSS样式整理
2015/01/30 Javascript
angularjs学习笔记之双向数据绑定
2015/09/26 Javascript
jQuery实现的个性化返回底部与返回顶部特效代码
2015/10/30 Javascript
Jquery EasyUI实现treegrid上显示checkbox并取选定值的方法
2016/04/29 Javascript
JavaScript操作表单实例讲解(上)
2016/06/20 Javascript
Angular.js与node.js项目里用cookie校验账户登录详解
2017/02/22 Javascript
浅析node.js的模块加载机制
2018/05/25 Javascript
jQuery阻止事件冒泡实例分析
2018/07/03 jQuery
json数据传到前台并解析展示成列表的方法
2018/08/06 Javascript
JS二级菜单不同实现方法分析【4种方法】
2018/12/21 Javascript
微信小程序使用for循环动态渲染页面操作示例
2018/12/25 Javascript
Vuex的各个模块封装的实现
2020/06/05 Javascript
对Python的多进程锁的使用方法详解
2019/02/18 Python
python 画二维、三维点之间的线段实现方法
2019/07/07 Python
python通过移动端访问查看电脑界面
2020/01/06 Python
Tensorflow实现部分参数梯度更新操作
2020/01/23 Python
python GUI库图形界面开发之PyQt5菜单栏控件QMenuBar的详细使用方法与实例
2020/02/28 Python
HTML5 WebGL 实现民航客机飞行监控系统
2019/07/25 HTML / CSS
美国一家主打母婴用品的团购网站:zulily
2017/09/19 全球购物
中专生毕业自我鉴定
2013/11/01 职场文书
21岁生日感言
2014/02/27 职场文书
我的职业生涯规划:打造自己的运动帝国
2014/09/18 职场文书
房屋租赁合同补充协议
2014/10/11 职场文书
学校政风行风评议工作总结
2014/10/21 职场文书
和谐家庭事迹材料
2014/12/20 职场文书
学习保证书怎么写
2015/02/26 职场文书
婚礼父母致辞
2015/07/28 职场文书
校运会广播稿
2015/08/19 职场文书
党员干部学习十八届五中全会精神心得体会
2016/01/05 职场文书
python如何利用cv2模块读取显示保存图片
2021/06/04 Python
Python实现视频自动打码的示例代码
2022/04/08 Python