【pytorch】学习笔记(四)-搭建神经网络进行关系拟合
【pytorch学习笔记】-搭建神经网络进行关系拟合
目标
1.创建一些围绕y=x^2+噪声这个函数的散点
2.用神经网络模型来建立一个可以代表他们关系的线条
建立数据集
import torch
from torch.autograd import Variable
import torch.nn.functional as F
import matplotlib.pyplot as plt
x=torch.unsqueeze(torch.linspace(-1,1,100),dim=1)#一维变二维,x从-1到1,切分为100份
y=x.pow(2)+0.2*torch.rand(x.size())#创建一些围绕着这y=x^2的随机点的散点
# plt.scatter(x.data.numpy(),y.data.numpy())#画图
# plt.show()
x,y=Variable(x),Variable(y)#构造神经网络要使用Variable类型
建立神经网络
1.继承torch.nn.Module模块
2.定义__init__函数,在初始化函数中定义输入层到隐藏层,从隐藏层再到输出层各个层的神经元个数
3.再一层层搭建(forward(x))层于层的关系链接
class Net(torch.nn.Module):
def __init__(self,n_feature,n_hidden,n_ouput):#初始化信息
super(Net, self).__init__()
self.hidden=torch.nn.Linear(n_feature,n_hidden,n_ouput)#隐藏层线性输出
self.predict=torch.nn.Linear(n_hidden,n_ouput)#输出层线性输出
def forward(self,x):#前向传递的过程
#正向传播输入值,神经网络输出预测值
x=F.relu(self.hidden(x))#激励函数加工一下
x=self.predict(x)#输出值预测值
return x
训练神经网络
1.定义训练工具optimizer,输入神经网络参数和学习效率
2.定义误差函数,使用均方差来计算实际值y和训练输出值之间的误差
3.每次训练向神经网络输入x,得到预测值,计算误差
4.注意要清空上一步的残余更新参数值
5.误差反向传播, 计算参数更新值
6.将参数更新值施加到 net 的 parameters 上
for t in range(200):#训练200次
prediction=net(x)#输入输入值
loss=loss_func(prediction,y)#计算误差预测值和真实值之间的误差,注意位置
optimizer.zero_grad()#梯度清零
loss.backward()#反向传递
optimizer.step()#优化梯度
可视化训练过程
for t in range(200):#训练200次
prediction=net(x)#输入输入值
loss=loss_func(prediction,y)#计算误差预测值和真实值之间的误差,注意位置
optimizer.zero_grad()#梯度清零
loss.backward()#反向传递
optimizer.step()#优化梯度
# 接着上面来
if t % 5 == 0:
# plot and show learning process
plt.cla()
plt.scatter(x.data.numpy(), y.data.numpy())
plt.plot(x.data.numpy(), prediction.data.numpy(), 'r-', lw=5)
plt.text(0.5, 0, 'Loss=%.4f' % loss.data.numpy(), fontdict={'size': 20, 'color': 'red'})
plt.pause(0.1)
完整代码
import torch
from torch.autograd import Variable
import torch.nn.functional as F
import matplotlib.pyplot as plt
x=torch.unsqueeze(torch.linspace(-1,1,100),dim=1)#一维变二维
y=x.pow(2)+0.2*torch.rand(x.size())
# plt.scatter(x.data.numpy(),y.data.numpy())
# plt.show()
x,y=Variable(x),Variable(y)#构造神经网络的是琥珀要使用Variable类型的
class Net(torch.nn.Module):
def __init__(self,n_feature,n_hidden,n_ouput):#初始化信息
super(Net, self).__init__()
self.hidden=torch.nn.Linear(n_feature,n_hidden,n_ouput)#隐藏层线性输出
self.predict=torch.nn.Linear(n_hidden,n_ouput)#输出层线性输出
def forward(self,x):#前向传递的过程
#正向传播输入值,神经网络输出预测值
x=F.relu(self.hidden(x))#激励函数加工一下
x=self.predict(x)#输出值预测值
return x
net=Net(n_feature=1,n_hidden=10,n_ouput=1)#输入值是一个,隐藏层有10个神经元,输出值为y值
print(net)
optimizer=torch.optim.SGD(net.parameters(),lr=0.5)#输入神经网络的所有参数,学习效率,这个是训练工具
loss_func=torch.nn.MSELoss()#误差处理均方差
plt.ion() # 画图
plt.show()
for t in range(200):#训练200次
prediction=net(x)#输入输入值
loss=loss_func(prediction,y)#计算误差预测值和真实值之间的误差,注意位置
optimizer.zero_grad()#梯度清零
loss.backward()#反向传递
optimizer.step()#优化梯度
# 接着上面来
if t % 5 == 0:
# plot and show learning process
plt.cla()
plt.scatter(x.data.numpy(), y.data.numpy())
plt.plot(x.data.numpy(), prediction.data.numpy(), 'r-', lw=5)
plt.text(0.5, 0, 'Loss=%.4f' % loss.data.numpy(), fontdict={'size': 20, 'color': 'red'})
plt.pause(0.1)
过程结果






中间过程省略一部分...

【pytorch】学习笔记(四)-搭建神经网络进行关系拟合的更多相关文章
- 莫烦PyTorch学习笔记(四)——回归
下面的代码说明个整个神经网络模拟回归的过程,代码含有详细注释,直接贴下来了 import torch from torch.autograd import Variable import torch. ...
- ensorflow学习笔记四:mnist实例--用简单的神经网络来训练和测试
http://www.cnblogs.com/denny402/p/5852983.html ensorflow学习笔记四:mnist实例--用简单的神经网络来训练和测试 刚开始学习tf时,我们从 ...
- Go语言学习笔记四: 运算符
Go语言学习笔记四: 运算符 这章知识好无聊呀,本来想跨过去,但没准有初学者要学,还是写写吧. 运算符种类 与你预期的一样,Go的特点就是啥都有,爱用哪个用哪个,所以市面上的运算符基本都有. 算术运算 ...
- kvm虚拟化学习笔记(四)之kvm虚拟机日常管理与配置
KVM虚拟化学习笔记系列文章列表----------------------------------------kvm虚拟化学习笔记(一)之kvm虚拟化环境安装http://koumm.blog.51 ...
- MySql学习笔记四
MySql学习笔记四 5.3.数据类型 数值型 整型 小数 定点数 浮点数 字符型 较短的文本:char, varchar 较长的文本:text, blob(较长的二进制数据) 日期型 原则:所选择类 ...
- 官网实例详解-目录和实例简介-keras学习笔记四
官网实例详解-目录和实例简介-keras学习笔记四 2018-06-11 10:36:18 wyx100 阅读数 4193更多 分类专栏: 人工智能 python 深度学习 keras 版权声明: ...
- ZooKeeper学习笔记四:使用ZooKeeper实现一个简单的分布式锁
作者:Grey 原文地址: ZooKeeper学习笔记四:使用ZooKeeper实现一个简单的分布式锁 前置知识 完成ZooKeeper集群搭建以及熟悉ZooKeeperAPI基本使用 需求 当多个进 ...
- C#可扩展编程之MEF学习笔记(四):见证奇迹的时刻
前面三篇讲了MEF的基础和基本到导入导出方法,下面就是见证MEF真正魅力所在的时刻.如果没有看过前面的文章,请到我的博客首页查看. 前面我们都是在一个项目中写了一个类来测试的,但实际开发中,我们往往要 ...
- IOS学习笔记(四)之UITextField和UITextView控件学习
IOS学习笔记(四)之UITextField和UITextView控件学习(博客地址:http://blog.csdn.net/developer_jiangqq) Author:hmjiangqq ...
随机推荐
- ISO15765
常用的缩略词 ISO15765网络层服务 协议功能 a)发送/接收最多4095个字节的数据信息: b)报告发送/接收完成状态. 网络层内部传输服务,CAN总线上的数据帧没帧只能传输8个字节,ISO 为 ...
- 背景(background)
背景(background) 背景家族由5个主要的背景属性组成 background-color背景颜色 background-color:colorNome(取值 如:颜色名 red green. ...
- JS基础_toString()
当我们直接在页面中打印一个对象时,实际上是输出的对象的toString()方法的返回值 如果我们希望在输出对象时不输出[ object Object ],可以为对象添加一个toString()方法或者 ...
- Druid连接池(无框架)
关于连接池有不少技术可以用,例如c3p0,druid等等,因为druid有监控平台,性能在同类产品中算top0的.所以我采用的事druid连接池. 首先熟悉一个技术,我们要搞明白,为什么要用他, 他能 ...
- golang mysql 如何设置最大连接数和最大空闲连接数
本文介绍golang 中连接MySQL时,如何设置最大连接数和最大空闲连接数. 关于最大连接数和最大空闲连接数,是定义在golang标准库中database/sql的. 文中例子连接MySQL用的SQ ...
- Cortex-M3 在C中上报入栈的寄存器和各fault状态寄存器
因为在标准C语音中是不能获取SP指针的.因而,如果想通过C代码来获取入栈的寄存器值,需要配合一小段汇编代码来获取当前的SP值,然后再把这个SP值以参数形式传送给C代码,最后以指针的形式把栈中的各寄存器 ...
- JAVA处理链表经典问题
定义链表节点Node class Node { private int Data;// 数据域 private Node Next;// 指针域 public Node(int Data) { // ...
- Spark3.0 preview预览版尝试GPU调用(本地模式不支持GPU)
Spark3.0 preview预览版可以下载使用,地址:https://archive.apache.org/dist/spark/spark-3.0.0-preview/,pom.xml也可以进行 ...
- delphi怎么一次性动态删除(释放)数个动态创建的组件?
比如procedure TForm1.Button1Click(Sender: TObject);vari:Integer;lbl: TLabel;beginfor i:=1 to 3 dobegin ...
- golang(08)接口介绍
原文链接 http://www.limerence2017.com/2019/09/12/golang13/#more 接口简介 golang 中接口是常用的数据结构,接口可以实现like的功能.什么 ...