【问题标题】:torch - subsample each dataset differently and concatenate them火炬 - 对每个数据集进行不同的子采样并将它们连接起来
【发布时间】:2022-10-01 20:12:29
【问题描述】:

我有两个数据集,但一个比另一个大,我想对它进行子采样(在每个时期重新采样)。

我可能无法使用 dataloader 参数采样器,因为我会将已经连接的数据集传递给 Dataloader。

我如何简单地实现这一目标?

我认为一种解决方案是编写一个类 SubsampledDataset(IterableDataset) ,每次调用 __iter__ 时(每个时期)都会重新采样。

(或者更好地使用地图样式的数据集,但是是否有一个钩子会在每个时期都被调用,比如__iter__gets?)

    标签: python-3.x torch


    【解决方案1】:

    这是我到目前为止所拥有的(未经测试)。用法:

    dataset1: Any = ...
    # subsample original_dataset2, so that it is equally large in each epoch
    dataset2 = RandomSampledDataset(original_dataset2, num_samples=len(dataset1))
    
    concat_dataset = ConcatDataset([dataset1, dataset2])
    
    data_loader = torch.utils.data.DataLoader(
        concat_dataset,
        sampler=RandomSamplerWithNewEpochHook(dataset2.new_epoch_hook, concat_dataset)
    )
    

    结果是 concat_dataset 将在每个 epoch (RandomSampler) 中进行混洗,此外,dataset2 组件是(可能更大) original_dataset2 的新样本,在每个 epoch 中都不同。

    您可以通过执行以下操作添加更多要进行子采样的数据集:

    sampler=RandomSamplerWithNewEpochHook(dataset2.new_epoch_hook
    

    这个:

    sampler=RandomSamplerWithNewEpochHook(lambda: dataset2.new_epoch_hook and dataset3.new_epoch_hook and dataset4.new_epoch_hook, ...
    

    代码:

    class RandomSamplerWithNewEpochHook(RandomSampler):
        """ Wraps torch.RandomSampler and calls supplied new_epoch_hook before each epoch. """
        
        def __init__(self, new_epoch_hook: Callable, data_source: Sized, replacement: bool = False,
                     num_samples: Optional[int] = None, generator=None):
            super().__init__(data_source, replacement, num_samples, generator)
            self.new_epoch_hook = new_epoch_hook
    
        def __iter__(self):
            self.new_epoch_hook()
            return super().__iter__()
    
    
    class RandomSampledDataset(Dataset):
        """ Subsamples a dataset. The sample is different in each epoch.
    
        This helps when concatenating datasets, as the subsampling rate can be different for each dataset.
        
        Call new_epoch_hook before each epoch. (This can be done using e.g. RandomSamplerWithNewEpochHook.)
    
        This would be arguably harder to achieve with a concatenated dataset and a sampler argument to Dataloader. The
        sampler would have to be aware of the indices of subdatasets' items in the concatenated dataset, of the subsampling 
        for each subdataset."""
        def __init__(self, dataset, num_samples, transform=lambda im: im):
            self.dataset = dataset
            self.transform = transform
            self.num_samples = num_samples
    
            self.sampler = RandomSampler(dataset, num_samples=num_samples)
            self.current_epoch_samples = None
    
        def new_epoch_hook(self):
            self.current_epoch_samples = torch.tensor(iter(self.sampler), dtype=torch.int)
    
        def __len__(self):
            return self.num_samples
    
        def __getitem__(self, item):
            if item < 0 or item >= len(self):
                raise IndexError
    
            img = self.dataset[self.current_epoch_samples[item].item()]
    
            return self.transform(img)
    

    【讨论】:

      【解决方案2】:

      您可以通过提高StopIteration 来停止迭代。这个错误被Dataloader 捕获并简单地停止迭代。所以你可以做这样的事情:

      class SubDataset(Dataset):
          """SubDataset class."""
          def __init__(self, dataset, length):
              self.dataset = dataset
              self.elem = 0
              self.length = length
      
          def __getitem__(self, index):
              self.elem += 1
              if self.elem > self.length:
                  self.elem = 0
                  raise StopIteration  # caught by DataLoader
              return self.dataset[index]
      
          def __len__(self):
              return len(self.dataset)
      
      
      if __name__ == '__main__':
          torch.manual_seed(0)
          dataloader = DataLoader(SubDataset(torch.arange(10), 5), shuffle=True)
          for _ in range(3):
              for x in dataloader:
                  print(x)
          print(len(dataloader))  # 10!!
      

      输出:

      请注意,将 __len__ 设置为 self.length 会导致问题,因为 dataloader 将仅使用 0 到 length-1 之间的索引(这不是您想要的)。不幸的是,如果没有这种行为(由于Dataloader 限制),我没有发现可以设置实际长度。因此要小心:len(dataset) 是原始长度,dataset.length 是新长度。

      【讨论】:

      • 这已经在torch.utils.data.Subset(Dataset) 中实现,不满足每个时期不同采样的要求
      • 在引发错误之前我忘记了self.elem = 0(请参阅编辑的代码)。现在我正在测试多个时期,并且数据集在每个时期都正确地重新洗牌
      猜你喜欢
      • 2013-06-23
      • 2015-01-31
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2017-03-21
      • 2011-06-27
      • 2018-10-29
      相关资源
      最近更新 更多