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

【刘二老师】pytorch深度学习笔记【08加载数据集】

【刘二老师】pytorch深度学习笔记【08加载数据集】

一、概念

  1. 用一个样本的随机梯度下降。
    优点:可以克服鞍点问题,性能好;缺点:是计算速度慢,优化时间长。
  2. 全部样本都用的Batch。
    优点:最大化的利用向量并行计算的优势,计算速度快;缺点:无法克服鞍点,性能会较差。
  3. 把上面两个综合起来,得出mini-Batch,综合了速度和性能。

二、架构原理

(一)

  • 内层每一次循环执行一个mini-Batch,迭代;外层表示训练的周期;两层组成嵌套循环。
  • Epoch:完整跑完整个数据集所有样本一遍 = 1 个 Epoch
  • Batch-Size:单次前向 + 反向传播,一次性扔进模型的样本(部分样本) 数量。
  • Iteration:内层循环跑一次 = 1 次 Iteration,也就是有多少个Batch(=总样本数 / Batch-Size)

(二)DataLoader

  1. Dataset(数据集类)
    存原始数据,负责「取单个样本」&「告知总共有多少样本」
  2. DataLoader(数据加载器)
    批量打包、打乱、多线程加载,负责把 Dataset 组装成训练用的一批一批数据
  3. 数据集 Dataset 需要支持索引,使 Dataloader 能访问到里面的每一个元素。
  • 如果数据集 Dataset 不能下标索引(比如只能从头到尾流式读取、不能跳着取),DataLoader 就没法随机采样、打乱数据。
  1. Dataloader 还需要知道 Dataset 的长度
  • DataLoader 计算一个 epoch 需要跑多少个批次:总批次数 = 总样本数 ÷ batch_size
  • 限制随机下标范围:不会生成超过数据集总数的索引,防止取数据时报错。
  1. shuffle:打乱数据集顺序。
  2. 分组:因为如图batch-size为 2,意味着两个一组,做成可迭代的Loader。
    第一次迭代给Batch1,第二次迭代给Batch2……

三、代码及详细讲解

importtorchimportnumpyasnpfromtorch.utils.dataimportDatasetfromtorch.utils.dataimportDataLoader# prepare datasetclassDiabetesDataset(Dataset):def__init__(self,filepath):xy=np.loadtxt(filepath,delimiter=',',dtype=np.float32)self.len=xy.shape[0]# shape(多少行,多少列),取0就是把有多少行拿出来让我们知道self.x_data=torch.from_numpy(xy[:,:-1])#要前八列self.y_data=torch.from_numpy(xy[:,[-1]])#要最后一列def__getitem__(self,index):returnself.x_data[index],self.y_data[index]#python中的return x,y就是返回一个元组(x,y)def__len__(self):returnself.lendataset=DiabetesDataset('diabetes.csv')#括号中是数据文件路径train_loader=DataLoader(dataset=dataset,batch_size=32,shuffle=True,num_workers=0)#num_workers 多线程#与上节课的一样classModel(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()defforward(self,x):x=self.sigmoid(self.linear1(x))x=self.sigmoid(self.linear2(x))x=self.sigmoid(self.linear3(x))returnx model=Model()# construct loss and optimizercriterion=torch.nn.BCELoss(reduction='mean')optimizer=torch.optim.SGD(model.parameters(),lr=0.01)# training cycle forward, backward, updateif__name__=='__main__':#不写这行会报错,要把下面的迭代代码封装到一个if语句中(或函数中),不能直接写这个循环。forepochinrange(100):fori,datainenumerate(train_loader,0):# train_loader 是先shuffle后mini_batch#enumerate是为了获得当前是第几次迭代#train_loader中的(x,y)元组就直接放到data中,而且train_loader直接把x,y转换成张量,所以不用加tensor。inputs,labels=data#inputs---x;labels----y,都是张量。y_pred=model(inputs)loss=criterion(y_pred,labels)print(epoch,i,loss.item())#backwardoptimizer.zero_grad()loss.backward()#updateoptimizer.step()

(一)

fromtorch.utils.dataimportDatasetfromtorch.utils.dataimportDataLoader
  • Dataset 是抽象类,无法实例化,实例化 ds=Dataset()会报错。因此要通过子类去继承 Dataset 来使用。
  • Dataloader 可以实例化。Dataloader 的功能是加载数据,分批、打乱数据,因此通过实例化来实现这个功能。

(二)

def__getitem__(self,index):returnself.x_data[index],self.y_data[index]
  • getitiem魔法方法:索引取样本,index是样本下标。

(三)

def__len__(self):returnself.len
  • len魔法方法:返回数据集总长度,调用len(dataset)时执行,返回总样本数量。

(四)

def__init__(self,filepath):
  1. 在构造数据集时,init下面有两种选择
  • 小数据(csv、小文本):__init__全量加载数据到内存,取样本直接内存读取,快、费内存;
  • 超大图像 / 分割数据:__init__只存文件路径,不取真实数据,取样本时临时读硬盘,慢、省内存。
  1. _init_:创建数据集对象时只运行 1 次.
  2. _getitem_(index)
  • 在小数据时:不用读硬盘,直接从内存取出第 i 组 x、y 返回。
  • 在大数据时:根据 index 拿到第 i 张图片路径,临时从硬盘读取图片。

(五)

train_loader=DataLoader(dataset=dataset,batch_size=32,shuffle=True,num_workers=0)
  1. Dataloader初始化代码,写四个方面:
  • 传递数据集,把定义的数据集对象dataset传进去。
  • batch_size:32,定义一个数据集的小批量有多少。
  • 是否要shuffle(打乱)
  • num_workers:读数据集构成mini-Batch时,是否要用多线程。也就是要不要并行,要几个并行。

四、MNIST数据集举例

  • datasets里面有MNIST类,用这个类来构造MNIST实例。
  • root:路径;要训练集还是测试集;ToTensor:转张量,缩放到(0,1)或(-1,1)这样的区间;download:如果没有这个数据集,要连线下载。
  • 训练数据集中,通常要shuffle;测试时不shuffle,输出的每一次顺序一样,方便观察结果。
  • 最后一行就是对loader进行迭代。

五、kaggle作业

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

相关文章:

  • 德州摩托车D本增驾全流程详解:从报名到拿证避坑指南
  • Containerlab实战系列之四:自动加载配置
  • FastAdmin仓库出入库管理插件|高效物资进销存系统(支持扫码打单与二次开发)
  • 2026保定财税管理公司选择指南:十大机构差异化能力深度解析 - 增长观测局
  • 2026邢台监控安装、监控维修厂家哪家好?本地实用选购指南与避坑要点 - mobible
  • 系规论文太难写?金老师团队帮你破局
  • UE5横板2D游戏开发:AI行为树与碰撞检测实战指南
  • 量化交易策略工程化实践:从双均线策略构建到回测验证
  • Python构建咖啡销售数据分析系统:从数据处理到智能预测
  • 基于具身智能体与专用分割模型的细粒度车辆损伤评估技术实践
  • 国内网络友好游戏平台盘点:无需加速器即可流畅使用 - 资讯综合
  • 深圳网站建设哪家口碑好:拒绝被割韭菜,教你从行业乱象中选出真正靠谱的服务商
  • Zookeeper集群部署与分布式锁实现实战指南
  • 2026年济宁大颗粒尿素批发商推荐哪家建议参考青州市天企源化肥有限公司 - 热点品牌推荐
  • FPS游戏外挂与吞子弹问题诊断:从网络同步到反作弊的全面解析
  • 2026沙河市网络布线,无线覆盖厂家推荐:安防监控与弱电工程怎么选?实用选购指南 - mobible
  • 2026年8月四川白酒品牌大挑选,哪家能脱颖而出引关注? - 企业推荐官
  • PSO-MPPT算法在光伏系统遮阴条件下的优化应用
  • 别瞎找Java培训了!3个狠招,一眼揪出烂机构
  • 2026桐乡外墙装修内墙装修避坑指南:5个常见坑+5条硬标准,靠谱公司推荐 - mobible
  • 【开源普惠・助力国产 AI】基于元初混沌熵控理论 —— AI 语料有序度智能清洗系统 完整开源
  • 基于Coze平台的火柴人心理学视频自动化生成工作流搭建指南
  • 2026肇庆浴室柜厂家哪家好,淋浴花洒厂家推荐避坑指南:5个挑选要点,帮你绕开90%的坑 - mobible
  • 零代码AI开发:DeepSeek与Cursor实战千问API设计
  • 第 7 章 舵机控制的高级话题 速度曲线、扭矩管理、通信可靠性、寿命维护——那些规格书不会告诉你的真相
  • Re:Linux系统篇(七) 开发工具篇 Chapter3:Makefile 从入门到精通 —— 依赖关系、伪目标、栈式推导与自动化构建全解
  • 2026内丘县电脑维修,电脑回收厂家推荐:5个避坑要点+4条实用选购指南 - mobible
  • Socket 管理详解——从原理到高性能架构设计(C++/Qt 实战)
  • 基于FPGA的AM调制系统实现:从数字信号处理到硬件设计实战
  • 大连装修公司**2026:半包和整装到底哪家强? - 米諾