求助:从零开始加载MNIST数据集并划分训练-验证-测试集
从零加载并划分MNIST数据集(仅用Python内置+NumPy)
前置说明
MNIST原始数据集包含4个二进制压缩文件:
- 训练集图像(train-images-idx3-ubyte.gz)
- 训练集标签(train-labels-idx1-ubyte.gz)
- 测试集图像(t10k-images-idx3-ubyte.gz)
- 测试集标签(t10k-labels-idx1-ubyte.gz)
你需要先从MNIST官方页面下载这4个文件,放在同一个目录下。
加载数据集函数
用Python内置的gzip模块解析MNIST二进制格式,结合NumPy处理数据结构:
import gzip import numpy as np def load_mnist_images(file_path): # 读取图像文件 with gzip.open(file_path, 'rb') as f: # 跳过前16字节的文件头(MNIST图像文件固定格式) f.read(16) # 读取所有图像字节数据,转为uint8数组 data = np.frombuffer(f.read(), dtype=np.uint8) # 重塑为(样本数, 28, 28)的二维图像格式 return data.reshape(-1, 28, 28) def load_mnist_labels(file_path): # 读取标签文件 with gzip.open(file_path, 'rb') as f: # 跳过前8字节的文件头(MNIST标签文件固定格式) f.read(8) # 读取所有标签数据 return np.frombuffer(f.read(), dtype=np.uint8)
划分训练集与验证集
原始训练集共60000个样本,可从中拆分出部分样本作为验证集(示例取10000个),先打乱数据再拆分:
def split_train_val(images, labels, val_size=10000): # 生成随机索引打乱数据顺序 indices = np.random.permutation(len(images)) # 拆分索引,划分训练/验证集 train_indices = indices[val_size:] val_indices = indices[:val_size] return (images[train_indices], labels[train_indices]), (images[val_indices], labels[val_indices])
完整使用示例
# 替换为你的实际文件路径 train_images_path = 'train-images-idx3-ubyte.gz' train_labels_path = 'train-labels-idx1-ubyte.gz' test_images_path = 't10k-images-idx3-ubyte.gz' test_labels_path = 't10k-labels-idx1-ubyte.gz' # 加载原始数据 train_images = load_mnist_images(train_images_path) train_labels = load_mnist_labels(train_labels_path) test_images = load_mnist_images(test_images_path) test_labels = load_mnist_labels(test_labels_path) # 划分训练集和验证集 (train_imgs, train_lbls), (val_imgs, val_lbls) = split_train_val(train_images, train_labels) # 验证数据形状 print(f"训练集: {train_imgs.shape}, 标签: {train_lbls.shape}") print(f"验证集: {val_imgs.shape}, 标签: {val_lbls.shape}") print(f"测试集: {test_images.shape}, 标签: {test_labels.shape}")
注意事项
- 确保文件路径正确,否则会抛出文件不存在错误
- 若需要固定划分结果,可在
split_train_val函数开头添加np.random.seed(你的种子值) - 图像数据为0-255的uint8格式,如需归一化到0-1范围,可执行
train_imgs = train_imgs / 255.0
内容的提问来源于stack exchange,提问作者tail
相关产品推荐
相关产品推荐

