当前位置: 首页 > news >正文

04-pytorch构建线性回归

1. 过程说明

在pytorch中进行模型构建的整个流程一般分为四个步骤:

  • 准备训练集数据

  • 构建要使用的模型

  • 设置损失函数和优化器

  • 模型训练

要使用的API:

  • 使用 PyTorch 的 nn.MSELoss() 代替平方损失函数

  • 使用 PyTorch 的 data.DataLoader 代替数据加载器

  • 使用 PyTorch 的 optim.SGD 代替优化器

  • 使用 PyTorch 的 nn.Linear 代替假设函数

2. 代码

import torch from torch.utils.data import TensorDataset # 构造数据集对象 from torch.utils.data import DataLoader # 数据加载器 from torch import nn # nn模块中有平方损失函数和假设函数 from torch import optim # optim模块中有优化器函数 from sklearn.datasets import make_regression # 创建线性回归模型数据集 import matplotlib.pyplot as plt plt.rcParams['font.sans-serif'] = ['SimHei'] # 用来正常显示中文标签 plt.rcParams['axes.unicode_minus'] = False # 用来正常显示负号 # 构造数据集 def create_dataset(): x, y, coef = make_regression(n_samples=100, n_features=1, noise=10, coef=True, bias=14.5, random_state=0) # 将构建数据转换为张量类型 x = torch.tensor(x) y = torch.tensor(y) return x, y, coef # 训练模型 def train(): # 构造数据集 x, y, coef = create_dataset() # 构造数据集对象 dataset = TensorDataset(x, y) # 构造数据加载器 # dataset=:数据集对象 # batch_size=:批量训练样本数据 # shuffle=:样本数据是否进行乱序 dataloader = DataLoader(dataset=dataset, batch_size=16, shuffle=True) # 构造模型 # in_features指的是输入的二维张量的大小,即输入的[batch_size, size]中的size # out_features指的是输出的二维张量的大小,即输出的[batch_size,size]中的size model = nn.Linear(in_features=1, out_features=1) # 构造平方损失函数 criterion = nn.MSELoss() # 构造优化函数 # params=model.parameters():训练的参数,w和b # lr=1e-2:学习率, 1e-2为10的负二次方 print("w和b-->", list(model.parameters())) print("w-->", model.weight) print("b-->", model.bias) optimizer = optim.SGD(params=model.parameters(), lr=1e-2) # 初始化训练次数 epochs = 100 # 损失的变化 epoch_loss = [] total_loss=0.0 train_sample=0.0 for _ in range(epochs): for train_x, train_y in dataloader: # 将一个batch的训练数据送入模型 y_pred = model(train_x.type(torch.float32)) # 计算损失值,均方误差,当前批次所有样本的平均误差 loss = criterion(y_pred, train_y.reshape(-1, 1).type(torch.float32)) total_loss += loss.item() # loss是平均误差,所以样本数+1 train_sample += 1 # 梯度清零 optimizer.zero_grad() # 自动微分(反向传播) loss.backward() # 更新参数 optimizer.step() # 计算所有batch的平均误差作为当前epoch的误差 epoch_loss.append(total_loss/train_sample) # 打印回归模型的w print(model.weight) # 打印回归模型的b print(model.bias) # 绘制损失变化曲线 plt.plot(range(epochs), epoch_loss) plt.title('损失变化曲线') plt.grid() plt.show() # 绘制拟合直线 plt.scatter(x, y) x = torch.linspace(x.min(), x.max(), 1000) y1 = torch.tensor([v * model.weight + model.bias for v in x]) y2 = torch.tensor([v * coef + 14.5 for v in x]) plt.plot(x, y1, label='训练') plt.plot(x, y2, label='真实') plt.grid() plt.legend() plt.show() if __name__ == '__main__': train()

http://www.jsqmd.com/news/1272086/

相关文章:

  • TMS570LS0914 ePWM/eCAP/eQEP模块实战:电机控制中的配置、联动与调试
  • 大语言模型自我笔记机制:提升复杂推理稳定性的关键技术
  • 嵌入式音频AGC算法:动态VAD与混合增益实现语音清晰度与自然度平衡
  • TMS320DM35x USB控制器编程实战:从架构解析到DMA优化
  • 问卷互填平台横向评测:问卷星、腾讯问卷、球球问卷,到底怎么选
  • 《Web前端工程师修炼之道》学习笔记:第二部分
  • 长三角注塑机工业设计优选 深耕设备外观结构全案服务,塑胶设备外观设计/设备外观设计/半导体设备外观设计,工业设计企业案例 - 品牌推荐师
  • 90% 的人都搞错过的国外 AI 名词,一篇给你全理清楚
  • Em Dash与AI:提升技术文档可读性的实用指南
  • C#开发OPC UA客户端:工业数据采集实战指南
  • AI量化投资平台架构解析与风险管理实践
  • C语言分支与循环结构详解与应用实践
  • 工业级AI智能体的关键技术架构与落地实践
  • Python实战:用LangChain构建高效RAG工作流
  • 2026 年新发布:吉安诚信的小C中C,大C护栏批发厂家哪个好,打破认知:小C中C的护栏,竟是保护大C的秘密?-鼎泽橡塑科技 - 实业推荐官【官方】
  • TI C6678 DSP时钟与复位系统配置详解:从PLL原理到实战避坑
  • C++实现Windows开机自启动:注册表操作与Wow64兼容性详解
  • EKF-SLAM可观测性分析与Matlab实现
  • Gemma 4开源大模型:工程化实践与多模态技术解析
  • C++内存管理与性能优化实战:从智能指针到并发编程的进阶指南
  • 二分查找解决LeetCode 1283最小除数问题
  • Windows 11无线网卡驱动安装与优化全指南
  • 【python】开发了一个电子桌面桌宠,会要饭,满屏跑,效果太棒了
  • 2026黄鹤杯网络安全人才创新大赛学生组(Misc部分)
  • ChatGPT优化服务商如何助力企业AI价值转化
  • RAG技术解析:如何提升大模型的事实准确性
  • 从SVN迁移到Git:原理、工具与自动化实践指南
  • C++编译原理全解析:从源码到可执行文件的完整流程与Java对比
  • 杭州燕壹画室值不值得去?一份来自实地走访的真实测评! - 资讯报道
  • 小论文/大论文必备 | RT-DETR多模态目标检测、绘制曲线对比图 | 引入多种绘制曲线对比图,包括mAP0.5,mAP0.5:0.95,Loss损失变化的曲线对比