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

复现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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 15:36:02