Advertisement

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)

还没有任何评论哟~