PyTorch中DataLoader、Sampler与Generator的关系及结果差异疑问
假设存在数据集:
datasets = [0,1,2,3,4]
场景一
代码:
torch.manual_seed(1) ran_sampler = RandomSampler(data_source=datasets) for data in ran_sampler: print(data)
输出结果:1,3,4,0,2
场景二
代码:
torch.manual_seed(1) seed=1234 G = torch.Generator() G.manual_seed(seed) ran_sampler = RandomSampler(data_source=datasets) dataloader = DataLoader(dataset=datasets, sampler=ran_sampler, generator=G) for data in ran_sampler: print(data)
输出结果仍为1,3,4,0,2,且无论seed取何值,结果不变。
场景三
代码:
torch.manual_seed(1) ran_sampler = RandomSampler(data_source=datasets) dataloader = DataLoader(dataset=datasets, sampler=ran_sampler) for data in dataloader: print(data)
输出结果为4,1,3,0,2
查看RandomSampler源码可知,未指定generator时,Sampler会自行生成随机种子并创建Generator:
seed = int(torch.empty((), dtype=torch.int64).random_().item()) generator = torch.Generator() generator.manual_seed(seed)
按理论三个场景结果应一致,但场景三结果不同,现解析三者的关系及原因:
核心逻辑梳理
Sampler的独立性:手动创建RandomSampler实例时,它会立即初始化自身的
generator(除非显式指定)。场景一和场景二中直接遍历ran_sampler,用的都是它自身初始化的Generator,所以输出固定。场景二中给DataLoader传的generator=G不会影响已初始化好的Sampler,因为Sampler的Generator在创建时就已确定,DataLoader不会覆盖已存在的Sampler内部Generator。场景三的关键:DataLoader对Sampler的隐式修改:当手动创建的RandomSampler传入DataLoader且未指定DataLoader的generator时,DataLoader会执行特殊操作:复用自身默认创建的Generator,替换掉Sampler原有的Generator。
具体来说,DataLoader初始化时,若发现传入了自定义Sampler且未指定generator,会生成一个默认Generator(该Generator受全局种子
torch.manual_seed(1)影响),然后将这个Generator赋值给Sampler的generator属性。这就导致Sampler原本的随机顺序被改变,所以场景三遍历DataLoader时,用的是被替换后的Generator生成的顺序,和直接遍历Sampler的结果不同。场景二不受影响的原因:场景二中显式给DataLoader指定了
generator=G,此时Sampler已有自身的Generator,DataLoader不会修改已初始化Sampler的Generator——只有当DataLoader使用自身默认Generator时,才会触发这个替换逻辑。
验证方式
可在场景三中添加代码,打印DataLoader创建前后Sampler的Generator种子:
torch.manual_seed(1) ran_sampler = RandomSampler(data_source=datasets) print("创建DataLoader前Sampler的Generator种子:", ran_sampler.generator.initial_seed()) dataloader = DataLoader(dataset=datasets, sampler=ran_sampler) print("创建DataLoader后Sampler的Generator种子:", ran_sampler.generator.initial_seed())
会发现两个种子值不同,证明DataLoader确实替换了Sampler的Generator。
内容的提问来源于stack exchange,提问作者liaoming999

