Skip to content

PyTorch - 加载数据

高效的数据加载对于训练深度学习模型至关重要。PyTorch 通过其 torch.utils.data 模块为此提供了强大而灵活的工具,主要包括 Dataset 和 DataLoader 类。torchvision 包还为计算机视觉任务提供了方便的预置数据集和图像转换实用程序(image transformation utilities)。

torch.utils.data.Dataset 是一个抽象类,代表一个数据集(dataset)。你的自定义数据集应该继承自 Dataset 并重写两个关键方法:

  • __len__(self):使得 len(dataset) 返回数据集的大小。
  • __getitem__(self, idx):支持索引,例如可以使用 dataset[i] 获取第 i 个样本。

这是一个自定义 Dataset 的基本框架:

import torch
from torch.utils.data import Dataset
class MyCustomDataset(Dataset):
def __init__(self, data_source, targets, transform=None):
self.data = data_source # 例如,文件路径列表,或预加载的数据
self.targets = targets
self.transform = transform
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
sample_data = self.data[idx] # 加载或访问你的数据点
sample_target = self.targets[idx]
# 如果有转换(transform),则应用
if self.transform:
sample_data = self.transform(sample_data)
return sample_data, sample_target

torchvision.datasets 提供了几个预加载的数据集,如 CIFAR10、MNIST 等,它们都是 torch.utils.data.Dataset 的子类。

例如,使用 CIFAR10 数据集:

import torchvision
import torchvision.transforms as transforms
# 定义转换(transformations):将 PIL Image 转换为 Tensor,然后进行归一化
transform_pipeline = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize(mean=(0.4914, 0.4822, 0.4465), std=(0.2023, 0.1994, 0.2010)) # CIFAR10 三通道的均值/标准差
])
# 加载 CIFAR10 训练集
trainset = torchvision.datasets.CIFAR10(root='./data',
train=True,
download=True,
transform=transform_pipeline)
# 加载 CIFAR10 测试集
testset = torchvision.datasets.CIFAR10(root='./data',
train=False,
download=True,
transform=transform_pipeline)

transforms.ToTensor() 将 PIL Image 或 NumPy ndarray(形状为 H x W x C,范围在 [0, 255])转换为 torch.FloatTensor(形状为 C x H x W,范围在 [0.0, 1.0])。transforms.Normalize() 使用给定的每个通道的均值和标准差对 Tensor 图像进行归一化。

torch.utils.data.DataLoader 是一个迭代器,它为数据加载提供了许多基本功能:

  • 批处理数据(Batching the data)。
  • 在每个 Epoch(轮次)打乱(Shuffling)数据。
  • 使用多进程 worker 并行加载数据。
import torch
# 假设 trainset 是一个 Dataset 的实例(例如上面创建的 CIFAR10 训练集)
trainloader = torch.utils.data.DataLoader(trainset,
batch_size=64, # 每批次的样本数量
shuffle=True, # 在每个 epoch 打乱数据
num_workers=2) # 用于数据加载的子进程(worker)数量
# 测试集类似,通常不打乱
testloader = torch.utils.data.DataLoader(testset,
batch_size=100,
shuffle=False,
num_workers=2)

batch_size 决定了在更新模型权重之前处理多少样本。shuffle=True 对于训练很重要,可以防止模型学习样本的顺序。num_workers > 0 启用多进程数据加载,这可以在 GPU 忙于计算时,在 CPU 上并行准备数据,从而显著加快训练速度。通常的做法是先从 num_workers=0 开始(以便于调试,因为数据加载在主进程中进行),然后根据系统性能和特定工作负载,增加 num_workers 的值,例如设置为 os.cpu_count(),以找到最优值。

示例:用于 CSV 文件的自定义 Dataset

Section titled “示例:用于 CSV 文件的自定义 Dataset”

通常,你的数据可能存储在 CSV 文件中。你可以使用 Pandas 等库来读取 CSV 文件,然后将其包装在自定义的 Dataset 中。

考虑一个名为 my_data.csv 的 CSV 文件,其中包含特征和标签:

my_data.csv 文件内容示例: feature1,feature2,label 1.5,2.3,0 4.1,3.9,1 0.8,1.1,0

import pandas as pd
import torch
from torch.utils.data import Dataset
class CSVTabularDataset(Dataset):
def __init__(self, csv_path, transform=None):
self.dataframe = pd.read_csv(csv_path)
# 假设最后一列是目标/标签,其余是特征
self.features = self.dataframe.iloc[:, :-1].values.astype('float32')
self.labels = self.dataframe.iloc[:, -1].values.astype('int64') # 标签通常使用 int64,例如用于 CrossEntropyLoss
self.transform = transform
def __len__(self):
return len(self.features)
def __getitem__(self, idx):
feature_sample = torch.tensor(self.features[idx], dtype=torch.float32)
label_sample = torch.tensor(self.labels[idx], dtype=torch.long)
if self.transform:
# 示例:transform 可以是一个用于特征的归一化函数
feature_sample = self.transform(feature_sample)
return feature_sample, label_sample
# 使用方法(假设 'my_data.csv' 存在或已创建):
# # 创建一个用于示例的 my_data.csv 文件:
# # with open('my_data.csv', 'w') as f:
# # f.write('feature1,feature2,label\n1.5,2.3,0\n4.1,3.9,1\n0.8,1.1,0')
#
# tabular_dataset = CSVTabularDataset(csv_path='my_data.csv')
# tabular_loader = torch.utils.data.DataLoader(tabular_dataset, batch_size=2, shuffle=True)
#
# # 遍历 DataLoader
# # for features, labels in tabular_loader:
# # print(f'批量特征形状: {features.shape}, 批量标签: {labels}')
# # # 示例输出: Batch Features Shape: torch.Size([2, 2]), Batch Labels: tensor([0, 1])
# # break

处理复杂的 CSV 数据(例如,图像关键点)

Section titled “处理复杂的 CSV 数据(例如,图像关键点)”

如果你的 CSV 文件包含更复杂的数据,例如图像文件的路径和相关的关键点坐标(landmark coordinates),你的 __getitem__ 方法将涉及更多步骤,比如从路径加载图像和解析坐标字符串。例如,更新原始教程中的关键点处理逻辑:

import pandas as pd
import numpy as np # 或者使用 torch 进行形状重塑/Tensor 转换
import torch
# from PIL import Image # 用于加载图像
# 此代码段展示了可放入自定义 Dataset 的 __getitem__ 方法中的逻辑。
# 假设 'self.dataframe' 是一个从 'faces/face_landmarks.csv' 等 CSV 文件加载的 pandas DataFrame。
# 每行可能包含:image_filename, landmark1_x, landmark1_y, landmark2_x, ...
# 并且 'idx' 是传递给 __getitem__ 的索引。
# 在自定义 Dataset 的 __getitem__(self, idx) 方法内部:
# current_row = self.dataframe.iloc[idx]
# img_path = current_row.iloc[0] # 第一个元素,例如 'path/to/image_001.jpg'
# landmark_values_str = current_row.iloc[1:].values # 其余是关键点坐标字符串
# 将关键点坐标字符串(从 CSV 读取)转换为数值数组并重塑
# landmarks_numeric = landmark_values_str.astype('float32').reshape(-1, 2) # 假设是二维关键点 (x, y)
# landmarks_tensor = torch.tensor(landmarks_numeric, dtype=torch.float32)
# 然后通常会使用 img_path 加载图像:
# image = Image.open(img_path).convert('RGB') # 转换为 RGB 以保持一致性
# 对图像和关键点应用必要的转换(如果适用)。
# if self.transform:
# image, landmarks_tensor = self.transform(image, landmarks_tensor) # 转换可能需要同时处理图像和关键点
# return image, landmarks_tensor
# 注意:原始代码段:
# landmarks_frame = pd.read_csv('faces/face_landmarks.csv')
# n = 65
# img_name = landmarks_frame.iloc[n, 0]
# landmarks = landmarks_frame.iloc[n, 1:].values # .values 替代了 .as_matrix()
# landmarks = landmarks.astype('float').reshape(-1, 2)
# 这段逻辑用于从整个 DataFrame 中提取特定行 'n' 的数据。
# 在 Dataset 中,你需要将其改编为用于当前 'idx' 的逻辑。

torchvision.transforms 提供了一系列常用的图像转换。这些转换可以使用 transforms.Compose 串联起来。一些常用的转换包括:

  • transforms.ToTensor():将 PIL Image 或 NumPy ndarray 转换为 torch.Tensor 并将值缩放到 [0,1] 范围。

  • transforms.ToPILImage():将 torch.Tensor 或 ndarray 转换为 PIL Image。

  • transforms.Normalize(mean, std):使用每个通道的均值和标准差对 Tensor 图像进行归一化。

  • transforms.Resize(size):将输入图像的大小调整为给定尺寸。

  • transforms.CenterCrop(size):从给定图像的中心裁剪指定尺寸。

  • transforms.RandomResizedCrop(size):将图像裁剪为随机大小和宽高比,然后调整到目标尺寸。

  • transforms.RandomHorizontalFlip(p=0.5):以概率 p(默认为 0.5)随机水平翻转给定图像。

  • 数据增强转换,如 transforms.ColorJitter、transforms.RandomRotation、transforms.RandomAffine 等,通过向模型提供更多样化的训练数据,有助于提高模型的鲁棒性(robustness)。

适当的数据增强和预处理是训练鲁棒模型的关键。更多详情请参阅 torchvision.transforms 文档。

  • 调试 DataLoader:从 num_workers=0 开始。这会在主进程中加载数据,使调试更容易(例如,堆栈跟踪更清晰)。确认无误后,增加 num_workers 以提高性能。

  • Batch Size(批量大小):选择适合你的 GPU 内存并能提供稳定梯度的 batch_size。由于硬件优化,通常使用 2 的幂次方(例如 32、64、128、256),但最优值需要根据经验确定。

  • 可复现性(Reproducibility):为了获得可复现的结果,尤其是在使用 shuffle=True 或随机转换时,需要为 PyTorch (torch.manual_seed())、NumPy (np.random.seed()) 和 Python 的 random 模块设置随机种子。对于 num_workers > 0 的 DataLoader,如果 worker 独立使用随机性,你可能还需要自定义 worker_init_fn 来确保 worker 的种子也设置正确。

  • 大型数据集:对于太大而无法完全加载到内存中的数据集,确保你的 Dataset 在 __getitem__ 中动态加载数据(例如,从磁盘读取图像)。对于真正海量的数据集或数据流式传输场景,考虑使用 torch.utils.data.IterableDataset,它适用于无法进行随机访问或数据动态生成的情况。像 WebDataset 这样的库对于大规模分布式训练也很有帮助。

  • Collate Function(整理函数):DataLoader 使用一个 collate_fn 来将单个样本(由 Dataset.__getitem__ 返回)组合成一个批次。默认的 collate_fn 对于许多标准数据类型(张量、数字、字符串)工作良好。然而,对于复杂的数据结构,如可变长度的序列或自定义对象,你可能需要提供自定义的 collate_fn 来指定样本应如何打包在一起(例如,将序列填充到相同长度)。

在图像分类任务中,你的文件夹结构可能每个子文件夹代表一个类别(例如,/data/cats/、/data/dogs/)。你的自定义 Dataset 将会:

  1. 扫描这些文件夹以找到所有图像路径并分配相应的标签(例如,猫为 0,狗为 1)。
  2. 在 __getitem__ 中,给定一个索引,从其路径加载图像文件(例如,使用 PIL.Image.open())。
  3. 应用 torchvision.transforms:调整到标准的输入尺寸,可能执行随机增强(如翻转或旋转),通过 transforms.ToTensor() 将图像转换为 torch.Tensor,然后使用预先计算的数据集的均值和标准差,通过 transforms.Normalize() 进行归一化。

然后,DataLoader 接收这个数据集,打乱数据(用于训练),并创建 (image_tensor_batch, label_tensor_batch) 的批次。这些批次被依次馈送到你的神经网络进行训练。这种高效的数据管线对于防止 GPU 空闲并最大化训练吞吐量至关重要。