python如何用pytorch实现线性回归?
Admin 2021-05-19 群英技术资讯 1591 次浏览
pytorch是一个python优先的深度学习框架,更够在强大的GPU加速基础上实现张量和动态神经网络。这篇文章主要带大家了解使用pytorch实现线性回归的步骤,感兴趣的朋友可以参考学习。
线性回归都是包括以下几个步骤:定义模型、选择损失函数、选择优化函数、 训练数据、测试
import torch
import matplotlib.pyplot as plt
# 构建数据集
x_data= torch.Tensor([[1.0],[2.0],[3.0],[4.0],[5.0],[6.0]])
y_data= torch.Tensor([[2.0],[4.0],[6.0],[8.0],[10.0],[12.0]])
#定义模型
class LinearModel(torch.nn.Module):
def __init__(self):
super(LinearModel, self).__init__()
self.linear= torch.nn.Linear(1,1) #表示输入输出都只有一层,相当于前向传播中的函数模型,因为我们一般都不知道函数是什么形式的
def forward(self, x):
y_pred= self.linear(x)
return y_pred
model= LinearModel()
# 使用均方误差作为损失函数
criterion= torch.nn.MSELoss(size_average= False)
#使用梯度下降作为优化SGD
# 从下面几种优化器的生成结果图像可以看出,SGD和ASGD效果最好,因为他们的图像收敛速度最快
optimizer= torch.optim.SGD(model.parameters(),lr=0.01)
# ASGD
# optimizer= torch.optim.ASGD(model.parameters(),lr=0.01)
# optimizer= torch.optim.Adagrad(model.parameters(), lr= 0.01)
# optimizer= torch.optim.RMSprop(model.parameters(), lr= 0.01)
# optimizer= torch.optim.Adamax(model.parameters(),lr= 0.01)
# 训练
epoch_list=[]
loss_list=[]
for epoch in range(100):
y_pred= model(x_data)
loss= criterion(y_pred, y_data)
epoch_list.append(epoch)
loss_list.append(loss.item())
print(epoch, loss.item())
optimizer.zero_grad() #梯度归零
loss.backward() #反向传播
optimizer.step() #更新参数
print("w= ", model.linear.weight.item())
print("b= ",model.linear.bias.item())
x_test= torch.Tensor([[7.0]])
y_test= model(x_test)
print("y_pred= ",y_test.data)
plt.plot(epoch_list, loss_list)
plt.xlabel("epoch")
plt.ylabel("loss_val")
plt.show()
使用SGD优化器图像:

使用ASGD优化器图像:

使用Adagrad优化器图像:

使用Adamax优化器图像:

以上就是关于pytorch实现线性回归的步骤以及代码,上述代码仅供大家参考学习,希望文本对大家熟悉pytorch的使用有帮助,更多pytorch相关的内容可以关注其他文章。
文本转载自脚本之家
免责声明:本站发布的内容(图片、视频和文字)以原创、转载和分享为主,文章观点不代表本网站立场,如果涉及侵权请联系站长邮箱:mmqy2019@163.com进行举报,并提供相关证据,查实之后,将立刻删除涉嫌侵权内容。
猜你喜欢
本篇文章给大家带来了关于Python的相关知识,其中主要整理了解析参数的三种方法相关问题,第一个选项是使用 argparse,它是一个流行的 Python 模块,专门用于命令行解析;另一种方法是读取 JSON 文件,我们可以在其中放置所有超参数;第三种也是鲜为人知的方法是使用 YAML 文件,下面一起来看一下,希望对大家有帮助。
这篇文章主要为大家介绍了python闭包和装饰器,具有一定的参考价值,感兴趣的小伙伴们可以参考一下,希望能够给你带来帮助
对于扑克牌21点相信是不少朋友的童年记忆吧,那么我们如果想要用python来实现这样一个小游戏,我们要怎样做呢?下面小编就给大家分享怎样用python写一个扑克牌21点小游戏的代码,感谢的朋友可以参考看看。
这篇文章主要介绍python实现跳表的内容,一些朋友可能对跳表是什么不是很了解,对此,接下来我们现了解一下跳表,再看python如何实现跳表,感兴趣的朋友就继续往下看吧。
我们知道OpenCV是一个用于图像处理、分析、机器视觉方面的开源函数库,这篇文章就主要给大家分享的是有关OpenCv库怎样实现绘制简单的图,小编觉得挺实用的,对新手认识OpenCv库有一定的帮助,因此分享给大家做个参考,接下来一起跟随小编看看吧。
成为群英会员,开启智能安全云计算之旅
立即注册关注或联系群英网络
7x24小时售前:400-678-4567
7x24小时售后:0668-2555666
24小时QQ客服
群英微信公众号
CNNIC域名投诉举报处理平台
服务电话:010-58813000
服务邮箱:service@cnnic.cn
投诉与建议:0668-2555555
Copyright © QY Network Company Ltd. All Rights Reserved. 2003-2020 群英 版权所有
增值电信经营许可证 : B1.B2-20140078 ICP核准(ICP备案)粤ICP备09006778号 域名注册商资质 粤 D3.1-20240008