【问题标题】:How to create a train-val split in custom image datasets using PyTorch?如何使用 PyTorch 在自定义图像数据集中创建训练验证拆分?
【发布时间】:2020-06-19 23:12:22
【问题描述】:

我想从我原来的 trainset 创建一个 train+val set。该目录分为训练和测试。我加载了原始训练集并希望将其拆分为训练集和验证集,以便我可以使用 train_loaderval_loader 在训练期间评估验证损失。

没有很多关于此的文档可以清楚地解释事情。

【问题讨论】:

    标签: machine-learning deep-learning computer-vision pytorch torchvision


    【解决方案1】:

    查看答案here

    我也把它贴在下面了。

    ================================================ =======

    使用ImageFolder 读取数据。任务是二值图像分类,数据集中有 498 幅图像,平均分布在两个类别中(每个类别 249 幅图像)。

    img_dataset = ImageFolder(..., transforms=t)
    

    1。 SubsetRandomSampler

    dataset_size = len(img_dataset)
    dataset_indices = list(range(dataset_size))
    
    np.random.shuffle(dataset_indices)
    
    val_split_index = int(np.floor(0.2 * dataset_size))
    
    train_idx, val_idx = dataset_indices[val_split_index:], dataset_indices[:val_split_index]
    
    train_sampler = SubsetRandomSampler(train_idx)
    val_sampler = SubsetRandomSampler(val_idx)
    
    
    train_loader = DataLoader(dataset=img_dataset, shuffle=False, batch_size=8, sampler=train_sampler)
    validation_loader = DataLoader(dataset=img_dataset, shuffle=False, batch_size=1, sampler=val_sampler)
    

    2。 random_split

    在这 498 张图片中,随机分配 400 张用于训练,其余 98 张用于验证。

    dataset_train, dataset_valid = random_split(img_dataset, (400, 98))
    
    train_loader = DataLoader(dataset=dataset_train, shuffle=True, batch_size=8)
    val_loader = DataLoader(dataset=dataset_valid, shuffle=False, batch_size=1)
    

    3。 WeightedRandomSampler

    如果有人在这里偶然发现WeightedRandomSampler,请查看@ptrblck 的答案here,以了解下面发生的情况。

    现在,WeightedRandomSampler 如何适合创建 train+val 集?因为与SubsetRandomSamplerrandom_split() 不同,我们不会在这里拆分train 和val。我们只是确保每批在训练期间获得相同数量的类。

    所以,我猜我们需要使用WeightedRandomSampler after random_split()SubsetRandomSampler。但这并不能确保 train 和 val 在类之间具有相似的比率。

    target_list = []
    
    for _, t in imgdataset:
        target_list.append(t)
    
    target_list = torch.tensor(target_list)
    target_list = target_list[torch.randperm(len(target_list))]
    
    # get_class_distribution() is a function that takes in a dataset and 
    # returns a dictionary with class count. In this case, the 
    # get_class_distribution(img_dataset)  returns the following - 
    # {'class_0': 249, 'class_0': 249}
    class_count = [i for i in get_class_distribution(img_dataset).values()]
    class_weights = 1./torch.tensor(class_count, dtype=torch.float) 
    
    class_weights_all = class_weights[target_list]
    
    weighted_sampler = WeightedRandomSampler(
        weights=class_weights_all,
        num_samples=len(class_weights_all),
        replacement=True
    )
    

    【讨论】:

      猜你喜欢
      • 2023-04-02
      • 2020-02-29
      • 1970-01-01
      • 2018-08-28
      • 1970-01-01
      • 2023-01-21
      • 2019-05-01
      • 1970-01-01
      • 2016-09-13
      相关资源
      最近更新 更多