是否需要同时向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
相关产品推荐
相关产品推荐

