如何使用 PyTorch 加载 CIFAR-100 数据集
使用 PyTorch 加载 CIFAR-100 数据集
本教程将演示如何使用 PyTorch 的 torchvision 库加载 CIFAR-100 数据集。CIFAR-100 是一个用于图像分类的常用数据集,包含 100 个类别,每个类别有 600 张图像。
我们将使用 torchvision.datasets.CIFAR100 类来加载数据集。此类允许您轻松下载和加载 CIFAR-100 数据集,并返回可用于构建训练和验证数据加载器的数据集对象。
以下是加载 CIFAR-100 数据集的 Python 代码示例:
import torchvision
import torchvision.transforms as transforms
import torch
# 定义数据变换
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])
# 加载训练数据集
train_dataset = torchvision.datasets.CIFAR100(root='./data', train=True, download=True, transform=transform)
# 加载验证数据集
val_dataset = torchvision.datasets.CIFAR100(root='./data', train=False, download=True, transform=transform)
# 创建训练数据加载器
train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=32, shuffle=True)
# 创建验证数据加载器
val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=32, shuffle=False)
代码解释:
- 首先,我们导入必要的库:
torchvision用于加载数据集,torchvision.transforms用于数据预处理,torch用于创建数据加载器。 - 然后,我们定义数据变换。数据变换是应用于数据集的预处理步骤。在这里,我们使用
transforms.Compose将两个变换链接在一起:transforms.ToTensor()将图像从 PIL 图像转换为 PyTorch 张量。transforms.Normalize()通过减去均值并除以标准差来标准化图像张量的像素值。
- 接下来,我们使用
torchvision.datasets.CIFAR100类加载 CIFAR-100 数据集。我们传递以下参数:root='./data'指定数据集的存储位置。如果数据集不存在,它将被下载到此位置。train=True表示我们要加载训练数据集。要加载验证数据集,请将此设置为False。download=True表示如果数据集不存在,则应下载数据集。transform=transform指定要应用于数据集的数据变换。
- 最后,我们使用
torch.utils.data.DataLoader类创建训练和验证数据加载器。数据加载器用于迭代训练模型的数据集。我们传递以下参数:dataset是要迭代的数据集。batch_size=32指定每个批次的大小。shuffle=True表示在每次 epoch 后对训练数据进行随机排序。这有助于提高模型的泛化能力。对于验证数据加载器,我们将shuffle设置为False,因为我们不需要对验证数据进行随机排序。
通过执行这些步骤,您可以加载 CIFAR-100 数据集并创建相应的训练和验证数据加载器,以用于训练和评估您的机器学习模型。
原文地址: https://www.cveoy.top/t/topic/byfn 著作权归作者所有。请勿转载和采集!