PyTorch: Rewriting/D reworking the Dataset into the DataLoader
发布时间
阅读量:
阅读量
前言
众所周知,在PyTorch框架中广泛使用的Dataset和Dataloder类被设计用于高效管理数据加载任务。为了执行深度学习模型的训练任务,在获得一批数据后才开始处理后续操作。在PyTorch的一些案例教学中,默认采用torchvision.datasets模块中的MNIST、CIFAR-10等标准数据集作为训练素材;其基本操作流程大致如下:
# 下载并存放数据集
train_dataset = torchvision.datasets.CIFAR10(root="数据集存放位置",download=True)
# load数据
train_loader = torch.utils.data.DataLoader(dataset=train_dataset)
然而,在我们的模型训练过程中,需要采用自制的数据集而非官方来源。这时应该具体该如何操作呢?
我们为了实现自定义数据加载功能可以通过重写torch.utils.data.Dataset类中的__getitem__与__len__方法来导入我们的数据集。其中:
__getitem__方法负责返回数据集中对应索引位置的数据对象__len__方法则用于返回整个数据集的具体样本总数
改写
基于pytorch官方提供
全部评论 (0)
还没有任何评论哟~
