PyTorch中ConcatDataset如何实现不同数据集的非均匀采样?
实现方案
PyTorch 中数据集的采样逻辑由Sampler组件控制,你不需要修改ConcatDataset的定义,只需要自定义符合采样规则的采样器传入 DataLoader 即可,具体实现如下:
1. 严格固定循环规律的实现
如果你需要严格按照「CustomDataset1、CustomDataset1、CustomDataset2」的顺序循环采样,用自定义采样器实现:
import torch from torch.utils.data import Sampler, ConcatDataset from typing import List class FixedRatioCycleSampler(Sampler[int]): def __init__(self, concat_dataset: ConcatDataset, ratios: List[int], shuffle_subset: bool = True): """ 按指定比例循环从ConcatDataset的子数据集采样 :param concat_dataset: 拼接后的ConcatDataset实例 :param ratios: 各子数据集的采样比例,你的需求对应传[2, 1] :param shuffle_subset: 子数据集内部是否打乱采样顺序 """ self.sub_datasets = concat_dataset.datasets self.sub_cum_sizes = concat_dataset.cumulative_sizes self.ratios = ratios self.shuffle_subset = shuffle_subset # 各子数据集在ConcatDataset中的全局索引偏移量 self.offset = [0] + self.sub_cum_sizes[:-1].tolist() # 计算单epoch总采样数,这里取最长子数据集对齐,你也可以自定义为固定数值 max_sub_len = max([len(ds) for ds in self.sub_datasets]) self.total_samples = sum(ratios) * (max_sub_len // min(ratios) + 1) def __iter__(self): # 生成每个子数据集的采样索引池,长度不够时自动循环采样 sub_index_pools = [] for ds_idx in range(len(self.sub_datasets)): # 生成子数据集内部索引 inner_indices = torch.arange(len(self.sub_datasets[ds_idx])) if self.shuffle_subset: inner_indices = inner_indices[torch.randperm(len(inner_indices))] # 循环填充到足够长度 repeat = (self.total_samples // self.ratios[ds_idx]) + 1 inner_indices = inner_indices.repeat(repeat)[:self.total_samples // self.ratios[ds_idx]] # 加上偏移量得到全局索引 inner_indices += self.offset[ds_idx] sub_index_pools.append(inner_indices.tolist()) # 按固定比例拼接索引,实现[ds1, ds1, ds2]循环 final_indices = [] ptrs = [0] * len(self.sub_datasets) while len(final_indices) < self.total_samples: for ds_idx in range(len(self.sub_datasets)): for _ in range(self.ratios[ds_idx]): if len(final_indices) >= self.total_samples: break final_indices.append(sub_index_pools[ds_idx][ptrs[ds_idx]]) ptrs[ds_idx] += 1 if len(final_indices) >= self.total_samples: break return iter(final_indices) def __len__(self): return self.total_samples
2. 自定义采样器的使用方法
创建 DataLoader 时传入自定义采样器即可:
# 原有数据集拼接逻辑 dataset = ConcatDataset([CustomDataset1, CustomDataset2]) # 初始化采样器,设置2:1的采样比例 sampler = FixedRatioCycleSampler(concat_dataset=dataset, ratios=[2, 1], shuffle_subset=True) # 传入DataLoader,注意设置sampler后不要重复设置shuffle参数 # batch_size设为3的话,每个batch刚好对应[ds1, ds1, ds2]的采样规律 dataloader = torch.utils.data.DataLoader(dataset, batch_size=3, sampler=sampler)
可选简化方案(仅保证整体比例,不要求固定循环顺序)
如果你不需要严格的固定循环顺序,只需要整体采样比例趋近于2:1,可以直接用PyTorch内置的WeightedRandomSampler,实现更简单:
dataset = ConcatDataset([CustomDataset1, CustomDataset2]) len1 = len(CustomDataset1) len2 = len(CustomDataset2) # 给所有样本设置权重:CustomDataset1每个样本权重为2,CustomDataset2每个为1 weights = [2.0] * len1 + [1.0] * len2 # num_samples可以自定义为你需要的单epoch采样总数 sampler = torch.utils.data.WeightedRandomSampler(weights, num_samples=len(weights), replacement=True) dataloader = torch.utils.data.DataLoader(dataset, batch_size=3, sampler=sampler)
内容的提问来源于stack exchange,提问作者moon
相关产品推荐
相关产品推荐

