Keras ImageDataGenerator.fit()秩4/reshape尺寸不匹配报错修复
LFW数据集运行LeNet模型维度适配问题修复方案
报错根因
- 首次触发的rank 4报错:TensorFlow/Keras框架的卷积神经网络输入要求为4维张量,维度顺序为
(样本总数, 图像高度, 图像宽度, 通道数),加载得到的原始训练集shape为(704, 227, 227),属于3维数组,缺失了最后一维的通道数维度(灰度图通道数为1,RGB彩色图通道数为3),不符合输入要求。 - 手动reshape触发的size不匹配报错:手动指定reshape目标形状的第一维(样本数)为96,但实际训练集总样本数为704,总元素数校验不通过。计算可得704227227=36276416,和报错提示的数组总元素数完全一致,硬编码样本数会直接导致维度计算错误。
正确修复步骤
- 对训练、验证、测试三个数据集,动态补全通道维度,禁止硬编码样本数参数:
若加载的是灰度格式LFW数据,执行如下代码完成维度转换:
若加载的是RGB三通道格式LFW数据,将上述代码中reshape的最后一个参数从1改为3即可。# 补全训练集通道维度 xtrain = xtrain.reshape(xtrain.shape[0], xtrain.shape[1], xtrain.shape[2], 1) # 补全验证集通道维度 xval = xval.reshape(xval.shape[0], xval.shape[1], xval.shape[2], 1) # 补全测试集通道维度 xtest = xtest.reshape(xtest.shape[0], xtest.shape[1], xtest.shape[2], 1) - 核对LeNet模型输入层的
input_shape参数,确保和单张图片维度匹配:单通道227*227尺寸输入对应参数值为(227,227,1),三通道则为(227,227,3)。 - 维度转换完成后,直接将4维格式的xtrain传入
train_generator.fit(xtrain)即可正常执行,无需额外修改fit方法的其他参数。
注意事项
如果加载数据集后做了训练/验证/测试集拆分,不需要手动统计各子集的样本数量,直接用数组自带的shape[0]属性动态获取当前子集的样本总数,就能完全避免reshape时元素数不匹配的问题。
内容的提问来源于stack exchange,提问作者Josh
相关产品推荐
相关产品推荐

