【问题标题】:Splitting custom PyTorch dataset into train loader and validation loader: Length of both same, even though dataset was split?将自定义 PyTorch 数据集拆分为训练加载器和验证加载器:两者的长度相同,即使数据集已拆分?
【发布时间】:2020-11-15 10:03:12
【问题描述】:

我正在尝试将其中一个 Pytorch 自定义数据集 (MNIST) 拆分为训练集和验证集,如下所示:

def get_train_valid_splits(data_dir,
                           batch_size,
                           random_seed=1,
                           valid_size=0.2,
                           shuffle=True,
                           num_workers=4,
                           pin_memory=False):

    normalize = transforms.Normalize((0.1307,), (0.3081,))  # MNIST

    # define transforms
    valid_transform = transforms.Compose([
            transforms.ToTensor(),
            normalize
        ])

        train_transform = transforms.Compose([
            transforms.ToTensor(),
            normalize
        ])

    # load the dataset
    train_dataset = datasets.MNIST(root=data_dir, train=True,
                download=True, transform=train_transform)

    valid_dataset = datasets.MNIST(root=data_dir, train=True,
                download=True, transform=valid_transform)

    dataset_size = len(train_dataset)
    indices = list(range(dataset_size))
    split = int(np.floor(valid_size * dataset_size))

    
    if shuffle == True:
        np.random.seed(random_seed)
        np.random.shuffle(indices)
    

    train_idx, valid_idx = indices[split:], indices[:split]

    train_sampler = sampler.SubsetRandomSampler(train_idx)
    valid_sampler = sampler.SubsetRandomSampler(valid_idx)

    print(len(train_sampler))
    print(len(valid_sampler))

    train_loader = torch.utils.data.DataLoader(train_dataset,
                    batch_size=batch_size, sampler=train_sampler,
                    num_workers=num_workers, pin_memory=pin_memory)

    valid_loader = torch.utils.data.DataLoader(valid_dataset,
                    batch_size=batch_size, sampler=valid_sampler,
                    num_workers=num_workers, pin_memory=pin_memory)

    print(len(train_loader.dataset))
    print(len(valid_loader.dataset))

    return (train_loader, valid_loader)

调用该函数后,我注意到要采样的索引结果看起来正确,48000 和 12000:

print(len(train_sampler))
print(len(valid_sampler))

但是当我查看与 train_loader 和 valid_loader 关联的数据集的长度时:

print(len(train_loader.dataset))
print(len(valid_loader.dataset))

两者的长度相同:60000!知道这里发生了什么吗?为什么它给两者的长度相同,即使我清楚地按索引分割它?

【问题讨论】:

    标签: python validation pytorch mnist dataloader


    【解决方案1】:

    train_loadervalid_loader 长度相同的原因是因为您对 train_datasetvalid_dataset 使用了相同的数据。

    你想要

    valid_dataset = datasets.MNIST(root=data_dir, train=False,
                                   download=True, transform=valid_transform)
    

    (不是train=True)下载验证集。

    【讨论】:

      【解决方案2】:

      这是因为数据加载器不会修改您传递给它的数据集,而是在您尝试通过迭代访问数据时“应用”诸如批量大小、采样器等之类的内容。你的问题是你使用len(loader.dataset) 它给你提供的数据集的长度没有修改,当你真的想要len(loader) 这是“应用”诸如批量大小和采样器之类的东西之后数据集的长度。

      import torch
      import numpy as np
      
      dataset = np.random.rand(100,200)
      sampler = torch.utils.data.SubsetRandomSampler(list(range(70)))
      
      loader = torch.utils.data.DataLoader(dataset, sampler=sampler)
      print(len(loader)) 
      >>> 70
      print(len(loader.dataset))
      >>> 100
      

      注意:len的结果会受batch size的影响:

      # with batch size
      loader = torch.utils.data.DataLoader(dataset, sampler=sampler, batch_size=2)
      print(len(loader)) 
      >>> 35
      print(len(loader.dataset))
      >>> 100
      

      【讨论】:

      • 谢谢!所以 len(loader) 会给你 num_samples/batch_size 对吗?那么要获取完整数据集的大小,您需要执行 len(loader)*batch_size,还是有更简单的方法来执行此操作?
      • @user6496380 刚刚通过响应更新,但是是的,这是正确的,您需要将 len 乘以 batch_size。但是,请注意,如果 num_samples 不能完美地划分为批大小,那么当您将 batch_size 重新乘以 len 时,您将无法得到准确的数字。
      • 我明白了,谢谢!最后一个问题:除了使用 train_loader、valid_loader 之外,您是否可以自己在 DataLoader 所做的相同索引上拆分数据集(使用 sampler=...),以获得 train_dataset 和 valid_dataset?基本上所有的数据。
      • 您始终可以使用torch.utils.data.random_split() 之类的内容。在这种情况下,您将使用随机采样器而不是子集随机采样器,因为数据集在传递给数据加载器之前已经被拆分。
      猜你喜欢
      • 2023-04-02
      • 1970-01-01
      • 1970-01-01
      • 2020-02-29
      • 2016-09-13
      • 2019-05-01
      • 2020-10-01
      • 2019-04-22
      • 2018-11-05
      相关资源
      最近更新 更多