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

PyTorch中DataLoader、Sampler与Generator的关系及结果差异疑问

PyTorch中DataLoader、RandomSampler与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)

按理论三个场景结果应一致,但场景三结果不同,现解析三者的关系及原因:


核心逻辑梳理

  1. Sampler的独立性:手动创建RandomSampler实例时,它会立即初始化自身的generator(除非显式指定)。场景一和场景二中直接遍历ran_sampler,用的都是它自身初始化的Generator,所以输出固定。场景二中给DataLoader传的generator=G不会影响已初始化好的Sampler,因为Sampler的Generator在创建时就已确定,DataLoader不会覆盖已存在的Sampler内部Generator。

  2. 场景三的关键: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的结果不同。

  3. 场景二不受影响的原因:场景二中显式给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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 11:05:17