使用SubsetRandomSampler后平行数据集无法同步加载的问题
平行数据集同步加载问题解决方法
我有两个平行数据集dataset1和dataset2,尝试用SubsetRandomSampler传入train_indices实现同步加载。即使设置了num_workers=0,也给numpy和torch设了随机种子,样本还是没法同步加载。以下是我的代码、实际输出和期望输出:
原代码
import torch, numpy as np from torch.utils.data import Dataset, DataLoader, SubsetRandomSampler dataset1 = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9]) dataset2 = torch.tensor([10, 11, 12, 13, 14, 15, 16, 17, 18, 19]) train_indices = list(range(len(dataset1))) torch.manual_seed(12) np.random.seed(12) np.random.shuffle(train_indices) sampler = SubsetRandomSampler(train_indices) dataloader1 = DataLoader(dataset1, batch_size=2, num_workers=0, sampler=sampler) dataloader2 = DataLoader(dataset2, batch_size=2, num_workers=0, sampler=sampler) for i, (data1, data2) in enumerate(zip(dataloader1, dataloader2)): x = data1 y = data2 print(x, y)
实际输出
tensor([5, 1]) tensor([15, 18]) tensor([0, 2]) tensor([14, 12]) tensor([4, 6]) tensor([16, 10]) tensor([8, 9]) tensor([11, 19]) tensor([7, 3]) tensor([17, 13])
期望输出
tensor([5, 1]) tensor([15, 11]) tensor([0, 2]) tensor([10, 12]) tensor([4, 6]) tensor([14, 16]) tensor([8, 9]) tensor([18, 19]) tensor([7, 3]) tensor([17, 13])
问题原因
SubsetRandomSampler的__iter__方法每次被调用时,都会对传入的indices重新生成随机排列的迭代器。两个DataLoader各自调用sampler的迭代器,会得到不同的采样序列,导致样本无法同步。
解决方法
方法一:使用SequentialSampler固定采样序列
预先生成shuffle后的固定索引序列,用SequentialSampler包装,确保两个DataLoader使用完全相同的采样顺序:
import torch, numpy as np from torch.utils.data import DataLoader, SequentialSampler dataset1 = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9]) dataset2 = torch.tensor([10, 11, 12, 13, 14, 15, 16, 17, 18, 19]) # 生成固定的shuffle索引 train_indices = list(range(len(dataset1))) torch.manual_seed(12) np.random.seed(12) np.random.shuffle(train_indices) # 使用SequentialSampler按固定序列采样 sampler = SequentialSampler(train_indices) dataloader1 = DataLoader(dataset1, batch_size=2, num_workers=0, sampler=sampler) dataloader2 = DataLoader(dataset2, batch_size=2, num_workers=0, sampler=sampler) for data1, data2 in zip(dataloader1, dataloader2): print(data1, data2)
方法二:自定义联合数据集(推荐)
将两个平行数据集包装成一个自定义Dataset,让DataLoader每次同时返回对应样本,从根源避免同步问题:
import torch, numpy as np from torch.utils.data import Dataset, DataLoader class ParallelDataset(Dataset): def __init__(self, data1, data2): self.data1 = data1 self.data2 = data2 assert len(data1) == len(data2), "两个数据集长度必须一致" def __len__(self): return len(self.data1) def __getitem__(self, idx): return self.data1[idx], self.data2[idx] dataset1 = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9]) dataset2 = torch.tensor([10, 11, 12, 13, 14, 15, 16, 17, 18, 19]) parallel_dataset = ParallelDataset(dataset1, dataset2) torch.manual_seed(12) np.random.seed(12) dataloader = DataLoader(parallel_dataset, batch_size=2, num_workers=0, shuffle=True) for data1, data2 in dataloader: print(data1, data2)
方法三:手动控制采样索引
直接预先生成分好batch的索引序列,手动按索引取数,完全控制同步逻辑:
import torch, numpy as np dataset1 = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9]) dataset2 = torch.tensor([10, 11, 12, 13, 14, 15, 16, 17, 18, 19]) train_indices = list(range(len(dataset1))) torch.manual_seed(12) np.random.seed(12) np.random.shuffle(train_indices) batch_size = 2 # 按batch分割索引 batches = [train_indices[i:i+batch_size] for i in range(0, len(train_indices), batch_size)] for batch_idx in batches: data1 = dataset1[batch_idx] data2 = dataset2[batch_idx] print(data1, data2)
内容的提问来源于stack exchange,提问作者kkgarg
相关产品推荐
相关产品推荐

