本文共 3313 字,大约阅读时间需要 11 分钟。
PyTorch中的Datasets和DataLoaders
在深入学习PyTorch的过程中,数据处理是不可忽视的重要环节。传统的数据处理代码常常会让人感到复杂且难以维护,这使得数据集代码与模型训练代码的耦合性较高。然而,PyTorch为我们提供了一套强大的数据处理工具——Datasets和DataLoaders,使得数据处理变得更加高效和灵活。
PyTorch中的Datasets类用于存储数据样本及其对应标签,而DataLoaders类则将数据集包装成一个可迭代的对象。Datasets类的主要作用是管理数据样本,DataLoaders则负责分批加载数据,适合在训练过程中高效处理大规模数据集。
Fashion-MNIST是由Zalando提供的一个常见数据集,包含60,000张训练图像和10,000张测试图像。每张图像尺寸为28×28,附带10种类别的标签。以下是使用PyTorch加载该数据集的示例代码:
import torchfrom torch.utils.data import Datasetfrom torchvision import datasetsfrom torchvision.transforms import ToTensor# 加载训练数据集training_data = datasets.FashionMNIST( root="data", train=True, download=True, transform=ToTensor())# 加载测试数据集test_data = datasets.FashionMNIST( root="data", train=False, download=True, transform=ToTensor())
通过索引操作,我们可以访问数据集中任意一个样本。以下代码展示了如何使用matplotlib可视化训练数据集中的部分样本:
import matplotlib.pyplot as plt# 定义标签映射labels_map = { 0: "T-Shirt", 1: "Trouser", 2: "Pullover", 3: "Dress", 4: "Coat", 5: "Sandal", 6: "Shirt", 7: "Sneaker", 8: "Bag", 9: "Ankle Boot",}# 设置绘图布局plt.figure(figsize=(8, 8))cols, rows = 3, 3# 遍历样本for i in range(1, cols * rows + 1): sample_idx = torch.randint(len(training_data), size=(1,)).item() img, label = training_data[sample_idx] plt.subplot(rows, cols, i) plt.title(labels_map[label]) plt.axis("off") plt.imshow(img.squeeze(), cmap="gray") 如果需要处理自定义数据集,需要创建一个继承自Dataset的自定义类。该类必须实现以下三个方法:
以下是一个简化的实现示例:
import osimport pandas as pdfrom torchvision.io import read_imageclass CustomImageDataset(Dataset): def __init__(self, annotations_file, img_dir, transform=None, target_transform=None): self.img_labels = pd.read_csv(annotations_file, names=['file_name', 'label']) self.img_dir = img_dir self.transform = transform self.target_transform = target_transform def __len__(self): return len(self.img_labels) def __getitem__(self, idx): img_path = os.path.join(self.img_dir, self.img_labels.iloc[idx, 0]) image = read_image(img_path) label = self.img_labels.iloc[idx, 1] if self.transform: image = self.transform(image) if self.target_transform: label = self.target_transform(label) return image, label
在模型训练阶段,我们通常希望通过DataLoader来批量加载数据集,这样可以提高数据加载的效率。此外,DataLoader还支持随机打乱数据集以减少模型过拟合。
from torch.utils.data import DataLoader# 创建训练数据集的DataLoader实例train_dataloader = DataLoader( training_data, batch_size=64, shuffle=True, num_workers=4)# 创建测试数据集的DataLoader实例test_dataloader = DataLoader( test_data, batch_size=64, shuffle=False, num_workers=4)
通过遍历DataLoader,可以逐批加载数据集中的样本。以下代码展示了如何从训练数据集中获取一批数据并进行可视化:
# 获取训练数据集的第一批样本train_features, train_labels = next(iter(train_dataloader))# 打印批量样本的形状print(f"Feature batch shape: {train_features.size()}")print(f"Labels batch shape: {train_labels.size()}")# 查看样本内容img = train_features[0].squeeze()label = train_labels[0]plt.imshow(img, cmap="gray")plt.show()print(f"Label: {label}") 通过以上代码,我们可以清晰地看到PyTorch如何通过Datasets和DataLoaders来简化数据处理流程,从而提高数据加载的效率和训练效果。
转载地址:http://vgxfk.baihongyu.com/