适合:需要处理自定义数据集、想规范数据加载流程的 PyTorch 新手 核心内容:加载官方数据集、自定义 Dataset 类、DataLoader 批量加载数据
一、前言
处理数据样本的代码往往杂乱且难以维护;理想情况下,我们希望数据集代码与模型训练代码解耦,以提高可读性和模块化程度。PyTorch 提供了两个数据核心组件:torch.utils.data.DataLoader 和 torch.utils.data.Dataset,让你既能使用预加载的数据集,也能处理自定义数据。
- Dataset:存储样本及其对应的标签;
- DataLoader:围绕 Dataset 封装可迭代对象,方便访问样本。
PyTorch 领域库(如 TorchText、TorchVision、TorchAudio)提供了大量预加载的数据集,这些数据集均继承自 torch.utils.data.Dataset,并实现了特定于数据类型的函数,可用于模型原型开发和基准测试。你可以在这里找到它们:
- 图像数据集:Image Datasets
- 文本数据集:Text Datasets
- 音频数据集:Audio Datasets
二、加载官方数据集(以 FashionMNIST 为例)
2.1 数据集介绍(原文翻译)
FashionMNIST 是 Zalando 服装图片数据集,包含 60000 个训练样本和 10000 个测试样本。每个样本是 28×28 的灰度图像,对应 10 个类别中的一个标签。
2.2 加载代码(附详细注释)
加载时的核心参数:
- root:训练/测试数据的存储路径;
- train:指定是训练集还是测试集;
- download=True:如果 root 路径下无数据,则从互联网下载;
- transform/target_transform:分别指定对特征和标签的变换。
import torch
from torch.utils.data import Dataset
from torchvision import datasets
from torchvision.transforms import ToTensor
import matplotlib.pyplot as plt
# 加载训练集
training_data = datasets.FashionMNIST(
root=\”data\”, # 数据存储根目录
train=True, # 训练集
download=True, # 自动下载缺失数据
transform=ToTensor() # 将图像转为张量
)
# 加载测试集
test_data = datasets.FashionMNIST(
root=\”data\”,
train=False, # 测试集
download=True,
transform=ToTensor()
)
注:运行代码时会显示数据下载进度,如下(原文进度条翻译):
0%| | 0.00/26.4M [00:00<?, ?B/s]
1%| | 164k/26.4M [00:00<00:55, 469kB/s]
31%|███▏ | 8.29M/26.4M [00:00<00:01, 14.5MB/s]
100%|██████████| 26.4M/26.4M [00:01<00:00, 18.2MB/s]
2.3 遍历并可视化数据集(原文翻译+代码)
我们可以像列表一样手动索引 Dataset(如 training_data[index]),并使用 matplotlib 可视化训练集中的样本:
# 标签映射(数字→类别名称)
labels_map = {
0: \”T-Shirt\”, # T恤
1: \”Trouser\”, # 裤子
2: \”Pullover\”, # 套头衫
3: \”Dress\”,



