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

是否需要同时向DataLoader与RandomSampler传入数据集?

问题修正与解释

首先,你的代码里有一个关键错误:在DataLoader中同时设置shuffle=True和sampler是冲突的。PyTorch规定,当指定了sampler参数时,shuffle参数会被直接忽略,而且这种写法可能触发警告,必须把shuffle设为False(或者不写,默认就是False)。

两种采样器的用法说明

  • RandomSampler需要传入数据集:它的作用是生成数据集索引的随机排列,所以必须通过传入数据集获取总长度,进而生成0到len(dataset)-1的随机索引序列。
  • WeightedRandomSampler不需要传入数据集:它基于你提供的权重列表采样,采样范围由权重列表的长度决定。这里必须保证权重列表的长度和数据集长度完全一致,否则会出现索引越界或样本不匹配的问题。

修正后的代码

def train_dataloader(self):
    if self._is_weighted_sampler:
        weights = list(self.label_weight_by_name.values())
        # 确保权重长度与数据集长度匹配,避免采样错误
        assert len(weights) == len(self._train_dataset), "权重列表长度必须与数据集长度匹配"
        sampler = torch.utils.data.sampler.WeightedRandomSampler(
            torch.tensor(weights), len(self._train_dataset)
        )
    else:
        sampler = torch.utils.data.RandomSampler(self._train_dataset)
    # 使用sampler时,shuffle必须设为False
    return DataLoader(self._train_dataset, batch_size=self._batch_size, shuffle=False, sampler=sampler)

关于“数据集被传入两次”的误解

RandomSampler接收数据集只是为了获取它的长度,并没有复制或重复加载数据集;DataLoader传入数据集是为了根据采样器生成的索引获取对应样本。两者职责完全不同,不存在“重复传入”的问题,你的误解主要来自对采样器作用和DataLoader参数优先级的不熟悉。

内容的提问来源于stack exchange,提问作者Gulzar

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 06:06:08