大数跨境

Torchvision(上):数据读取,训练开始的第一步

Torchvision(上):数据读取,训练开始的第一步 知识代码AI
2026-08-22
2

Torchvision(上):数据读取,训练开始的第一步


一、PyTorch中的数据读取机制

PyTorch 提供了一套方便的数据读取机制,使用 Dataset 类与 DataLoader 类的组合来得到数据迭代器,在训练或预测时输出每一批次所需的数据,并对数据进行相应的预处理与数据增强操作。[pdf_6]


二、Dataset 类

1. 基本概念

  • Dataset 是一个抽象类,用来表示数据集
  • 通过继承 Dataset 类来自定义数据集的格式、大小和其它属性,供 DataLoader 类直接使用
  • 无论自定义数据集还是官方封装好的数据集,本质都是继承了 Dataset 类[pdf_6]

2. 必须重写的方法

方法
说明
__init__()
构造函数,可自定义数据读取方法及进行数据预处理
__len__()
返回数据集大小
__getitem__()
索引数据集中的某一个数据

3. 自定义 Dataset 示例

import torch
from torch.utils.data import Dataset

classMyDataset(Dataset):
# 构造函数
def__init__(self, data_tensor, target_tensor):
        self.data_tensor = data_tensor
        self.target_tensor = target_tensor
# 返回数据集大小
def__len__(self):
return self.data_tensor.size(0)
# 返回索引的数据与标签
def__getitem__(self, index):
return self.data_tensor[index], self.target_tensor[index]

调用示例:

# 生成数据
data_tensor = torch.randn(103)
target_tensor = torch.randint(2, (10,))

# 将数据封装成Dataset
my_dataset = MyDataset(data_tensor, target_tensor)

# 查看数据集大小
print('Dataset size:', len(my_dataset))      # 输出:Dataset size: 10

# 使用索引调用数据
print('tensor_data[0]: ', my_dataset[0])
# 输出: tensor_data[0]: (tensor([0.4931, -0.0697, 0.4171]), tensor(0))

三、DataLoader 类

1. 设计目的

实际项目中数据量很大,受限于内存、I/O 速度等问题,无法一次性将所有数据加载到内存,也不能只用一个进程加载,因此需要多进程、迭代加载

2. 基本概念

  • DataLoader 是一个迭代器
  • 最基本的使用方法:传入一个 Dataset 对象,根据 batch_size 参数的值生成一个 batch 的数据
  • 可节省内存,同时支持多进程数据打乱等处理[pdf_6]

3. DataLoader 参数说明

参数
类型
说明
dataset
Dataset
输入的数据集,必须参数
batch_size
int
每个 batch 有多少个样本
shuffle
bool
在每个 epoch 开始的时候,是否对数据进行重新打乱
num_workers
int
加载数据的进程数,0 意味着所有数据都会被加载进主进程,默认为 0

4. 代码示例

from torch.utils.data import DataLoader

tensor_dataloader = DataLoader(dataset=my_dataset,  # 传入的数据集, 必须参数
                               batch_size=2,        # 输出的batch大小
                               shuffle=True,        # 数据是否打乱
                               num_workers=0)       # 进程数, 0表示只有主进程

# 以循环形式输出
for data, target in tensor_dataloader:
    print(data, target)

# 输出一个batch
print('One batch tensor data: ', iter(tensor_dataloader).next())

四、Torchvision 库

1. 什么是 Torchvision

  • Torchvision 是一个和 PyTorch 配合使用的 Python 包,包含很多图像处理工具
  • 三大组成部分
    • ✅ 常用数据集
    • ✅ 常见网络模型
    • ✅ 常用图像处理方法

2. 安装方式

# conda 安装
conda install torchvision -c pytorch

# pip 安装
pip install torchvision

3. 依赖库 Pillow

  • Torchvision 默认使用的图像加载器是 PIL,因此需要安装 Pillow
  • 提供广泛的文件格式支持,功能包括图像储存、图像显示、格式转换以及基本的图像处理操作
# 安装 Pillow
conda install pillow
pip install pillow

五、torchvision.datasets 包

1. 基本概念

  • torchvision.datasets 包提供了丰富的图像数据集接口
  • 常用的图像数据集(如 MNIST、COCO 等)均有封装
  • 重要提醒:该包本身不包含数据集的文件本身,工作方式是先从网络把数据集下载到用户指定目录,再用加载器将数据集加载到内存中,最后将加载后的数据集作为对象返回给用户[pdf_6]

2. 支持的数据集(部分列举)

分类
数据集
手写数字
MNIST、EMNIST、KMNIST、QMNIST
时尚物品
Fashion-MNIST
街景门牌
SVHN
自然图像
CIFAR、STL10、Places365、ImageNet
目标检测
COCO、VOC
人脸
CelebA
场景文字
SEMEION
视频行为
HMDB51、UCF101、Kinetics-400

各数据集详细说明与接口参见官方文档:https://pytorch.org/vision/stable/datasets.html


六、MNIST 数据集详解

1. 简介

MNIST 是一个著名的手写数字数据集,在深度学习领域是经典的学习入门样例,是 NIST 数据集的一个子集。[pdf_6]

2. 数据集组成

内容
文件名
大小
训练集图片
train-images-idx3-ubyte.gz
9.9MB→47MB,6万个样本
训练集标签
train-labels-idx1-ubyte.gz
29KB→60KB,6万个标签
测试集图片
t10k-images-idx3-ubyte.gz
1.6MB→7.8MB,1万个样本
测试集标签
t10k-labels-idx1-ubyte.gz
5KB→10KB,1万个标签

3. 读取 MNIST 数据集

import torchvision

mnist_dataset = torchvision.datasets.MNIST(root='./data',
                                           train=True,
                                           transform=None,
                                           target_transform=None,
                                           download=True)

4. 构造函数参数说明

参数
说明
root
指定保存 MNIST 数据集的位置;download=False 时从该位置读取
download
是否下载数据集;为 True 时自动下载;已存在文件则不会重复下载
train
True 加载训练集,False 加载测试集
transform
对图像进行预处理操作,如数据增强、归一化、旋转或缩放等
target_transform
对图像标签进行预处理操作

5. 数据类型

  • mnist_dataset 的类型是 torchvision.datasets.mnist.MNIST
  • 该类是 Dataset 类的派生类——torchvision.datasets 已帮我们写好了对 Dataset 类的继承,直接使用即可
  • 其他数据集使用方式类似:只需将类名换成其它数据集名字即可
  • 对于没有官方接口的图像数据集,可以使用 torchvision.datasets.ImageFolder 接口来自行定义[pdf_6]

6. 数据预览

mnist_dataset_list = list(mnist_dataset)
print(mnist_dataset_list)

转换后的数据集对象变成了一个元组列表,每个元组有两个元素:

  • 第一个元素:图像数据(PIL.Image.Image 类型)
  • 第二个元素:图像的标签
display(mnist_dataset_list[0][0])
print("Image label is:", mnist_dataset_list[0][1])
# 示例结果:第一条数据是手写数字"7",对应标签是"7"

七、两种读取数据的方法对比

方法
适用场景
说明
自定义 Dataset
任何自定义数据集
继承 Dataset 类,重写三个方法,最通用的方法
Torchvision 内置数据集
常用图像数据集
直接实例化即可,如 MNIST、CIFAR、COCO 等

torchvision.datasets 继承了 Dataset 类,在预定义许多常用数据集的同时,还预留了数据预处理与数据增强的接口。[pdf_6]


八、每课一练

问题:在 PyTorch 中定义一个数据集,应继承哪个类?

解答:应继承 Dataset 类torch.utils.data.Dataset),并重写 __init__()__len__()__getitem__() 三个方法。


九、小结

  1. 数据读取机制:Dataset 类 + DataLoader 类 → 数据迭代器
  2. Dataset 类:抽象类,自定义数据集需继承并重写 __init____len____getitem__
  3. DataLoader 类:迭代器,支持 batch 加载、多进程、数据打乱
  4. Torchvision 库:包含常用数据集 + 常见网络模型 + 图像处理方法
  5. MNIST 数据集:经典手写数字数据集,6万训练+1万测试,接口封装了下载、解压、读取、解析全过程
  6. 核心原则torchvision.datasets 是 Dataset 的派生类,预留了 transform 接口,方便数据预处理与增强


【声明】内容源于网络
0
0
知识代码AI
技术基底 机器视觉全栈 × 光学成像 × 图像处理算法 编程栈 C++/C#工业开发 | Python智能建模 工具链 Halcon/VisionPro工业部署 | PyTorch/TensorFlow模型炼金术 | 模型压缩&嵌入式移植
内容 403
粉丝 0
知识代码AI 技术基底 机器视觉全栈 × 光学成像 × 图像处理算法 编程栈 C++/C#工业开发 | Python智能建模 工具链 Halcon/VisionPro工业部署 | PyTorch/TensorFlow模型炼金术 | 模型压缩&嵌入式移植
总阅读6.6k
粉丝0
内容403