PyTorch中多数据集差异化子采样(按Epoch重采样)拼接实现问询
针对动态子采样需求的两种实现方案
方案1:基于IterableDataset的自动Epoch重采样
直接实现一个包装类继承IterableDataset,利用其__iter__方法在每个Epoch被调用时自动重新采样,完美契合你的需求:
import torch from torch.utils.data import IterableDataset, ConcatDataset, Dataset class SubsampledIterableDataset(IterableDataset): def __init__(self, base_dataset, target_sample_count): self.base_dataset = base_dataset self.target_size = target_sample_count # 与小数据集规模对齐 def __iter__(self): # 每个Epoch随机生成采样索引 sample_indices = torch.randperm(len(self.base_dataset))[:self.target_size].tolist() for idx in sample_indices: yield self.base_dataset[idx] # 实际使用 small_ds = YourSmallDataset() large_ds = YourLargeDataset() # 将大数据集包装为动态采样版本 subsampled_large = SubsampledIterableDataset(large_ds, len(small_ds)) # 拼接两个数据集 combined_ds = ConcatDataset([subsampled_large, small_ds]) # 传入DataLoader train_loader = torch.utils.data.DataLoader(combined_ds, batch_size=32)
每次DataLoader启动一个新Epoch时,都会触发__iter__方法重新采样,无需额外手动操作。
方案2:Map-style Dataset的Epoch级采样更新
Map-style Dataset没有内置的Epoch钩子,但可以通过自定义包装类+手动触发更新的方式实现:
import torch from torch.utils.data import Dataset, ConcatDataset class EpochSubsampledMapDataset(Dataset): def __init__(self, base_dataset, target_sample_count): self.base_dataset = base_dataset self.target_size = target_sample_count self.current_sample_indices = [] def refresh_samples(self): # 每个Epoch开始前调用,更新采样索引 self.current_sample_indices = torch.randperm(len(self.base_dataset))[:self.target_size].tolist() def __len__(self): return self.target_size def __getitem__(self, idx): return self.base_dataset[self.current_sample_indices[idx]] # 实际使用 small_ds = YourSmallDataset() large_ds = YourLargeDataset() subsampled_large = EpochSubsampledMapDataset(large_ds, len(small_ds)) combined_ds = ConcatDataset([subsampled_large, small_ds]) train_loader = torch.utils.data.DataLoader(combined_ds, batch_size=32) # 训练循环中手动触发更新 for epoch in range(10): subsampled_large.refresh_samples() for batch in train_loader: # 训练逻辑 pass
这种方案保留了Map-style Dataset的随机访问特性,适合需要依赖该特性的场景,代价是每个Epoch开始前需要手动调用刷新方法。
补充说明
- 如果需要采样的可复现性,可以在采样时固定随机种子(比如
torch.manual_seed(epoch + 固定值)) - 两种方案都能满足你"拼接后传入DataLoader"的要求,无需修改DataLoader的sampler参数
内容的提问来源于stack exchange,提问作者Adam
相关产品推荐
相关产品推荐

