Skip to content

数据管道 Dataset 与 DataLoader

模型吃的不是整份数据集,而是一批一批(batch)的张量。Dataset 定义"单个样本长什么样",DataLoader 负责"打乱、组批、多线程搬运"。这篇讲清数据怎么从磁盘流进模型。

Dataset:一个"可按下标取样本"的协议

PyTorch 对数据集的要求极其朴素——实现两个东西即可:

python
import torch
from torch.utils.data import Dataset

class NumberDataset(Dataset):
    """演示用:0~99 的整数,标签是奇偶性(偶=0,奇=1)"""

    def __init__(self):
        self.x = torch.arange(100, dtype=torch.float32).unsqueeze(1) / 100  # [100, 1]
        self.y = (torch.arange(100) % 2).long()                             # [100]

    def __len__(self):                 # 有多少个样本
        return len(self.x)

    def __getitem__(self, idx):        # 第 idx 个样本长什么样
        return self.x[idx], self.y[idx]

约定就这两条:__len__ 返回样本数,__getitem__ 返回 (输入, 标签) 元组。几乎所有教程和源码项目都遵循这个协议。

现成的:torchvision 等内置数据集

MNIST、CIFAR-10 等经典数据集已内置,第一次运行会自动下载到 root 目录(MNIST 约 11MB):

python
from torchvision import datasets, transforms

transform = transforms.Compose([
    transforms.ToTensor(),          # PIL图像/数组 → [0,1] 的 float32 张量,必做!
    transforms.Normalize((0.1307,), (0.3081,)),  # 标准化为均值0方差1(MNIST统计量)
])

train_set = datasets.MNIST(root="./data", train=True,  download=True, transform=transform)
test_set  = datasets.MNIST(root="./data", train=False, download=True, transform=transform)

img, label = train_set[0]           # 像列表一样取第 0 个样本
print(img.shape, label)              # torch.Size([1, 28, 28]) tensor(5)

为什么要 Normalize? 神经网络对输入尺度敏感,均值 0、方差 1 的输入让梯度更稳定、收敛更快。不同数据集用各自的统计量(网上可查到 MNIST/CIFAR-10 的标准值),训练集和测试集必须用同一组参数。

DataLoader:把样本变成批次

python
from torch.utils.data import DataLoader

train_loader = DataLoader(train_set, batch_size=64, shuffle=True)
test_loader  = DataLoader(test_set,  batch_size=256, shuffle=False)  # 测试集不打乱

for images, labels in train_loader:      # 迭代即得批次,这就是训练循环里的一行
    print(images.shape, labels.shape)     # torch.Size([64, 1, 28, 28])  torch.Size([64])
    break

关键参数:

参数作用建议
batch_size每批样本数32/64/128/256,显存越大可以越大
shuffle每 epoch 重新打乱训练集 True,测试集 False
num_workers子进程预加载数据Linux 下 2~8 提速明显;Windows 下必须把主逻辑包进 if __name__ == "__main__":,否则进程无限递归报错
drop_last丢弃最后不足一批的数据BatchNorm 模型建议 True,避免 batch=1 时统计量出错

训练/验证/测试:为什么分三份

  • 训练集(train):更新参数用,模型"见过"。
  • 验证集(validation):每个 epoch 结束时评估,用来调超参数、早停——相当于"模拟考"。
  • 测试集(test):项目最终交付前只碰一次的"高考",提前反复看会导致调参过拟合。

没有现成划分时,用随机种子切分:

python
from torch.utils.data import random_split

full = datasets.MNIST(root="./data", train=True, download=True, transform=transform)
train_set, val_set = random_split(full, [55000, 5000],
                                  generator=torch.Generator().manual_seed(42))

注意 manual_seed(42):固定种子保证每次切分一致,实验可复现。"结果可复现"是深度学习工程的第一美德,凡调实验先设种子:torch.manual_seed(42)

自定义 Dataset 的实战模板

处理自己的数据(比如一堆图片文件)时,继承 Dataset 是标准做法:

python
from pathlib import Path
from PIL import Image

class ImageFolderDataset(Dataset):
    """假设目录结构: root/类别名/xxx.jpg"""

    def __init__(self, root, transform=None):
        self.paths = sorted(Path(root).rglob("*.jpg"))
        self.classes = sorted({p.parent.name for p in self.paths})
        self.class_to_idx = {c: i for i, c in enumerate(self.classes)}
        self.transform = transform

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

    def __getitem__(self, idx):
        path = self.paths[idx]
        img = Image.open(path).convert("RGB")
        label = self.class_to_idx[path.parent.name]
        if self.transform:
            img = self.transform(img)
        return img, label

坑预警:__getitem__ 里不要做重活(读大文件、解压),DataLoader 每个 batch 会调用它 batch_size 次,重活会拖慢整个训练。能预处理的一律提前预处理(如把大图先缩到模型输入尺寸存盘)。

小结

  • Dataset 协议:__len__ + __getitem__ 返回 (输入, 标签)
  • DataLoader 负责 shuffle / 组批 / 多进程;训练集 shuffle=True,测试集 False
  • 数据预处理记得固定随机种子,保证实验可复现。

概念篇到此结束。现在动手:去环境搭建把 PyTorch 装起来,然后跑最小可运行示例