使用 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)

代码解释:

  1. 首先,我们导入必要的库:torchvision 用于加载数据集,torchvision.transforms 用于数据预处理,torch 用于创建数据加载器。
  2. 然后,我们定义数据变换。数据变换是应用于数据集的预处理步骤。在这里,我们使用 transforms.Compose 将两个变换链接在一起:
    • transforms.ToTensor() 将图像从 PIL 图像转换为 PyTorch 张量。
    • transforms.Normalize() 通过减去均值并除以标准差来标准化图像张量的像素值。
  3. 接下来,我们使用 torchvision.datasets.CIFAR100 类加载 CIFAR-100 数据集。我们传递以下参数:
    • root='./data' 指定数据集的存储位置。如果数据集不存在,它将被下载到此位置。
    • train=True 表示我们要加载训练数据集。要加载验证数据集,请将此设置为 False
    • download=True 表示如果数据集不存在,则应下载数据集。
    • transform=transform 指定要应用于数据集的数据变换。
  4. 最后,我们使用 torch.utils.data.DataLoader 类创建训练和验证数据加载器。数据加载器用于迭代训练模型的数据集。我们传递以下参数:
    • dataset 是要迭代的数据集。
    • batch_size=32 指定每个批次的大小。
    • shuffle=True 表示在每次 epoch 后对训练数据进行随机排序。这有助于提高模型的泛化能力。对于验证数据加载器,我们将 shuffle 设置为 False,因为我们不需要对验证数据进行随机排序。

通过执行这些步骤,您可以加载 CIFAR-100 数据集并创建相应的训练和验证数据加载器,以用于训练和评估您的机器学习模型。


原文地址: https://www.cveoy.top/t/topic/byfn 著作权归作者所有。请勿转载和采集!

免费AI点我,无需注册和登录