下载 MNIST、查看张量形状、可视化样本

🎉摘要:本文提供Python代码,用于下载MNIST手写数字数据集,统计各数字样本数,并可视化0~9各一个样本。代码输出包括训练集60000张、测试集10000张图片,以及张量形状和像素值范围。适合机器学习初学者学习数据加载与张量操作。

把下面代码保存为 chapter02_data.py,运行后会在 ./data/ 目录下载 MNIST,并弹窗显示第一张图片。

代码如下:

"""
第 2 章:认识数据 —— 下载 MNIST、查看张量形状、可视化样本

运行方式:python chapter02_data.py
"""
import torch
from torchvision import datasets, transforms
import matplotlib.pyplot as plt

# 解决中文乱码 这两行加在plt导入之后
plt.rcParams["font.family"] = ["Microsoft YaHei"]
plt.rcParams["axes.unicode_minus"] = False  # 解决负号显示异常


# 1. 下载 MNIST 数据集
print("正在下载/加载 MNIST 数据集...")
train_dataset = datasets.MNIST(  # 下载训练数据集
    root="./data",
    train=True,
    download=True,
    transform=transforms.ToTensor()  # 自动把图片转成 Tensor(值范围 0~1)
)
test_dataset = datasets.MNIST(  # 下载测试数据集
    root="./data",
    train=False,
    download=True,
    transform=transforms.ToTensor()
)

print(f"训练集:{len(train_dataset)} 张图片")
print(f"测试集:{len(test_dataset)} 张图片")


# 2. 查看数据信息
# 看看各数字有多少样本
# Counter 是 Python 内置的计数器工具,用来快速统计可迭代对象(列表、元组、字符串等)里面各
# 个元素出现的次数,返回一个类似字典的对象:{元素:出现次数}。
from collections import Counter

label_counts = Counter()
for _, label in train_dataset:
    label_counts[int(label)] += 1
    
print("\n各数字样本数:")
for i in range(10):
    print(f"  数字 {i}: {label_counts[i]:5d} 张")


# 3. 看几个样本长什么样
# 创建画布,生成2行5列的子图;figsize设置整张图的宽10英寸、高4英寸
fig, axes = plt.subplots(2, 5, figsize=(10, 4))
# flatten():把二维(2,5)的axes数组摊平成一维数组,方便用下标i直接访问每一张子图
axes = axes.flatten()

for i in range(10):
    # 找到第一个标签为 i 的图片
    for img, label in train_dataset:
        if label == i:
            axes[i].imshow(img.squeeze(), cmap='gray')
            axes[i].set_title(f"数字 {i}")
            axes[i].axis('off')
            break

plt.suptitle("MNIST 数据集:0~9 各一个样本")
plt.tight_layout()
plt.show()


# 4. 理解张量操作
# 取第一张图,看看它的数值范围
img, _ = train_dataset[0]
print(f"\n图片 Tensor 形状: {img.shape}")
print(f"像素值范围: {img.min().item():.4f} ~ {img.max().item():.4f}")
print(f"数据类型: {img.dtype}")
# 注意:ToTensor() 会自动把像素从 0~255 缩放到 0~1

运行示例,输出如下:

正在下载/加载 MNIST 数据集...
训练集:60000 张图片
测试集:10000 张图片

各数字样本数:
  数字 0:  5923 张
  数字 1:  6742 张
  数字 2:  5958 张
  数字 3:  6131 张
  数字 4:  5842 张
  数字 5:  5421 张
  数字 6:  5918 张
  数字 7:  6265 张
  数字 8:  5851 张
  数字 9:  5949 张

图片 Tensor 形状: torch.Size([1, 28, 28])
像素值范围: 0.0000 ~ 1.0000
数据类型: torch.float32

其中,各数字样本数显示如下图:

   

说说我的看法
全部评论(
没有评论
关于
本网站专注于 Java、数据库(MySQL、Oracle)、Linux、软件架构及大数据等多领域技术知识分享。涵盖丰富的原创与精选技术文章,助力技术传播与交流。无论是技术新手渴望入门,还是资深开发者寻求进阶,这里都能为您提供深度见解与实用经验,让复杂编码变得轻松易懂,携手共赴技术提升新高度。如有侵权,请来信告知:hxstrive@outlook.com
其他应用
公众号