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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 01:52:52