1、生成高斯分布的随机数

导入numpy模块,通过numpy模块内的方法生成一组在方程

y = 2 * x + 3

周围小幅波动的随机坐标。代码如下:

 import numpy as np
import matplotlib.pyplot as plot def getRandomPoints(count):
xList = []
yList = []
for i in range(count):
x = np.random.normal(0, 0.5)
y = 2 * x + 3 + np.random.normal(0, 0.3)
xList.append(x)
yList.append(y)
return xList, yList if __name__ == '__main__':
X, Y = getRandomPoints(1000)
plot.scatter(X, Y)
plot.show()

运行上述代码,输出图形如下:

2、采用TensorFlow来获取上述方程的系数

  首先搭建基本的预估模型y = w * x + b,然后再采用梯度下降法进行训练,通过最小化损失函数的方法进行优化,最终训练得出方程的系数。

  在下面的例子中,梯度下降法的学习率为0.2,训练迭代次数为100次。

 def train(x, y):
# 生成随机系数
w = tf.Variable(tf.random_uniform([1], -1, 1))
# 生成随机截距
b = tf.Variable(tf.random_uniform([1], -1, 1))
# 预估值
preY = w * x + b # 损失值:预估值与实际值之间的均方差
loss = tf.reduce_mean(tf.square(preY - y))
# 优化器:梯度下降法,学习率为0.2
optimizer = tf.train.GradientDescentOptimizer(0.2)
# 训练:最小化损失函数
trainer = optimizer.minimize(loss) with tf.Session() as sess:
sess.run(tf.global_variables_initializer())
# 打印初始随机系数
print('init w:', sess.run(w), 'b:', sess.run(b))
# 先训练个100次:
for i in range(100):
sess.run(trainer)
# 每10次打印下系数
if i % 10 == 9:
print('w:', sess.run(w), 'b:', sess.run(b)) if __name__ == '__main__':
X, Y = getRandomPoints(1000)
train(X, Y)

  运行上面的代码,某次的最终结果为:

w = 1.9738449
b = 3.0027733

仅100次的训练迭代,得出的结果已十分接近方程的实际系数。

  某次模拟训练中的输出结果如下:

init w: [-0.6468966] b: [0.52244043]
w: [1.0336646] b: [2.9878206]
w: [1.636582] b: [3.0026987]
w: [1.8528996] b: [3.0027785]
w: [1.930511] b: [3.0027752]
w: [1.9583567] b: [3.0027738]
w: [1.9683474] b: [3.0027735]
w: [1.9719319] b: [3.0027733]
w: [1.9732181] b: [3.0027733]
w: [1.9736794] b: [3.0027733]
w: [1.9738449] b: [3.0027733]

3、完整代码和结果

完整测试代码:

 import numpy as np
import matplotlib.pyplot as plot
import tensorflow as tf def getRandomPoints(count, xscale=0.5, yscale=0.3):
xList = []
yList = []
for i in range(count):
x = np.random.normal(0, xscale)
y = 2 * x + 3 + np.random.normal(0, yscale)
xList.append(x)
yList.append(y)
return xList, yList def train(x, y, learnrate=0.2, cycle=100):
# 生成随机系数
w = tf.Variable(tf.random_uniform([1], -1, 1))
# 生成随机截距
b = tf.Variable(tf.random_uniform([1], -1, 1))
# 预估值
preY = w * x + b # 损失值:预估值与实际值之间的均方差
loss = tf.reduce_mean(tf.square(preY - y))
# 优化器:梯度下降法
optimizer = tf.train.GradientDescentOptimizer(learnrate)
# 训练:最小化损失函数
trainer = optimizer.minimize(loss) with tf.Session() as sess:
sess.run(tf.global_variables_initializer())
# 打印初始随机系数
print('init w:', sess.run(w), 'b:', sess.run(b))
for i in range(cycle):
sess.run(trainer)
# 每10次打印下系数
if i % 10 == 9:
print('w:', sess.run(w), 'b:', sess.run(b))
return sess.run(w), sess.run(b) if __name__ == '__main__':
X, Y = getRandomPoints(1000)
w, b = train(X, Y)
plot.scatter(X, Y)
plot.plot(X, w * X + b, c='r')
plot.show()

  最终效果图如下,蓝色为高斯随机分布数据,红色为最终得出的直线:

本文地址:https://www.cnblogs.com/laishenghao/p/9571343.html

TensorFlow 实现线性回归的更多相关文章

  1. tensorflow实现线性回归、以及模型保存与加载

    内容:包含tensorflow变量作用域.tensorboard收集.模型保存与加载.自定义命令行参数 1.知识点 """ 1.训练过程: 1.准备好特征和目标值 2.建 ...

  2. TensorFlow简单线性回归

    TensorFlow简单线性回归 将针对波士顿房价数据集的房间数量(RM)采用简单线性回归,目标是预测在最后一列(MEDV)给出的房价. 波士顿房价数据集可从http://lib.stat.cmu.e ...

  3. 深度学习入门实战(二)-用TensorFlow训练线性回归

    欢迎大家关注腾讯云技术社区-博客园官方主页,我们将持续在博客园为大家推荐技术精品文章哦~ 作者 :董超 上一篇文章我们介绍了 MxNet 的安装,但 MxNet 有个缺点,那就是文档不太全,用起来可能 ...

  4. 利用TensorFlow实现线性回归模型

    准备数据: import numpy as np import tensorflow as tf import matplotlib.pylot as plt # 随机生成1000个点,围绕在y=0. ...

  5. tensorflow实现线性回归总结

    1.知识点 """ 模拟一个y = 0.7x+0.8的案例 报警: 1.initialize_all_variables (from tensorflow.python. ...

  6. 如何用TensorFlow实现线性回归

    环境Anaconda 废话不多说,关键看代码 import tensorflow as tf import os os.environ['TF_CPP_MIN_LOG_LEVEL']='2' tf.a ...

  7. TensorFlow多元线性回归实现

    多元线性回归的具体实现 导入需要的所有软件包:   因为各特征的数据范围不同,需要归一化特征数据.为此定义一个归一化函数.另外,这里添加一个额外的固定输入值将权重和偏置结合起来.为此定义函数 appe ...

  8. TensorFlow实现线性回归模型代码

    模型构建 1.示例代码linear_regression_model.py #!/usr/bin/python # -*- coding: utf-8 -* import tensorflow as ...

  9. 学习TensorFlow,线性回归模型

    学习TensorFlow,在MNIST数据集上建立softmax回归模型并测试 一.代码 <span style="font-size:18px;">from tens ...

  10. tensorflow 学习1——tensorflow 做线性回归

    . 首先 Numpy: Numpy是Python的科学计算库,提供矩阵运算. 想想list已经提供了矩阵的形式,为啥要用Numpy,因为numpy提供了更多的函数. 使用numpy,首先要导入nump ...

随机推荐

  1. 结合 spring 使用阿里 Druid 连接池配置方法

    1.数据源 <!-- 配置数据源 --> <bean name="dataSource" class="com.alibaba.druid.pool.D ...

  2. jboss eap 6.2 ear包 下使用log4j日志

    被jboss7/eap的日志问题搞死了,查了好多资料,都是war包的,基本上使用jboss-deployment-structure.xml放到WEB-INF下,文件内容如下: 是我总是没法成功,最后 ...

  3. UNIX高级环境编程(15)进程和内存分配 < 故宫角楼 >

    故宫角楼是很多摄影爱好者常去的地方,夕阳余辉下的故宫角楼平静而安详.   首先,了解一下进程的基本概念,进程在内存中布局和内容. 此外,还需要知道运行时是如何为动态数据结构(如链表和二叉树)分配额外内 ...

  4. 【转】MaBatis学习---源码分析MyBatis缓存原理

    [原文]https://www.toutiao.com/i6594029178964673027/ 源码分析MyBatis缓存原理 1.简介 在 Web 应用中,缓存是必不可少的组件.通常我们都会用 ...

  5. MySQL基础之 AND和OR运算符

    AND和OR运算符 作用:用于基于一个以上的条件对记录进行过滤 用法:可在WHERE子句中把两个或多个条件结合在一起. AND:如果第一个条件和第二个条件都成立,才会显示一条记录 OR:如果第一个条件 ...

  6. Redis系列三:reids常用命令

    全局命令 keys *  查看所有键 dbsize 查看的是当前所在redis数据库的键总数 如果存在大量键,线上禁止使用此指令 exists key 检查键是否存在,存在返回1,不存在返回0 del ...

  7. 关于java中的使用通配符错误,错误信息Diamond types are not supported at language level '5‘

    当时,我问了下大神,他们问我是不是jdk问题.因为jdk8才支持这样的棱形写法.当时自己的jdk版本是jdk8,然后就奇怪了,最后我发现原来在Language level中调成了5.0 5.0不支持6 ...

  8. Kafka学习之路 (一)Kafka的简介

    一.简介 1.1 概述 Kafka是最初由Linkedin公司开发,是一个分布式.分区的.多副本的.多订阅者,基于zookeeper协调的分布式日志系统(也可以当做MQ系统),常见可以用于web/ng ...

  9. windows 下配置 Nginx 常见问题

    因为最近的项目需要用到负载均衡,不用考虑,当然用大名鼎鼎的Nginx啦.至于Nginx的介绍,这里就不多说了,直接进入主题如何在Windows下配置. 我的系统是win7旗舰版的,到官网下载最新版本 ...

  10. java代码,在linux上删除文件

    1.其实在linux上和window是一样的 2.path 传入的路径(直接从根目录到你的文件的位置) public static boolean delFile(String path) { log ...