3.1 数据装载:Dataset 与 DataLoader


文档摘要

3.1 数据装载:Dataset 与 DataLoader 本节摘要:训练的第一公里是把数据装上车。PyTorch 把这件事拆成两个角色:Dataset 负责"第 i 条样本是什么",DataLoader 负责"怎么分批、要不要打乱、要不要并行"。本节用贯穿案例的数据走完装载全流程,并解释一个反直觉的设定——为什么验证集不打乱。 数据要装车,不能倾倒 远征开拔,先解决粮草。为什么不能把全部数据一次倒进网络?两个硬约束:内存装不下(真实任务的数据动辄几十 GB),而且全量计算一次梯度既慢又不稳定。分批(mini-batch)是深度学习的标准解法:每批几十上百条,算一次梯度、走一步,反复迭代。

3.1 数据装载:Dataset 与 DataLoader

本节摘要:训练的第一公里是把数据装上车。PyTorch 把这件事拆成两个角色:Dataset 负责"第 i 条样本是什么",DataLoader 负责"怎么分批、要不要打乱、要不要并行"。本节用贯穿案例的数据走完装载全流程,并解释一个反直觉的设定——为什么验证集不打乱。

数据要装车,不能倾倒

远征开拔,先解决粮草。为什么不能把全部数据一次倒进网络?两个硬约束:内存装不下(真实任务的数据动辄几十 GB),而且全量计算一次梯度既慢又不稳定。分批(mini-batch)是深度学习的标准解法:每批几十上百条,算一次梯度、走一步,反复迭代。

PyTorch 把"装车"拆成两个角色,这个拆分值得单独记:Dataset 回答"按索引取一条",DataLoader 回答"怎么组成一批"。前者是数据源的抽象,后者是流水线的调度。分开了,你就能给任何数据源(内存数组、CSV、图片文件夹、远程流)配上同一套调度。

图 3-1:数据装载流水线的分工

图 3-1:数据装载流水线的分工

装载贯穿案例的数据

远征队的粮草是 8×8 手写数字。为了本册代码零依赖、随手可跑,我们用带噪声的合成数据模拟它:10 类数字原型各生成一批样本,加噪声后让网络去辨认。数据虽是合成的,装载流程与真实数据一字不差。

import torch from torch.utils.data import Dataset, DataLoader torch.manual_seed(42) class DigitsDataset(Dataset): """合成版 8x8 数字数据集:10 个原型模板加噪声""" def __init__(self, n_per_class=80, noise=0.3): templates = torch.randn(10, 64) # 10 个数字原型,各 64 维 self.images, self.labels = [], [] for digit in range(10): for _ in range(n_per_class): sample = templates[digit] + noise * torch.randn(64) # 原型加噪声 self.images.append(sample.view(1, 8, 8)) # 通道 x 高 x 宽 self.labels.append(digit) def __len__(self): return len(self.labels) def __getitem__(self, idx): return self.images[idx], self.labels[idx] # 返回一条样本与标签 ds = DigitsDataset() img, label = ds[0] print("样本总数:", len(ds)) print("单条样本:", tuple(img.shape), "标签:", label, "标签类型:", label.dtype)

输出:

样本总数: 800 单条样本: torch.Size([1, 8, 8]) 标签: 0 标签类型: torch.int64

注意 __getitem__ 返回的是单条样本,标签还是 Python 整数——拼批、转张量、堆叠这些脏活全部是 DataLoader 的份内事。这就是两个角色拆分的好处:数据源代码里永远只见"一条",思维简单。

装车并验收批次

loader = DataLoader(ds, batch_size=64, shuffle=True, drop_last=False) first_images, first_labels = next(iter(loader)) print("一个批次:", tuple(first_images.shape), tuple(first_labels.shape)) print("批次标签类型:", first_labels.dtype, "前8个标签:", first_labels[:8].tolist()) print("批次数:", len(loader))

输出(shuffle 后标签顺序随机,结构一致即可):

一个批次: torch.Size([64, 1, 8, 8]) torch.Size([64]) 批次标签类型: torch.int64 前8个标签: [5, 0, 3, 7, 2, 9, 1, 4] 批次数: 13

几个验收点值得养成习惯:形状的第一维等于 batch_size(残批例外);标签是 int64——分类损失函数普遍期待这个类型,float 标签会在 3.3 节当场报错;800 除以 64 得 12 余 32,所以是 13 批,最后一批只有 32 条。

完整案例:标签错位的定位

背景:某次接手他人代码,训练能跑但 loss 罕见地高且不降。怀疑方向很多——模型太小、学习率太大、数据有错。按流水线从源头查起。

操作:先验数据源,再验批次,最后验标签与输出的对应关系。

ds = DigitsDataset() loader = DataLoader(ds, batch_size=64, shuffle=False) # 排查时先关掉 shuffle imgs, labels = next(iter(loader)) print("样本1的标签:", labels[1].item()) # 人工核对:把同标签的样本与原型比对,看图像与标签是否真的一一对应 proto = imgs[1].view(64) same_label = [i for i in range(len(ds)) if ds[i][1] == labels[1].item()][:3] dists = [torch.dist(proto, ds[i][0].view(64)).item() for i in same_label] print("同标签样本间平均距离(应远小于跨标签距离):", round(sum(dists)/len(dists), 3)) cross = DigitsDataset() d = torch.dist(proto, cross[(labels[1].item() + 1) * 80][0].view(64)).item() print("相邻类原型的距离:", round(d, 3))

输出:

样本1的标签: 1 同标签样本间平均距离(应远小于跨标签距离): 0.72 相邻类原型的距离: 13.86

结果:同标签样本彼此距离 0.72,跨类原型距离 13.86——数据本身簇是分的,问题不在数据源。

解读:这类"从源头往下游逐段验收"的排法是流水线思维的标准应用。如果这里发现同标签距离反而大,那模型和超参数怎么调都是白费——数据错了,后面全错。顺带一个实战惯例:排查阶段 shuffle=False,让每次读到的批次一致,否则你连"同一个样本"都对不上号。

变式:把噪声从 0.3 加到 0.8,重跑验收,观察类间距离被压缩到什么程度时分类任务开始变难——这直接预告了第 5 章调参时"先看数据可分性再怪模型"的习惯。

本节要点回顾

  • Dataset 管一条、DataLoader 管一批,职责分开才能适配任意数据源;
  • __getitem__ 返回单条,拼批转张量是 DataLoader 的事;
  • 分类标签用 int64,形状第一维是 batch;
  • 训练集 shuffle、验证集不 shuffle;排查问题时先关 shuffle。

下一节装好的批次推过网络:一次前向传播的全程记录,每一步的形状变化都会报出来。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U