复现Keras迁移学习教程:跳过缓存步骤及解决输入形状不匹配错误
问题解决方法
报错原因
你遇到的形状不匹配报错,本质是模型要求输入为带批次维度的批量数据(形状格式为(None, 150, 150, 3),第一个维度为批次大小),跳过原代码后,数据集返回的单张图片形状为(150, 150, 3),缺失批次维度才触发报错。原代码中的cache()和prefetch()仅用于优化数据读取速度,和形状匹配无关,你完全可以去掉这两个配置,只要保证数据集能输出带批次维度的数据即可。
可用修改方案
方案1:仅保留batch操作
直接删掉cache()和prefetch()配置,仅保留生成批次的batch()方法即可:
batch_size = 32 train_ds = train_ds.batch(batch_size) validation_ds = validation_ds.batch(batch_size) test_ds = test_ds.batch(batch_size)
方案2:完全不修改数据集处理代码,在fit时指定batch_size
如果你不想对数据集处理逻辑做任何修改,直接在调用model.fit()时传入batch_size参数即可,框架会自动对数据集按指定大小打包,补充批次维度:
# 示例fit调用,其余参数按你原有配置保留即可 model.fit( train_ds, validation_data=validation_ds, epochs=你的训练轮数, batch_size=32 )
注意事项
- 采用方案2时,要保证数据集没有提前执行过
batch()操作,否则会重复叠加批次维度导致新的形状错误。 - 去掉
cache()配置后,每轮训练都会重新从磁盘读取图片,训练速度会有所下降,但功能完全正常。
内容的提问来源于stack exchange,提问作者Francesco Grossetti
相关产品推荐
相关产品推荐

