Python实现MNIST各数字样本数统计与训练集图像可视化
首先确保你已经导入了所需依赖:
import numpy as np import matplotlib.pyplot as plt
统计各数字类别样本数量
直接使用numpy的bincount方法即可快速完成计数,针对你已经加载好的训练集、测试集标签数组,代码如下:
# 标签从loadtxt读取后默认是float类型,先转成int,指定minlength确保输出长度固定为10 train_class_count = np.bincount(t_train.astype(np.int32), minlength=10) test_class_count = np.bincount(t_test.astype(np.int32), minlength=10) # 打印统计结果 for num in range(10): print(f"类别{num}:训练集{train_class_count[num]}个样本,测试集{test_class_count[num]}个样本")
说明:np.bincount会统计数组中每个非负整数的出现频次,针对MNIST这种标签为0-9连续整数的场景,比通用的value_counts方法效率更高。
训练集每类至少9张图像可视化
采用10行9列的子图布局,每行对应一个数字类别,每行展示9张该类别的样本,代码如下:
# MNIST默认图像尺寸为28*28,如果你之前未定义该变量请补充 image_size = 28 # 初始化画布,10行对应0-9共10个数字,9列对应每类展示9张图 fig, axes = plt.subplots(nrows=10, ncols=9, figsize=(9, 10)) # 调整子图间距,避免标签、图像重叠 plt.subplots_adjust(wspace=0.03, hspace=0.25) for digit in range(10): # 筛选出当前数字在训练集中的所有样本索引 sample_idx = np.where(t_train == digit)[0] # 随机抽取9个不重复的样本索引,不需要随机可替换为 sample_idx[:9] 取前9个 selected_idx = np.random.choice(sample_idx, size=9, replace=False) # 逐张绘制图像 for col, idx in enumerate(selected_idx): ax = axes[digit, col] # 将一维784维像素向量重构为28*28的二维图像,用灰度色阶显示 ax.imshow(X_train[idx].reshape(image_size, image_size), cmap="gray") # 隐藏坐标轴刻度和边框 ax.axis("off") # 在每行第一列的子图上标注当前类别 if col == 0: ax.set_title(f"数字 {digit}", fontsize=10) plt.show()
说明:如果需要每类展示更多样本,只需修改子图的ncols参数和np.random.choice的size参数即可。
内容的提问来源于stack exchange,提问作者user15295241
相关产品推荐
相关产品推荐

