首页 / PyTorch 入门教程 / Dataset 与 DataLoader:数据加载管线

PyTorch 入门教程

Dataset 与 DataLoader:数据加载管线

本教程共 60 篇 · 第 11 篇 · 更新于 2026-08-17 · 约 3 分钟阅读

PyTorchDatasetDataLoader数据加载批处理TensorDataset

本节目标:搞懂 DatasetDataLoader 的分工,学会自定义数据集,理解批处理、打乱、多进程这几个关键参数。

先讲分工

训练时数据不能一股脑全塞进模型。原因有两个:数据太多,显存装不下;顺序太固定,模型容易「背答案」而不是学规律。于是 PyTorch 把数据环节拆成两个角色:

  • 数据集(Dataset):负责「我有什么」,回答两个问题:一共多少样本?第 i 个样本长什么样?
  • 数据加载器(DataLoader):负责「怎么喂」,把样本攒成一批批(batch),按需打乱、多进程加速。

可以把 Dataset 想成后厨的冰箱:食材都按编号存放。DataLoader 是配菜员:一次抓几份、抓之前摇匀、忙不过来多雇几个人。两者各管一摊,互不越界。

自定义 Dataset:两个方法

自定义数据集只需继承 torch.utils.data.Dataset,实现两个方法:

import torch
from torch.utils.data import Dataset

class MyDataset(Dataset):
    def __init__(self, data, labels):
        self.data = data
        self.labels = labels

    def __len__(self):
        return len(self.data)

    def __getitem__(self, idx):
        return self.data[idx], self.labels[idx]

data = torch.randn(100, 5)
labels = torch.randint(0, 2, (100,))
dataset = MyDataset(data, labels)

print(len(dataset))          # 100
print(dataset[0][0].shape)   # torch.Size([5])

写完 __len____getitem__,数据集就「活」了:len(dataset)dataset[3]for 循环全部可用。真实场景里,__getitem__ 常常是「现取现读」——比如从磁盘读一张图片。好处是内存友好,几万张图不用全部驻留内存。

Tip

数据本身就是张量的话别自己造轮子,TensorDataset 就是干这个的:TensorDataset(x, y) 一行搞定上面的类。

DataLoader:配菜员上岗

把数据集交给 DataLoader,设置几个关键参数:

from torch.utils.data import DataLoader

loader = DataLoader(dataset, batch_size=16, shuffle=True)
for batch_idx, (x, y) in enumerate(loader):
    print(batch_idx, x.shape, y.shape)
    # 0 torch.Size([16, 5]) torch.Size([16])
    if batch_idx == 1:
        break

参数逐个说:

  1. batch_size:一批多少个样本。太小训练慢;太大会吃显存甚至报 OOM。32 到 128 是常见起步区间。
  2. shuffle:每轮(epoch)是否打乱顺序。训练集建议 True,顺序太固定模型容易记住答案;验证集不需要,纯属浪费时间。
  3. num_workers:几个子进程并行读数据。数据预处理耗时(比如图像解码)时开 2-4 个明显提速;纯内存张量数据开 0 反而更快,进程间通信也有成本。
  4. drop_last:最后一批不满 batch_size 时是否丢弃。有些网络结构对批大小敏感,此时常设 True。
  5. pin_memory:用 GPU 训练时设 True,数据会放在锁页内存里,CPU 到 GPU 的拷贝快一截。

len(loader) 返回批次数,等于样本总数除以 batch_size 向上取整。想手动取一批数据,用 iter(loader) 拿到迭代器再 next()

data_iter = iter(loader)
x, y = next(data_iter)   # 拿第一批

内置数据集:不用自己造

常见数据集 torchvision 都帮你打包好了。MNIST、CIFAR-10 这些,一行代码下载并加载:

from torchvision import datasets, transforms

train_dataset = datasets.MNIST(
    root="./data", train=True,
    transform=transforms.ToTensor(),
    download=True,
)

它返回的就是标准的 Dataset 对象,可以直接丢给 DataLoader。图像分类还有一种常用数据集 ImageFolder:把图片按类别放文件夹,root/ants/xxx.jpgroot/bees/yyy.jpg,它自动把子文件夹名当标签,省去手写 CSV 解析。

多个数据集想合并?用 ConcatDataset([ds1, ds2]),拼接后当成一个数据集用。

三个常见坑

第一个坑:样本形状不统一。图片有的 300×300 有的 200×200,DataLoader 拼批次时会报错。解法是预处理统一尺寸,这正是下一章 transforms 的主场。

第二个坑:多进程卡死。Windows 下 num_workers 大于 0 时,数据加载代码必须放在 if __name__ == "__main__": 保护块里,否则子进程会递归启动、无限循环。我当年在这儿卡过一下午,教训深刻。

第三个坑:以为 shuffle=True 是「打乱一次就固定了」。其实每轮 epoch 都会重新洗牌,这是特性不是 bug。

进阶一点:collate_fn 与 getitems

DataLoader 默认把一批样本按「堆叠」方式合并:每个样本返回元组 (x, y),它就堆出一批 x 和一批 y。如果样本结构特殊(比如变长序列),默认合并会失败,这时传 collate_fn 自定义合并逻辑:

def my_collate(batch):
    xs = [item[0] for item in batch]
    ys = [item[1] for item in batch]
    return torch.nn.utils.rnn.pad_sequence(xs, batch_first=True), torch.tensor(ys)

loader = DataLoader(dataset, batch_size=8, collate_fn=my_collate)

collate_fn 拿到一个 batch 的样本列表,怎么拼你说了算。长度不等的序列可以用 pad_sequence 补成一样长,再堆成张量。

另一个优化点:数据集里可以实现 __getitems__(self, indices),让 DataLoader 一次取整批,而不是逐个调用 __getitem__。适合「一次查询取多行」的场景,比如批量读数据库,省下反复建连接的开销。官方测试里,它能把加载速度提升好几倍。

小结

  • Dataset 管「有什么」:__len__ + __getitem__ 是全部要求。
  • DataLoader 管「怎么喂」:batch_size、shuffle、num_workers、drop_last、pin_memory。
  • 常见数据集用 torchvision 内置的,多个数据集用 ConcatDataset 合并。
  • 形状不统一、Windows 多进程卡死、shuffle 每轮重洗,是新手三大坑。
  • 特殊批处理用 collate_fn,批量取数可以试试 __getitems__

下一章讲 transforms:数据集读出来的原始数据,怎么加工成模型爱吃的格式。