You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何让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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.19 09:22:05