torchvision.datasets.MNIST 是 PyTorch 框架中用于下载和处理 MNIST 手写数字数据集的模块。MNIST 数据集包含了大量的手写数字图像,每个图像都是 28 x 28 像素大小的灰度图像。

对于 MNIST 数据集中的每一个图像, torchvision.datasets.MNIST 返回一个由两个元素组成的元组。第一个元素是 PIL.Image.Image 对象类型的图像,表示该图像的像素矩阵。而第二个元素则是一个整数,表示该图像所代表的数字。

因此我们可以使用如下代码调用MNIST数据集中的一部分数据:

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

# 下载 MNIST 数据集
mnist_train = dset.MNIST('./', train=True, transform=transforms.ToTensor(), download=True)
mnist_test = dset.MNIST('./', train=False, transform=transforms.ToTensor(), download=True)

# 获取 MNIST 数据集中的第一张训练图像
image, label = mnist_train[0]

print(image.shape) # 打印图像尺寸 (1, 28, 28)
print(label) # 打印标签

通过这些代码,我们可以看到 MNIST 数据集中的每个样本都是一个尺寸为(1, 28, 28)的张量,其中 1 表示普通的灰度图像。同时,还能得到该图像所代表的数字。

更多推荐