PyTorch重写DataSet类

PyTorch重写DataSet类


文章目录


前言

在之前沐神的Cifar-10分类 课程学习中,沐神是用的将每一类创建一个文件夹去完成图片的导入。此外我们还可以通过重写DataSet类来完成!

一、如何重写?

通过查看官方文档我们可知。

需要去重写__getitem__这个方法,去以一种特定的方法拿到一个数据。并且选择性的重写__len__这个方法,去返回整个数据集的大小。

二、具体代码

1.数据集格式

这个数据集是沐神课程上讲过的cifar-10数据集。

train和test文件夹分别为要进行训练和测试的图片。而训练数据的标签以csv文件存在trainLabels.csv文件中。

2.获取标签

python 复制代码
def read_csv_labels(fname):
    with open(fname,'r') as f:
        lines = f.readlines()[1:]
    tokens = [l.rstrip().split(',') for l in lines]
    return dict(((name,label) for name,label in tokens))

这里通过一个read_csv_labels的方法 将图片名字和标签以一个字典的方式返回

3.重写dataset

python 复制代码
class MyDateset(Dataset):
    def __init__(self,root_dir,state,label_dict=None):
        self.root_dir = root_dir
        self.state = state
        if label_dict is not None:
            self.label_dict = label_dict
        self.img_path = os.listdir(os.path.join(root_dir,state))
        # os.listdir 将当前文件夹下的图片名称按列表返回

    def __getitem__(self, idx):
        img = Image.open(os.path.join(self.root_dir,self.state,self.img_path[idx]))
        if self.state == 'train':
            img_num =self.img_path[idx].split('.')[0]
            # 这个取出来是数字.jpg 所以需要将.jpg舍去
            label = self.label_dict[img_num]
            return img,label
        else:
            return img

    def __len__(self):
        return len(self.img_path)

state参数表示此时是训练数据集还是测试数据集。

4.调用

python 复制代码
root_dir = "D:\\PytorchLearn\\cifar-10"
label_dict = read_csv_labels(os.path.join(root_dir,"trainLabels.csv"))

train_dataset = MyDateset(root_dir,'train',label_dict)

test_dataset = MyDateset(root_dir,'test')

train_iter = torch.utils.data.DataLoader(train_dataset,batch_size=8,shuffle=True)

总结

以上就是重写DataSet的方法,有不足之处还望各位指出。

相关推荐
蒋星熠36 分钟前
实证分析:数据驱动决策的技术实践指南
大数据·python·数据挖掘·数据分析·需求分析
youngfengying42 分钟前
《轻量化 Transformers:开启计算机视觉新篇》
人工智能·计算机视觉
独行soc2 小时前
2025年渗透测试面试题总结-250(题目+回答)
网络·驱动开发·python·安全·web安全·渗透测试·安全狮
一晌小贪欢2 小时前
Pandas操作Excel使用手册大全:从基础到精通
开发语言·python·自动化·excel·pandas·办公自动化·python办公
搞科研的小刘选手3 小时前
【同济大学主办】第十一届能源资源与环境工程研究进展国际学术会议(ICAESEE 2025)
大数据·人工智能·能源·材质·材料工程·地理信息
MARS_AI_3 小时前
云蝠智能 VoiceAgent 2.0:全栈语音交互能力升级
人工智能·自然语言处理·交互·信息与通信·agi
top_designer3 小时前
Substance 3D Stager:电商“虚拟摄影”工作流
人工智能·3d·设计模式·prompt·技术美术·教育电商·游戏美术
雷神大青椒3 小时前
离别的十字路口: 是否还记得曾经追求的梦想
人工智能·程序人生·职场和发展·玩游戏
IT痴者4 小时前
《PerfettoSQL 的通用查询模板》---Android-trace
android·开发语言·python
m0_650108244 小时前
多模态大模型 VS. 图像视频生成模型浅析
人工智能·技术边界与协同·mllm与生成模型·技术浅谈