You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

求助:从零开始加载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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.19 17:50:25