Pytorch笔记之回归

文章目录


前言

以线性回归为例,记录Pytorch的基本使用方法。


一、导入库

python 复制代码
import numpy as np
import matplotlib.pyplot as plt
import torch
from torch.autograd import Variable # 定义求导变量
from torch import nn, optim # 定义网络模型和优化器

二、数据处理

将数据类型转为tensor,第一维度变为batch_size

python 复制代码
# 构建数据
x = np.random.rand(100)
noise = np.random.normal(0, 0.01, x.shape)
y = 0.1 * x + 0.2 + noise
# 数据处理
x_data = torch.FloatTensor(x.reshape(-1, 1))
y_data = torch.FloatTensor(y.reshape(-1, 1))
inputs = Variable(x_data)
target = Variable(y_data)

三、构建模型

1、继承nn.Module,定义一个线性回归模型。在__init__中定义连接层,定义前向传播的方法

2、实例化模型,定义损失函数与优化器

python 复制代码
# 继承模型
class LinearRegression(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc = nn.Linear(1, 1)
    def forward(self, x):
        out = self.fc(x)
        return out
# 定义模型
print('模型参数')
model = LinearRegression()
mse_loss = nn.MSELoss()
optimizer = optim.SGD(model.parameters(), lr=0.1)
for name, param in model.named_parameters():
    print('{}:{}'.format(name, param))

四、迭代训练

1、梯度清零:optimizer.zero_grad()

2、反向传播计算梯度值:loss.backward()

3、执行参数更新:optimizer.step()

循环迭代,定期输出损失值

python 复制代码
print('损失值')
for i in range(1001):
    out = model.forward(inputs)
    loss = mse_loss(out, target)
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()
    if i % 200 == 0:
        print(i, loss.item())

五、结果预测

绘制样本的散点图与预测值的折线图

python 复制代码
print('结果预测')
y_pred = model(x_data)
plt.plot(x, y, 'b.')
plt.plot(x, y_pred.data.numpy(), 'r-')
plt.show()

总结

使用Pytorch进行训练主要的三步:

(1)数据处理:将数据维度转换为(batch, *),数据类型转换为可训练的tensor;

(2)构建模型:继承nn.Module,定义连接层与运算方法,实例化,定义损失函数与优化器;

(3)迭代训练:循环迭代,依次执行梯度清零、梯度计算、参数更新。

相关推荐
丑小鸭是白天鹅1 小时前
嵌入式C语言学习笔记之枚举、联合体
c语言·笔记·学习
十一10243 小时前
FX10/20 (CYUSB401X)开发笔记5 固件架构
笔记
FakeOccupational3 小时前
【电路笔记 通信】AXI4-Lite协议 FPGA实现 & Valid-Ready Handshake 握手协议
笔记·fpga开发
奶黄小甜包4 小时前
C语言零基础第18讲:自定义类型—结构体
c语言·数据结构·笔记·学习
盼小辉丶4 小时前
PyTorch生成式人工智能——使用MusicGen生成音乐
pytorch·python·深度学习·生成模型
Moshow郑锴6 小时前
机器学习相关算法:回溯算法 贪心算法 回归算法(线性回归) 算法超参数 多项式时间 朴素贝叶斯分类算法
算法·机器学习·回归
Tiger Z6 小时前
《动手学深度学习v2》学习笔记 | 1. 引言
pytorch·深度学习·ai编程
rannn_1116 小时前
【MySQL学习|黑马笔记|Day7】触发器和锁(全局锁、表级锁、行级锁、)
笔记·后端·学习·mysql
草莓熊Lotso7 小时前
《详解 C++ Date 类的设计与实现:从运算符重载到功能测试》
开发语言·c++·经验分享·笔记·其他
_Kayo_13 小时前
node.js 学习笔记3 HTTP
笔记·学习