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

Keras ResNet50迁移学习预测输入形状不兼容问题求助

解决Keras ResNet50迁移学习预测时的维度不兼容问题

问题根源

  • 缺失Batch维度:Keras模型默认接收批量输入,形状要求为(batch_size, 224, 224, 3),但你传入的单张图像是(224,224,3),无batch维度。模型会误将第一个维度224识别为batch_size,剩余的(224,3)与期望输入形状不匹配,引发报错(报错中的32为显示异常,核心是维度缺失)。
  • 预处理冗余且顺序错误:模型构建时已包含preprocess_input和Rescaling层,预测时重复执行会破坏数据分布,甚至引发维度处理异常。

修复方案

1. 给输入添加Batch维度

在调用predict前,用tf.expand_dims为单张图像添加batch维度:

wg_model.predict(tf.expand_dims(im_arr_scaled, axis=0))

2. 移除冗余预处理步骤

模型内部已包含完整预处理流程,预测时无需重复执行,修改后的预测代码如下:

# 加载图像并转换为数组
ims = keras.utils.load_img(test_files[0], target_size=(224, 224))
im_arr = keras.utils.img_to_array(ims)

# 添加batch维度
im_arr_batch = tf.expand_dims(im_arr, axis=0)

WEIGHTS = "/home/app/src/experiments/exp_007/model.01-5.2777.h5"
wg_model = resnet_50.create_model(weights=WEIGHTS)

# 执行预测
wg_model.predict(im_arr_batch)

3. 修正模型构建中的预处理顺序

ResNet50的preprocess_input基于原始0-255像素值设计,需调整预处理顺序,确保先做缩放再做模型专属预处理:

# 模型构建代码中调整顺序
x = augmentation_layer(inputs)
# 先缩放再做preprocess_input
scale_layer = layers.Rescaling(scale=1./255)
x = scale_layer(x)
x = preprocess_input(x)

验证方法

打印关键步骤的张量形状,确认维度符合要求:

print("im_arr shape:", im_arr.shape)  # 预期输出:(224, 224, 3)
print("im_arr_batch shape:", im_arr_batch.shape)  # 预期输出:(1, 224, 224, 3)

内容的提问来源于stack exchange,提问作者agustin

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 17:52:13