pytorch导入数据集

1、概念:

Dataset:一种数据结构,存储数据及其标签

Dataloader:一种工具,可以将Dataset里的数据分批、打乱、批量加载并进行迭代等

(方便模型训练和验证)

Dataset就像一个大书架,存放着带有标签的数据书籍,并且这些书有编号(0,1,2...);

而Dataloader就像一个图书管理员,负责从书架上按需取出书籍并分批提供给读者。

2、Dataset的组织形式

train:训练集 val:验证集

一种方式是label作为数据文件夹的名字,

另一种方式是label和数据本身分开成两个文件夹(label文件夹里装的是和每个数据对应的.txt)

3、处理图像:PIL(Python Imaging Library).Image

|--------------------------------------|--------------------------------------|
| pip install Pillow | 安装PIL |
| from PIL import Image | 引入Image类(代表图像对象, 可以通过创建Image实例来操作图像) |
| img=Image.open('图像路径') 打开图像 | img.show() 显示图像 |
| print(img.size) 输出(宽度,高度) | print(img.format) 输出图像格式(JPEG、PNG等) |
| resized_img=img.resize((宽度,高度)) 调整大小 | |
| resized_img=img.save('新路径') 保存为新文件 | |

4、处理目录和文件:os

|---------------------------------------------|--------------------------------------|
| import os | |
| cur_dir=os.getcwd() | 获取当前工作目录 |
| files=os.listdir(cur_dir) | 列举当前目录下的所有子目录(文件和文件夹) |
| os.makedirs('new_folder') | 创建新文件夹(如果不存在) |
| os.remove('file.txt') | 删除文件(os.rmdir('empty_folder')删除空文件夹) |
| os.path.exists('some_path') | 检查路径是否存在 |
| file_path=os.path.join('folder','file.txt') | 拼接路径 |
| abs_path=os.path.abspath('file.txt) | 获取文件的绝对路径 |

5、代码

python 复制代码
from torch.utils.data import Dataset #从torch的常用工具箱utils中拿data工具,然后引入Dataset类
from PIL import Image #处理图片要用到
import os #访问目录、获取图片的地址要用到

class MyData(Dataset): #让MyData类继承Dataset类
    def __init__(self,root_dir,label_dir): #数据集的初始化:要用到根目录和标签目录(这里把label作为数据文件夹的名字了)
        self.root_dir=root_dir
        self.label_dir=label_dir
        self.path=os.path.join(self.root_dir,self.label_dir) #根目录+标签目录=数据集的路径
        self.img_dir_list=os.listdir(self.path) #列举数据集目录下的每个数据(文件)

    def __getitem__(self,idx): #获取索引对应的数据
        img_dir=self.img_dir_list[idx] #得到索引对应的数据文件
        img_path=os.path.join(self.root_dir,self.label_dir,img_dir) #数据集路径+数据文件=数据文件路径
        img=Image.open(img_path)
        label=self.label_dir
        return img,label

    def __len__(self):
        return len(self.img_dir_list) #数据长度=数据集目录下的子文件数量

root_dir=r"dataset/hymenoptera_data/train"
ants_label_dir="ants"
ants_dataset=MyData(root_dir,ants_label_dir)
bees_label_dir="bees"
bees_dataset=MyData(root_dir,bees_label_dir)

train_dataset=ants_dataset+bees_dataset
相关推荐
iiiiii112 分钟前
【论文阅读笔记】Reparameterization Proximal Policy Optimization (RPO):将PPO的稳定性引入可微分强化学习
论文阅读·人工智能·笔记·深度学习·机器学习·机器人·具身智能
GitCode官方12 分钟前
openPangu-2.0-Pro 模型及技术报告正式开源上线 AtomGit AI
人工智能·华为·开源·harmonyos·atomgit
少年最是14 分钟前
体协作系统,用于分析聊天记录中特定人物的荣格八维人格类型。 本文使用 CC-BY-NC-SA . 协议。转载或者 AI 模型/智能体 ...
人工智能
xiaobaihuoke16 分钟前
多平台买家账号批量管理技术方案:环境隔离与自动化操作实践
人工智能·亚马逊鲲鹏系统·多平台账号管理·跨境电商买家号
HiEV17 分钟前
Mobileye 3.0:自动驾驶科技问题基本解决,Shashua押注物理AI
人工智能·自动驾驶
段一凡-华北理工大学19 分钟前
AI推动工业智能化转型~系列文章05:特征工程:工业 AI 的第一生产力
数据库·人工智能·分布式·搜索引擎·特征工程·高炉智能化
落苜蓿蓝19 分钟前
Java 循环中对象复用导致属性覆盖?从 JVM 内存模型讲解原因
java·jvm·python
李伟_Li慢慢21 分钟前
凸包算法
人工智能·算法
徐礼昭|商派软件市场负责人26 分钟前
AI Agent 蜂群协作与经验基因化:从单体执行到自组织系统的工程路径
大数据·人工智能·算法·智能体·多agent
厚皮龙28 分钟前
相机标定_角点检测_重投影误差_总结
人工智能