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中用于求最小值的min()方法
May 15 Python
Python函数的周期性执行实现方法
Aug 13 Python
python中子类继承父类的__init__方法实例
Dec 15 Python
python安装教程 Pycharm安装详细教程
May 02 Python
Python线程创建和终止实例代码
Jan 20 Python
Python实现通讯录功能
Feb 22 Python
Django实现简单网页弹出警告代码
Nov 15 Python
谈谈Python:为什么类中的私有属性可以在外部赋值并访问
Mar 05 Python
python3+opencv 使用灰度直方图来判断图片的亮暗操作
Jun 02 Python
详解用Python爬虫获取百度企业信用中企业基本信息
Jul 02 Python
python的链表基础知识点
Sep 13 Python
解决pycharm 格式报错tabs和space不一致问题
Feb 26 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
Linux下安装oracle客户端并配置php5.3
2014/10/12 PHP
完美解决phpexcel导出到xls文件出现乱码的问题
2016/10/29 PHP
php实现的支付宝网页支付功能示例【基于TP5框架】
2019/09/16 PHP
jquery.validate使用攻略 第五步 正则验证
2010/07/01 Javascript
Web Inspector:关于在 Sublime Text 中调试Js的介绍
2013/04/18 Javascript
鼠标经过tr时,改变tr当前背景颜色
2014/01/13 Javascript
jQuery实现当按下回车键时绑定点击事件
2014/01/28 Javascript
JS实现带有抽屉效果的产品类网站多级导航菜单代码
2015/09/15 Javascript
javascript实现数组去重的多种方法
2016/03/14 Javascript
JavaScript动态生成二维码图片
2016/04/20 Javascript
jQuery自制提示框tooltip改进版
2016/08/01 Javascript
JS代码实现百度地图 画圆 删除标注
2016/10/12 Javascript
vue实现可增删查改的成绩单
2016/10/27 Javascript
vue使用watch 观察路由变化,重新获取内容
2017/03/08 Javascript
微信小程序 设置启动页面的两种方法
2017/03/09 Javascript
node.js中debug模块的简单介绍与使用
2017/04/25 Javascript
详解最新vue-cli 2.9.1的webpack存在问题
2017/12/16 Javascript
微信小程序实现给嵌套template模板传递数据的方式总结
2017/12/18 Javascript
js操作table中tr的顺序实现上移下移一行的效果
2018/11/22 Javascript
微信小程序的注册页面包含倒计时验证码、获取用户信息
2019/05/22 Javascript
30分钟用Node.js构建一个API服务器的步骤详解
2019/05/24 Javascript
vue实现的请求服务器端API接口示例
2019/05/25 Javascript
详解Python字符串切片
2019/05/20 Python
python 写一个性能测试工具(一)
2020/10/24 Python
python中复数的共轭复数知识点总结
2020/12/06 Python
德国高性价比网上药店:medpex
2017/07/09 全球购物
运行时异常与一般异常有何异同?
2014/01/05 面试题
大学生学习生活的自我评价
2013/11/01 职场文书
学习经验交流会主持词
2014/04/01 职场文书
人力资源职位说明书
2014/07/29 职场文书
2014副镇长民主生活会个人对照检查材料思想汇报
2014/09/30 职场文书
新郎结婚保证书
2015/02/26 职场文书
保护地球的宣传语
2015/07/13 职场文书
家长会感言
2015/08/01 职场文书
python 如何将两个实数矩阵合并为一个复数矩阵
2021/05/19 Python
python cv2图像质量压缩的算法示例
2021/06/04 Python