如何修改Keras CNN代码从本地文件夹加载图像并设置shuffle=True
Keras实现MNIST手写数字识别CNN代码修改方案
一、需修改的具体代码位置
总共两处核心修改:
- 导入模块段修改
删除原代码中内置数据集导入行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灰度图。 - 数据加载段替换
找到原代码中(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
相关产品推荐
相关产品推荐

