Appearance
数据管道 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。 - 数据预处理记得固定随机种子,保证实验可复现。