欢迎光临
我们一直在努力

PyTorch3 PyTorch Dataset & DataLoader 保姆级教程|自定义数据集必看

适合:需要处理自定义数据集、想规范数据加载流程的 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\”,

赞(0)
未经允许不得转载:171主机测评 » PyTorch3 PyTorch Dataset & DataLoader 保姆级教程|自定义数据集必看
分享到: 更多 (0)

评论 抢沙发

  • 昵称 (必填)
  • 邮箱 (必填)
  • 网址