码农知识堂 - 1000bd
  •   Python
  •   PHP
  •   JS/TS
  •   JAVA
  •   C/C++
  •   C#
  •   GO
  •   Kotlin
  •   Swift
  • Pytorch笔记之回归


    文章目录

    • 前言
    • 一、导入库
    • 二、数据处理
    • 三、构建模型
    • 四、迭代训练
    • 五、结果预测
    • 总结


    前言

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


    一、导入库

    import numpy as np
    import matplotlib.pyplot as plt
    import torch
    from torch.autograd import Variable # 定义求导变量
    from torch import nn, optim # 定义网络模型和优化器
    
    • 1
    • 2
    • 3
    • 4
    • 5

    二、数据处理

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

    # 构建数据
    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
    • 2
    • 3
    • 4
    • 5
    • 6
    • 7
    • 8
    • 9

    三、构建模型

    1、继承nn.Module,定义一个线性回归模型。在__init__中定义连接层,定义前向传播的方法
    2、实例化模型,定义损失函数与优化器

    # 继承模型
    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
    • 2
    • 3
    • 4
    • 5
    • 6
    • 7
    • 8
    • 9
    • 10
    • 11
    • 12
    • 13
    • 14
    • 15

    四、迭代训练

    1、梯度清零:optimizer.zero_grad()
    2、反向传播计算梯度值:loss.backward()
    3、执行参数更新:optimizer.step()
    循环迭代,定期输出损失值

    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())
    
    • 1
    • 2
    • 3
    • 4
    • 5
    • 6
    • 7
    • 8
    • 9

    五、结果预测

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

    print('结果预测')
    y_pred = model(x_data)
    plt.plot(x, y, 'b.')
    plt.plot(x, y_pred.data.numpy(), 'r-')
    plt.show()
    
    • 1
    • 2
    • 3
    • 4
    • 5


    总结

    使用Pytorch进行训练主要的三步:
    (1)数据处理:将数据维度转换为(batch, *),数据类型转换为可训练的tensor;
    (2)构建模型:继承nn.Module,定义连接层与运算方法,实例化,定义损失函数与优化器;
    (3)迭代训练:循环迭代,依次执行梯度清零、梯度计算、参数更新。

  • 相关阅读:
    闭关之 Vulkan 应用开发指南笔记(二):队列、命令、移动数据和展示
    关于go语言的那点事
    PTA题目 分段计算居民水费
    驱动开发:内核无痕隐藏自身分析
    c++ 友元函数 友元类
    数学建模之线性规划(含MATLAB代码)
    1677. 发票中的产品金额
    封装了几个CAPL发送诊断相关函数,具有较高的可复用性
    如何搭建Spring项目,修改目录,修改pom.xml文件?
    【JVM】JVisualVM的介绍、使用和GC过程
  • 原文地址:https://blog.csdn.net/qq_53715621/article/details/133607569
  • 最新文章
  • 沪漂五周年了:我越来越迷茫了
    Agentic Skill Routing 实战:别再把所有 Skill 塞进 AI Agent 上下文
    MySQL-Seconds_behind_master的精度误差
    [MAF预定义ChatClient中间件-03]CachingChatClient——利用缓存省钱省时间
    AI的至暗历史:从万众期待到被政府撤资,AI的两次死亡徘徊
    Agent OS :五种驯服不确定性的范式
    PortSwigger SQL注入LAB11
    数据库即时编译JIT
    [Begin]AI Learn Data Day 0
    深度学习进阶(二十七)现代 LLM 的核心架构设计其二:SwiGLU
  • 热门文章
  • 十款代码表白小特效 一个比一个浪漫 赶紧收藏起来吧!!!
    奉劝各位学弟学妹们,该打造你的技术影响力了!
    五年了,我在 CSDN 的两个一百万。
    Java俄罗斯方块,老程序员花了一个周末,连接中学年代!
    面试官都震惊,你这网络基础可以啊!
    你真的会用百度吗?我不信 — 那些不为人知的搜索引擎语法
    心情不好的时候,用 Python 画棵樱花树送给自己吧
    通宵一晚做出来的一款类似CS的第一人称射击游戏Demo!原来做游戏也不是很难,连憨憨学妹都学会了!
    13 万字 C 语言从入门到精通保姆级教程2021 年版
    10行代码集2000张美女图,Python爬虫120例,再上征途
小工具 小游戏
Copyright © 2022 侵权请联系2656653265@qq.com    京ICP备2022015340号-1

京公网安备 11010502049817号