把下面代码保存为 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其中,各数字样本数显示如下图:
