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

如何修改Keras CNN代码从本地文件夹加载图像并设置shuffle=True

Keras实现MNIST手写数字识别CNN代码修改方案

一、需修改的具体代码位置

总共两处核心修改:

  1. 导入模块段修改
    删除原代码中内置数据集导入行from keras.datasets import mnist,新增本地图像读取所需的依赖导入,替换后导入段如下:
    from __future__ import print_function
    import keras
    import os
    import numpy as np
    from keras.preprocessing.image import load_img, img_to_array
    from keras.models import Sequential
    from keras.layers import Dense, Dropout, Flatten
    from keras.layers import Conv2D, MaxPooling2D
    from keras import backend as K
    

    本地数据集目录要求:根目录下需建train、test两个子文件夹,每个子文件夹下分别建0-9共10个数字命名的文件夹,对应存放该数字的手写数字图片,图片尺寸建议提前调整为28*28灰度图。

  2. 数据加载段替换
    找到原代码中(x_train, y_train), (x_test, y_test) = mnist.load_data()这一行,将整行替换为自定义本地数据加载逻辑,替换后的代码段如下(注意把load_local_mnist的入参改成你自己的数据集根目录路径):
    batch_size = 128
    num_classes = 10
    epochs = 12
    
    # input image dimensions
    img_rows, img_cols = 28, 28
    
    # 自定义本地MNIST数据集加载函数
    def load_local_mnist(data_path):
        def load_single_split(split_dir):
            x_data, y_data = [], []
            for label in range(10):
                class_dir = os.path.join(split_dir, str(label))
                for img_file in os.listdir(class_dir):
                    # 读取为28*28灰度图
                    img = load_img(
                        os.path.join(class_dir, img_file),
                        color_mode="grayscale",
                        target_size=(img_rows, img_cols)
                    )
                    x_data.append(img_to_array(img))
                    y_data.append(label)
            return np.array(x_data), np.array(y_data)
        train_x, train_y = load_single_split(os.path.join(data_path, "train"))
        test_x, test_y = load_single_split(os.path.join(data_path, "test"))
        return (train_x, train_y), (test_x, test_y)
    
    # 替换原内置数据集加载接口,括号内填写本地数据集根目录路径
    (x_train, y_train), (x_test, y_test) = load_local_mnist("./mnist_data")
    
    if K.image_data_format() == 'channels_first':
        x_train = x_train.reshape(x_train.shape[0], 1, img_rows, img_cols)
        x_test = x_test.reshape(x_test.shape[0], 1, img_rows, img_cols)
        input_shape = (1, img_rows, img_cols)
    else:
        x_train = x_train.reshape(x_train.shape[0], img_rows, img_cols, 1)
        x_test = x_test.reshape(x_test.shape[0], img_rows, img_cols, 1)
        input_shape = (img_rows, img_cols, 1)
    

其余归一化、标签转换、模型搭建、模型编译的代码段不需要改动。

二、shuffle=True参数设置位置

有两个可选位置,按需选择即可:

  • 位置1:手动打乱(加载数据后立刻执行)
    在本地数据加载完成、reshape维度转换之前,加入如下代码即可一次性打乱训练集顺序:
    # 生成打乱后的索引
    permutation = np.random.permutation(x_train.shape[0])
    x_train = x_train[permutation]
    y_train = y_train[permutation]
    
  • 位置2:模型训练时自动打乱(推荐)
    在调用model.fit()执行训练时,直接传入shuffle=True参数,Keras会在每个训练epoch开始前自动打乱训练集顺序,训练段代码示例如下:
    model1.fit(
        x_train, y_train,
        batch_size=batch_size,
        epochs=epochs,
        verbose=1,
        validation_data=(x_test, y_test),
        shuffle=True  # 此处设置即可
    )
    

内容的提问来源于stack exchange,提问作者Awan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 08:24:15