Skip to content

PyTorch - 数据集

PyTorch - 处理数据集 (Datasets) 和数据加载器 (DataLoaders)

Section titled “PyTorch - 处理数据集 (Datasets) 和数据加载器 (DataLoaders)”

高效的数据处理对于训练机器学习模型至关重要。PyTorch 为此提供了两个主要的原语(Primitives):torch.utils.data.Dataset 和 torch.utils.data.DataLoader。这些工具类有助于抽象化访问和迭代数据的过程。

torchvision 库是 PyTorch 在计算机视觉任务方面的配套库,它在 torchvision.datasets 中内置了几个常用的数据集。

一个表示数据集的抽象类。你创建的任何自定义数据集都应该继承自 Dataset 并重写(override)两个方法:

  • __len__(self):应返回数据集中的总样本数。
  • __getitem__(self, index):应返回给定 index 处的样本(例如,图像及其标签)。

torchvision.datasets 为许多常用数据集提供了预构建的 Dataset 类。

一个迭代器,提供以下功能:

  • Batching (批量处理):将多个样本组合成一个批次(batch)。
  • Shuffling (打乱):在每个周期(epoch)随机打乱数据,以防止模型产生偏差(model bias)。
  • Parallel Loading (并行加载):使用多个子进程并行加载数据,加快训练速度。

DataLoader 将一个 Dataset 对象作为输入,使得迭代数据变得容易。

处理数据集时,尤其是在计算机视觉领域,你经常需要对数据应用变换(transforms)(例如,调整图像大小、将它们转换为张量(tensors)、归一化像素值)。torchvision.transforms 提供了常用的图像变换方法。

加载数据集时可以应用两种类型的变换:

  • transform:一个函数/变换,它接收一个样本(例如,PIL 图像)并返回一个变换后的版本。常用的变换包括 transforms.ToTensor()、transforms.Normalize()、transforms.Resize()。
  • target_transform:一个函数/变换,它接收目标(例如,标签)并返回一个变换后的版本。

这些变换可以使用 transforms.Compose() 组合在一起。

示例:使用 torchvision.datasets.MNIST

Section titled “示例:使用 torchvision.datasets.MNIST”

MNIST 是一个经典的包含手写数字的数据集。我们来看如何加载它。

import torch
import torchvision
import torchvision.transforms as transforms
from torch.utils.data import DataLoader
# Define a transform to convert images to tensors and normalize them
# Normalization values for MNIST are typically mean=0.1307, std=0.3081
# 定义一个变换:将图像转换为张量并进行归一化
# MNIST 的典型归一化值:均值=0.1307,标准差=0.3081
transform_mnist = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,))
])
# Download and load the training data
# root: directory where data is/will be stored
# train=True: get the training set
# download=True: download if not already present
# transform: apply the defined transformations
# 下载并加载训练数据
# root: 数据存储的目录
# train=True: 获取训练集
# download=True: 如果不存在则下载
# transform: 应用定义的变换
train_dataset_mnist = torchvision.datasets.MNIST(
root='./data',
train=True,
transform=transform_mnist,
download=True
)
# Download and load the test data
# 下载并加载测试数据
test_dataset_mnist = torchvision.datasets.MNIST(
root='./data',
train=False,
transform=transform_mnist,
download=True
)
# Create DataLoaders
# batch_size: number of samples per batch
# shuffle=True: shuffle training data to ensure randomness
# 创建数据加载器 (DataLoaders)
# batch_size: 每批次的样本数
# shuffle=True: 打乱训练数据以确保随机性
train_loader_mnist = DataLoader(dataset=train_dataset_mnist, batch_size=64, shuffle=True)
test_loader_mnist = DataLoader(dataset=test_dataset_mnist, batch_size=1000, shuffle=False)
print(f"Number of training samples: {len(train_dataset_mnist)}")
print(f"Number of test samples: {len(test_dataset_mnist)}")
# Example of iterating through the DataLoader
# 示例:遍历数据加载器
data_iter = iter(train_loader_mnist)
images, labels = next(data_iter)
print(f"Shape of a batch of images: {images.shape}") # Expected: torch.Size([64, 1, 28, 28])
print(f"Shape of a batch of labels: {labels.shape}") # Expected: torch.Size([64])

torchvision.datasets 类(如 MNIST)的参数:

  • root (str): 数据集 MNIST/processed/training.pt 和 MNIST/processed/test.pt 存在的或将要保存的根目录。
  • train (bool, optional): 如果为 True,则从 training.pt 创建数据集,否则从 test.pt 创建。
  • download (bool, optional): 如果为 True,则从互联网下载数据集并将其放置在 root 目录中。如果数据集已下载,则不再下载。
  • transform (callable, optional): 一个函数/变换,它接收一个 PIL 图像并返回一个变换后的版本。
  • target_transform (callable, optional): 一个函数/变换,它接收目标并对其进行变换。

示例:使用 torchvision.datasets.CIFAR10

Section titled “示例:使用 torchvision.datasets.CIFAR10”

CIFAR10 是另一个流行的数据集,包含 10 个类别的 32x32 彩色图像。

# Define transforms for CIFAR10 (3-channel images)
# Normalization values for CIFAR10 are typically mean=[0.4914, 0.4822, 0.4465], std=[0.2023, 0.1994, 0.2010]
# 定义 CIFAR10(3通道图像)的变换
# CIFAR10 的典型归一化值:均值=[0.4914, 0.4822, 0.4465],标准差=[0.2023, 0.1994, 0.2010]
transform_cifar10 = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010))
])
train_dataset_cifar10 = torchvision.datasets.CIFAR10(
root='./data',
train=True,
transform=transform_cifar10,
download=True
)
train_loader_cifar10 = DataLoader(dataset=train_dataset_cifar10, batch_size=64, shuffle=True)
print(f"\nNumber of CIFAR10 training samples: {len(train_dataset_cifar10)}")
data_iter_cifar = iter(train_loader_cifar10)
images_cifar, labels_cifar = next(data_iter_cifar)
print(f"Shape of a CIFAR10 batch of images: {images_cifar.shape}") # Expected: torch.Size([64, 3, 32, 32])

其他值得注意的 torchvision.datasets

Section titled “其他值得注意的 torchvision.datasets”

torchvision 还提供了许多其他数据集,包括:

  • ImageNet:一个大型图像数据集(由于大小和许可原因,需要手动下载)。
  • COCO (Common Objects in Context):用于目标检测、分割和图像描述(captioning)。例如:
# COCO Example (requires COCO API and data to be downloaded separately)
# try:
# coco_cap = torchvision.datasets.CocoCaptions(
# root = 'path/to/coco/images/train2017', # Directory with images
# annFile = 'path/to/coco/annotations/captions_train2017.json', # Annotation JSON file
# transform=transforms.ToTensor()
# )
# print(f'Number of COCO caption samples: {len(coco_cap)}')
# if len(coco_cap) > 0:
# img, target = coco_cap[0] # Get first sample (image and its captions)
# print(f'COCO Image shape: {img.shape}, Number of captions: {len(target)}')
# except Exception as e:
# print(f"Could not load COCO dataset. Ensure data and API are set up. Error: {e}")
# COCO 示例(需要单独下载 COCO API 和数据)
# try:
# coco_cap = torchvision.datasets.CocoCaptions(
# root = 'path/to/coco/images/train2017', # 包含图像的目录
# annFile = 'path/to/coco/annotations/captions_train2017.json', # 标注 JSON 文件
# transform=transforms.ToTensor()
# )
# print(f'COCO 图像描述样本数: {len(coco_cap)}')
# if len(coco_cap) > 0:
# img, target = coco_cap[0] # 获取第一个样本(图像及其图像描述)
# print(f'COCO 图像形状: {img.shape}, 图像描述数量: {len(target)}')
# except Exception as e:
# print(f"无法加载 COCO 数据集。请确保数据和 API 已设置。错误: {e}")

如果数据可用,COCO 的输出可能如下所示:

Number of COCO caption samples: 118287 (Example number for train2017 captions)
COCO Image shape: torch.Size([3, height, width]), Number of captions: 5 (Example number)
COCO 图像描述样本数: 118287 (train2017 图像描述的示例数量)
COCO 图像形状: torch.Size([3, height, width]), 图像描述数量: 5 (示例数量)

对于大多数实际项目,你需要创建自己的自定义 Dataset 类。这包括:

  1. 创建一个继承自 torch.utils.data.Dataset 的类。
  2. 实现 __init__(self, ...):初始化数据集,例如,加载文件路径、标注(annotations)。
  3. 实现 __len__(self):返回总样本数。
  4. 实现 __getitem__(self, idx):加载并返回给定索引 idx 的一个样本(例如,图像和标签)。这通常是应用变换(transforms)的地方。

欲了解更多详细信息,请参阅 PyTorch 官方关于 Dataset 和 DataLoader 的文档,并探索关于创建自定义数据集的教程。

额外资源: