码农知识堂 - 1000bd
  •   Python
  •   PHP
  •   JS/TS
  •   JAVA
  •   C/C++
  •   C#
  •   GO
  •   Kotlin
  •   Swift
  • 使用Dataset 和DataLoader 加载数据集


    一、加载数据过程

    PyTorch 数据加载实用程序的核心是 torch.utils.data.DataLoader 类。 它表示可在数据集上迭代的 Python,并支持

    • 映射式和迭代式的数据集,
    • 自定义数据加载顺序,
    • 自动批次,
    • 单进程和多进程数据加载,
    • 自动内存固定。

    这些选项由 DataLoader 的构造函数参数配置,构造函数的签名如下:

    如下如显示了dataLoader的过程,shuffle将Dataset里的数据打乱,batch_size=2

    二、模型建立流程

    1、准备数据集(Dataset和DataLoader)2、继承Module类设计自己的模型

    3、使用PyTorch APi 构造损失函数和优化器  4、采用前向传播、返向回馈、更新 反复训练。

    三、代码实现

    import torch.nn
    import numpy as np
    from torch.utils.data import Dataset, DataLoader

    class DiabetesDataset(Dataset):
       
    def __init__(self, filepath):
            xy = np.loadtxt(filepath,
    delimiter=',', dtype=np.float32)
           
    self.len = xy.shape[0]
           
    self.x_data = torch.from_numpy(xy[:, :-1])
           
    self.y_data = torch.from_numpy(xy[:, [-1]])

       
    def __getitem__(self, index):
           
    return self.x_data[index], self.y_data[index]

       
    def __len__(self):
           
    return self.len


    dataset = DiabetesDataset(
    'diabetes.csv.gz')
    train_loader = DataLoader(
    dataset=dataset, batch_size=64, shuffle=True, num_workers=2)


    # 继承类Module,自动会实现反向计算图
    class Model(torch.nn.Module):
       
    # 构造方法
       
    def __init__(self):
           
    super(Model, self).__init__()
           
    self.linear1 = torch.nn.Linear(8, 6)
           
    self.linear2 = torch.nn.Linear(6, 4)
           
    self.linear3 = torch.nn.Linear(4, 1)
           
    self.sigmoid = torch.nn.Sigmoid()

       
    def forward(self, x):
            x =
    self.sigmoid(self.linear1(x))
            x =
    self.sigmoid(self.linear2(x))
            x =
    self.sigmoid(self.linear3(x))
           
    return x


    model = Model()

    criterion = torch.nn.BCELoss(
    size_average=True)
    optimizer = torch.optim.SGD(model.parameters(),
    lr=0.1)

    if __name__=='__main__':
       
    for epoch in range(100):
           
    for i, data in enumerate(train_loader, 0):
               
    #1.prepare data
               
    inputs, labels = data
               
    #2.Forward
                
    y_pred = model(inputs)
                loss = criterion(y_pred, labels)
               
    print(epoch, loss.item())
               
    #3.Backward
               
    optimizer.zero_grad()
                loss.backward()
               
    #4.Update
               
    optimizer.step()

    四、运行结果

  • 相关阅读:
    【大禹DGC】1-基本介绍
    【MySQL 8.0新特性】窗口函数
    stack&queue&priority_queue
    2009-2018年31省份旅游收入(入境、国内、总收入;第三产值;GDP)
    058_末晨曦Vue技术_过渡 & 动画之过渡的类名
    深度剖析Istio共享代理新模式Ambient Mesh
    基于Monkey的稳定性测试
    go语言并发实战——日志收集系统(八) go语言操作etcd以及利用watch实现对键值的监控
    Javascript知识【jQuery:数组遍历和事件】
    Kubernetes来去今生与基础理论
  • 原文地址:https://blog.csdn.net/axiaoquan/article/details/127649061
  • 最新文章
  • 【JVM】编译执行与解释执行的区别是什么?JVM 使用哪种方式?
    用 Hashids 优雅解决 C 端自增 ID 暴露问题
    V8引擎 精品漫游指南--Ignition篇(上) 指令 栈帧 槽位 调用约定 内存布局 基础内容
    LLVM Pass快速入门(四):代码插桩
    milkup:桌面端 markdown AI续写和即时渲染
    基于项目工程构建SBOM(软件物料清单)的研究
    鸿蒙应用开发UI基础第二节:鸿蒙应用程序框架核心解析与实操
    .NET 中如何快速实现 List 集合去重?
    扣子Coze实战:从0到1打造抖音+小红书热点监控智能体
    浅谈数据访问层
  • 热门文章
  • 十款代码表白小特效 一个比一个浪漫 赶紧收藏起来吧!!!
    奉劝各位学弟学妹们,该打造你的技术影响力了!
    五年了,我在 CSDN 的两个一百万。
    Java俄罗斯方块,老程序员花了一个周末,连接中学年代!
    面试官都震惊,你这网络基础可以啊!
    你真的会用百度吗?我不信 — 那些不为人知的搜索引擎语法
    心情不好的时候,用 Python 画棵樱花树送给自己吧
    通宵一晚做出来的一款类似CS的第一人称射击游戏Demo!原来做游戏也不是很难,连憨憨学妹都学会了!
    13 万字 C 语言从入门到精通保姆级教程2021 年版
    10行代码集2000张美女图,Python爬虫120例,再上征途
小工具 小游戏
Copyright © 2022 侵权请联系2656653265@qq.com    京ICP备2022015340号-1

京公网安备 11010502049817号