PyTorch中DataLoader(shuffle=True)未正确打乱DBpedia数据集?
问题原因及解决方案
问题根源
torchtext提供的DBpedia数据集属于IterableDataset类型,这类数据集通过迭代器顺序输出数据,默认按类别排序返回。而DataLoader的shuffle=True仅对支持随机访问的普通Dataset生效,对IterableDataset不起作用,因此你看到的批次仍保持原数据集的类别顺序,导致单批次标签种类极少。
解决方法
方法一:转换为普通Dataset
将IterableDataset转为列表形式的普通Dataset,让shuffle=True正常生效:
import torchtext.datasets as d from torch.utils.data import DataLoader, Dataset # 将IterableDataset转为列表 train_data = list(d.DBpedia(split="train", root=".cache")) # 定义列表类Dataset class ListDataset(Dataset): def __init__(self, data): self.data = data def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx] # 创建支持shuffle的DataLoader train_loader = DataLoader( ListDataset(train_data), batch_size=10000, shuffle=True, ) # 验证标签分布 for labels, texts in train_loader: print(len(set(labels.tolist())))
方法二:使用torchdata的Shuffler变换(推荐)
针对IterableDataset,使用torchdata的Shuffler进行数据打乱,这是torchtext新版本推荐的处理方式:
from torchdata.datapipes.iter import Shuffler import torchtext.datasets as d from torch.utils.data import DataLoader # 获取数据集并添加shuffle操作 train_dp = d.DBpedia(split="train", root=".cache") train_dp = Shuffler(train_dp, buffer_size=100000) # buffer_size越大,打乱效果越充分 train_loader = DataLoader( train_dp, batch_size=10000, ) # 验证标签分布 for labels, texts in train_loader: print(len(set(labels.tolist())))
注意事项
- 方法一需将整个数据集加载到内存,DBpedia训练集共56万条数据,需确保内存足够。
- 方法二中
Shuffler的buffer_size参数决定打乱程度,建议设置为远大于batch_size的值,平衡打乱效果与内存占用。
内容的提问来源于stack exchange,提问作者fuji
相关产品推荐
相关产品推荐

