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(10, 3)
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 |
|
|
batch_size |
|
|
shuffle |
|
|
num_workers |
|
|
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. 支持的数据集(部分列举)
|
|
|
|---|---|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
★各数据集详细说明与接口参见官方文档:https://pytorch.org/vision/stable/datasets.html
六、MNIST 数据集详解
1. 简介
MNIST 是一个著名的手写数字数据集,在深度学习领域是经典的学习入门样例,是 NIST 数据集的一个子集。[pdf_6]
2. 数据集组成
|
|
|
|
|---|---|---|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
3. 读取 MNIST 数据集
import torchvision
mnist_dataset = torchvision.datasets.MNIST(root='./data',
train=True,
transform=None,
target_transform=None,
download=True)
4. 构造函数参数说明
|
|
|
|---|---|
root |
|
download |
|
train |
|
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 |
|
|
| Torchvision 内置数据集 |
|
|
★
torchvision.datasets继承了 Dataset 类,在预定义许多常用数据集的同时,还预留了数据预处理与数据增强的接口。[pdf_6]
八、每课一练
问题:在 PyTorch 中定义一个数据集,应继承哪个类?
解答:应继承 Dataset 类(torch.utils.data.Dataset),并重写 __init__()、__len__()、__getitem__() 三个方法。
九、小结
-
数据读取机制:Dataset 类 + DataLoader 类 → 数据迭代器 -
Dataset 类:抽象类,自定义数据集需继承并重写 __init__、__len__、__getitem__ -
DataLoader 类:迭代器,支持 batch 加载、多进程、数据打乱 -
Torchvision 库:包含常用数据集 + 常见网络模型 + 图像处理方法 -
MNIST 数据集:经典手写数字数据集,6万训练+1万测试,接口封装了下载、解压、读取、解析全过程 -
核心原则: torchvision.datasets是 Dataset 的派生类,预留了 transform 接口,方便数据预处理与增强

