PyTorch - 数据集
PyTorch - 处理数据集 (Datasets) 和数据加载器 (DataLoaders)
Section titled “PyTorch - 处理数据集 (Datasets) 和数据加载器 (DataLoaders)”高效的数据处理对于训练机器学习模型至关重要。PyTorch 为此提供了两个主要的原语(Primitives):torch.utils.data.Dataset 和 torch.utils.data.DataLoader。这些工具类有助于抽象化访问和迭代数据的过程。
torchvision 库是 PyTorch 在计算机视觉任务方面的配套库,它在 torchvision.datasets 中内置了几个常用的数据集。
核心概念:Dataset 和 DataLoader
Section titled “核心概念:Dataset 和 DataLoader”torch.utils.data.Dataset
Section titled “torch.utils.data.Dataset”一个表示数据集的抽象类。你创建的任何自定义数据集都应该继承自 Dataset 并重写(override)两个方法:
__len__(self):应返回数据集中的总样本数。__getitem__(self, index):应返回给定index处的样本(例如,图像及其标签)。
torchvision.datasets 为许多常用数据集提供了预构建的 Dataset 类。
torch.utils.data.DataLoader
Section titled “torch.utils.data.DataLoader”一个迭代器,提供以下功能:
- Batching (批量处理):将多个样本组合成一个批次(batch)。
- Shuffling (打乱):在每个周期(epoch)随机打乱数据,以防止模型产生偏差(model bias)。
- Parallel Loading (并行加载):使用多个子进程并行加载数据,加快训练速度。
DataLoader 将一个 Dataset 对象作为输入,使得迭代数据变得容易。
变换 (Transforms)
Section titled “变换 (Transforms)”处理数据集时,尤其是在计算机视觉领域,你经常需要对数据应用变换(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 torchimport torchvisionimport torchvision.transforms as transformsfrom 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.3081transform_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 (示例数量)创建自定义数据集
Section titled “创建自定义数据集”对于大多数实际项目,你需要创建自己的自定义 Dataset 类。这包括:
- 创建一个继承自
torch.utils.data.Dataset的类。 - 实现
__init__(self, ...):初始化数据集,例如,加载文件路径、标注(annotations)。 - 实现
__len__(self):返回总样本数。 - 实现
__getitem__(self, idx):加载并返回给定索引idx的一个样本(例如,图像和标签)。这通常是应用变换(transforms)的地方。
欲了解更多详细信息,请参阅 PyTorch 官方关于 Dataset 和 DataLoader 的文档,并探索关于创建自定义数据集的教程。
额外资源: