Python生成器接入Keras model.fit_generator报错求助及实现咨询
解决Keras fit_generator输入维度错误及生成器实现合理性分析
一、先搞定输入维度错误问题
你碰到的ValueError核心原因很直接:Keras的卷积层要求输入是4维张量,格式为(batch_size, height, width, channels),但你的生成器每次只返回单张3维的图片(28,28,1),少了关键的批量维度。
这里有两种解决思路,推荐第二种更贴合训练逻辑的方式:
快速修复:给单张图加批量维度
在生成器里用np.expand_dims给图片和标签扩展维度,修改你的gen方法:import numpy as np # 记得导入numpy def gen(self, feat, labels): i=0 total_samples = len(feat) while True: # 循环到末尾重置索引,避免越界报错 if i >= total_samples: i = 0 im = cv2.imread(feat[i],0) im = im.reshape(28,28,1) # 增加批量维度,变成(1,28,28,1) im = np.expand_dims(im, axis=0) label = np.expand_dims(labels[i], axis=0) yield im, label i+=1更优方案:生成批量样本
每次返回一个批次的图片和标签(比如常用的batch_size=32),这才是Keras训练的常规操作:def gen(self, feat, labels, batch_size=32): total_samples = len(feat) while True: # 随机打乱样本顺序,提升模型泛化能力 indices = np.random.permutation(total_samples) for start in range(0, total_samples, batch_size): end = min(start + batch_size, total_samples) batch_indices = indices[start:end] batch_imgs = [] batch_labels = [] for idx in batch_indices: im = cv2.imread(feat[idx], 0) if im is None: # 处理图片读取失败的情况 print(f"警告:无法读取图片 {feat[idx]}") continue im = im.reshape(28,28,1) batch_imgs.append(im) batch_labels.append(labels[idx]) yield np.array(batch_imgs), np.array(batch_labels)同时要修改
fit_generator的steps_per_epoch,设置为总样本数//batch_size(如果有余数就加1,确保所有样本都被训练到)。
二、你的生成器实现合理吗?
实话实说,原实现有不少可以优化的点:
- 类封装冗余:完全没必要专门写
Generator类,直接用生成器函数更简洁直观 - 无边界处理:原代码里
i会一直递增,超出feat长度后会直接抛出IndexError - 缺少数据打乱:固定顺序训练样本,会影响模型的泛化能力
- 过时API使用:
pd.get_dummies(...).as_matrix()已经被废弃,应该改用.values或者.to_numpy() - 无错误处理:没考虑图片读取失败的情况(比如文件损坏、路径错误)
三、优化后的完整代码
from keras.utils import to_categorical from keras.models import Sequential from keras.layers import Dense, Conv2D, Flatten import pandas as pd import os import cv2 import numpy as np def data_generator(feat, labels, batch_size=32): total_samples = len(feat) while True: # 随机打乱样本索引 indices = np.random.permutation(total_samples) for start in range(0, total_samples, batch_size): end = min(start + batch_size, total_samples) batch_indices = indices[start:end] batch_imgs = [] batch_labels = [] for idx in batch_indices: img_path = feat[idx] im = cv2.imread(img_path, 0) if im is None: print(f"Warning: Failed to read image {img_path}") continue # 强制调整图片尺寸为28x28,避免尺寸不一致问题 im = cv2.resize(im, (28,28)) im = im.reshape(28,28,1) # 这里可以添加数据增强操作,比如随机翻转、旋转等 batch_imgs.append(im) batch_labels.append(labels[idx]) yield np.array(batch_imgs), np.array(batch_labels) if __name__ == "__main__": input_dir = './mnist' output_file = 'dataset.csv' filename = [] label = [] for root,dirs,files in os.walk(input_dir): for file in files: full_path = os.path.join(root,file) filename.append(full_path) label.append(os.path.basename(os.path.dirname(full_path))) data = pd.DataFrame(data={'filename': filename, 'label':label}) data.to_csv(output_file,index=False) feat = data['filename'].values # 替换过时的as_matrix() labels = pd.get_dummies(data['label']).to_numpy() batch_size = 32 steps_per_epoch = len(feat) // batch_size # 处理样本数不能被batch_size整除的情况 if len(feat) % batch_size != 0: steps_per_epoch +=1 # 创建模型 model = Sequential() model.add(Conv2D(64, kernel_size=3, activation="relu", input_shape=(28,28,1))) model.add(Conv2D(32, kernel_size=3, activation="relu")) model.add(Flatten()) model.add(Dense(2, activation="softmax")) model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy']) # 使用生成器训练 model.fit_generator(data_generator(feat, labels, batch_size), steps_per_epoch=steps_per_epoch, epochs=5, verbose=1)
内容的提问来源于stack exchange,提问作者eng2019
相关产品推荐
相关产品推荐

