tensorflow 自定义损失函数示例代码
作者:陈阔 时间:2023-03-13 21:37:18
这个自定义损失函数的背景:(一般回归用的损失函数是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 自定义损失函数示例,希望对大家有所帮助!
来源:https://www.cnblogs.com/chenkuo/p/8087055.html
标签:tensorflow,自定义,损失,函数
![](/images/zang.png)
![](/images/jiucuo.png)
猜你喜欢
python队列queue模块详解
2023-03-28 17:26:02
Go语言文件开关及读写操作示例
2023-08-05 19:47:27
python自动登录12306并自动点击验证码完成登录的实现源代码
2021-07-08 12:50:29
SQL文本字段的数字排序问题
2008-11-18 16:47:00
![](https://img.aspxhome.com/file/UploadPic/200811/18/o200852912144-57s.jpg)
树型结构在ASP中的简单解决
2007-10-07 12:52:00
Oracle 数据库操作类
2009-08-12 12:06:00
MYSQL教程:表达式操作符和数据类型转换
2009-02-27 15:51:00
Sql Server、Oracle以及Access数据库 判断字段是否为空的办法 (From calmzeal's code life)
2011-02-24 19:44:00
Python实现对二维码数据进行压缩
2022-10-22 12:51:59
![](https://img.aspxhome.com/file/2023/2/92172_0s.png)
Python参数类型以及常见的坑详解
2023-04-16 13:52:33
![](https://img.aspxhome.com/file/2023/0/97020_0s.png)
Python爬虫基础之requestes模块
2022-04-24 20:20:15
![](https://img.aspxhome.com/file/2023/3/83943_0s.png)
Google的设计导引
2008-04-06 14:18:00
960px下的黄金分割
2009-07-31 14:33:00
![](https://img.aspxhome.com/file/UploadPic/20097/31/498022_1248973274hhmt-19s.jpg)
python用matplotlib可视化绘图详解
2022-04-20 02:38:42
![](https://img.aspxhome.com/file/2023/7/72067_0s.png)
Linux服务器网卡流量查看方法 shell和Python各一枚
2023-10-07 09:31:19
![](https://img.aspxhome.com/file/2023/9/66259_0s.gif)
简介Python中用于处理字符串的center()方法
2021-04-15 20:47:17
JS Object.preventExtensions(),Object.seal()与Object.freeze()用法实例分析
2023-09-04 22:44:33
微信小程序创建自定义全局函数以及其调用方法详解
2023-08-24 20:43:22
![](https://img.aspxhome.com/file/2023/5/56255_0s.png)
wxPython中wx.gird.Gird添加按钮的实现
2021-07-05 15:13:19
![](https://img.aspxhome.com/file/2023/9/86269_0s.png)
分享2个方便调试Python代码的实用工具
2021-08-26 18:50:08
![](https://img.aspxhome.com/file/2023/6/65856_0s.png)