如何让PyTorch DataLoader迭代顺序与random.shuffle种子序列一致?
问题:让PyTorch DataLoader迭代顺序与random.shuffle序列完全对应
我尝试通过设置随机种子1,让PyTorch DataLoader按照特定序列加载数据,以下是我的代码:
import random import torch.utils.data.dataset as Dataset import torch.utils.data.dataloader as DataLoader from torch.utils.data.sampler import Sampler class MyDataset(Dataset.Dataset): def __init__(self): self.Data = [x for x in range(10)] self.Label = [x for x in range(10)] def __getitem__(self, index): data = self.Data[index] label = self.Label[index] return data, label def __len__(self): return len(self.Data) class RandSeqSampler(Sampler): def __init__(self, data_source): super().__init__(data_source) self.data_source = data_source def __iter__(self): indices = list(range(len(self.data_source))) random.shuffle(indices) return iter(indices) def __len__(self): return len(self.data_source) random.seed(1) dataset = MyDataset() dataloader = DataLoader.DataLoader(dataset=dataset, batch_size=1, sampler=RandSeqSampler(dataset)) for i, (data, label) in enumerate(dataloader): print(data, label) print("\n\n\n\n\n") for i, (data, label) in enumerate(dataloader): print(data, label) random.seed(1) a = [x for x in range(10)] random.shuffle(a) print(a) random.shuffle(a) print(a)
输出结果为:
tensor([6]) tensor([6]) tensor([8]) tensor([8]) tensor([9]) tensor([9]) tensor([7]) tensor([7]) tensor([5]) tensor([5]) tensor([3]) tensor([3]) tensor([0]) tensor([0]) tensor([4]) tensor([4]) tensor([1]) tensor([1]) tensor([2]) tensor([2]) tensor([4]) tensor([4]) tensor([8]) tensor([8]) tensor([2]) tensor([2]) tensor([6]) tensor([6]) tensor([5]) tensor([5]) tensor([9]) tensor([9]) tensor([0]) tensor([0]) tensor([7]) tensor([7]) tensor([1]) tensor([1]) tensor([3]) tensor([3]) [6, 8, 9, 7, 5, 3, 0, 4, 1, 2] [5, 1, 9, 0, 3, 2, 6, 4, 8, 7]
可以看到,第一次迭代时DataLoader的加载顺序与random.shuffle的结果一致,但第二次迭代的加载顺序与random.shuffle的第二次结果不一致。我希望DataLoader的加载顺序能与random.shuffle的序列完全对应,请问该如何实现?
解决方案
问题原因
你的RandSeqSampler每次迭代时都会调用random.shuffle,但第一次遍历DataLoader后,全局随机种子的状态已经被第一次shuffle操作消耗并改变了。而你单独测试random.shuffle时是重新设置了种子1,所以第二次DataLoader的shuffle是基于第一次shuffle后的种子状态,自然和重新设种子后的shuffle结果无法匹配。
方案1:每次遍历前重置随机种子
如果希望每次遍历DataLoader都对应random.seed(1)后的第一次shuffle结果,只需在每次遍历前重新设置随机种子:
random.seed(1) for i, (data, label) in enumerate(dataloader): print(data, label) print("\n\n\n\n\n") random.seed(1) # 重置种子到初始状态 for i, (data, label) in enumerate(dataloader): print(data, label)
如果要对应第二次shuffle的结果,可以在第二次遍历前先执行一次无意义的shuffle来消耗种子状态,但这种方式不够直观。
方案2:预定义所有需要的序列(更可控)
预先生成所有需要的随机序列,让Sampler按顺序返回这些序列,完全不受全局种子状态影响:
import random import torch.utils.data.dataset as Dataset import torch.utils.data.dataloader as DataLoader from torch.utils.data.sampler import Sampler class MyDataset(Dataset.Dataset): def __init__(self): self.Data = [x for x in range(10)] self.Label = [x for x in range(10)] def __getitem__(self, index): data = self.Data[index] label = self.Label[index] return data, label def __len__(self): return len(self.Data) class PredefinedSeqSampler(Sampler): def __init__(self, sequences): super().__init__(None) self.sequences = sequences self.current_idx = 0 def __iter__(self): # 返回当前序列,迭代后切换到下一个(循环使用或按需停止) current_seq = self.sequences[self.current_idx] self.current_idx = (self.current_idx + 1) % len(self.sequences) return iter(current_seq) def __len__(self): return len(self.sequences[0]) # 预先生成匹配需求的序列 random.seed(1) seq1 = list(range(10)) random.shuffle(seq1) seq2 = list(range(10)) random.shuffle(seq2) dataset = MyDataset() # 传入预生成的序列列表 dataloader = DataLoader.DataLoader(dataset=dataset, batch_size=1, sampler=PredefinedSeqSampler([seq1, seq2])) print("第一次遍历:") for i, (data, label) in enumerate(dataloader): print(data, label) print("\n\n第二次遍历:") for i, (data, label) in enumerate(dataloader): print(data, label) print("\n预生成的目标序列:") print(seq1) print(seq2)
这种方式可以完全控制每次DataLoader迭代的顺序,确保和你预先生成的random.shuffle结果完全一致。
内容的提问来源于stack exchange,提问作者WillWu
相关产品推荐
相关产品推荐

