要获取 MNIST 数据集,你可以使用 PyTorch 提供的 'torchvision.datasets' 模块。下面是一个获取 MNIST 数据集的示例代码:

import torchvision.datasets as datasets
import torchvision.transforms as transforms

# 定义数据预处理的转换
transform = transforms.Compose([
    transforms.ToTensor(),  # 将图像转换为 Tensor
    transforms.Normalize((0.5,), (0.5,))  # 标准化图像数据
])

# 下载并加载 MNIST 训练集
train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)

# 下载并加载 MNIST 测试集
test_dataset = datasets.MNIST(root='./data', train=False, download=True, transform=transform)

# 获取训练集和测试集的 DataLoader
train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=64, shuffle=True)
test_loader = torch.utils.data.DataLoader(test_dataset, batch_size=64, shuffle=False)

在上面的代码中,我们使用 'datasets.MNIST' 函数来下载和加载 MNIST 数据集。'root' 参数指定了数据集文件的存储路径,'download=True' 表示如果数据集文件不存在,则自动下载。'transform' 参数用于定义预处理的转换操作,例如将图像转换为张量并进行标准化处理。

然后,我们使用 'torch.utils.data.DataLoader' 将训练集和测试集封装为 DataLoader 对象,以便于后续进行批量训练和测试。

请确保已安装 torchvision 库,可以使用以下命令进行安装:

pip install torchvision

希望这个解答对你有所帮助。如果还有其他问题,请随时提问。

如何使用 PyTorch 获取 MNIST 数据集

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

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