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

Keras训练后保存的.h5模型预测异常问题求助

分析你的猫狗分类CNN模型保存/加载后预测异常的问题

这种情况我之前做图像分类项目时也碰到过,结合你的描述,大概率是下面几个原因之一,咱们一个个排查:

1. 模型保存时机或实例被覆盖

你提到把保存代码移到脚本末尾后,当前脚本预测就正常了——这说明你最初保存的不是训练完成后的那个模型实例。比如脚本后面可能有重新初始化模型的代码(比如又写了model = build_cnn_model()),或者不小心用其他变量覆盖了训练好的model,导致保存的是一个未训练/半训练的空模型。

解决方法:

  • 把model.save('xxx.h5')放在所有训练代码的最后一行,确保之后没有任何修改或重新定义model的操作;
  • 保存前可以打印模型的训练指标(比如print(model.evaluate(test_dataset))),确认模型确实处于训练完成的状态。

2. 预测时的数据预处理和训练时不一致

这是图像分类最常见的“坑”!模型是基于训练时的预处理规则学习的,如果预测时预处理步骤不一样,输入数据的分布完全偏离模型预期,就会输出毫无意义的结果,甚至全偏向某一类。

比如常见的不一致场景:

  • 训练时用ImageDataGenerator(rescale=1./255)做了归一化,但预测时加载图像后没除以255;
  • 训练时图像尺寸是(224,224),预测时加载的图像没resize到对应大小;
  • 训练时做了随机翻转、裁剪等数据增强,但预测时错误地沿用了这些随机操作(预测只需要固定预处理)。

解决方法:

  • 把训练时的预处理逻辑封装成复用函数,比如:
    def preprocess_image(img_path, target_size=(224,224)):
        img = tf.keras.preprocessing.image.load_img(img_path, target_size=target_size)
        img_array = tf.keras.preprocessing.image.img_to_array(img)
        img_array = np.expand_dims(img_array, axis=0)
        img_array /= 255.0  # 和训练时的rescale参数保持一致
        return img_array
    
  • 预测时严格调用这个函数处理输入图像,确保和训练时的输入分布完全匹配。

3. 模型保存/加载方式错误

如果你的模型包含自定义层、自定义损失函数或度量,直接用model.save()和model.load_model()可能会出现加载不完整的情况,导致模型权重或编译参数异常,预测结果全偏。另外,如果错误地只用model.save_weights()保存权重,加载时又没先构建好和训练时一致的模型结构,也会得到一个“空模型”。

解决方法:

  • 若用标准Keras层,优先用model.save('my_model.h5')保存完整模型,加载时用:
    loaded_model = tf.keras.models.load_model('my_model.h5')
    
  • 若有自定义组件,加载时需要指定custom_objects参数,比如:
    loaded_model = tf.keras.models.load_model('my_model.h5', custom_objects={'my_custom_loss': my_custom_loss})
    
  • 加载后打印loaded_model.summary(),确认模型结构和训练时完全一致,再用训练集里的图像测试,看结果是否和训练完成时的预测一致。

4. 类别标签映射错误

训练时用flow_from_directory生成数据,类别顺序是按文件夹名的字母排序的。比如你的数据集是train/cat和train/dog,那么class_mode='binary'时,cat对应的标签是0,dog是1。如果预测时标签映射搞反了,或者模型输出的概率判断逻辑错误,就会出现全预测为某一类的情况。

比如错误的判断逻辑:

# 训练时0=cat,1=dog,但这里搞反了
pred_prob = loaded_model.predict(img_array)[0][0]
label = 'cat' if pred_prob > 0.5 else 'dog'

解决方法:

  • 训练时打印train_generator.class_indices,确认标签和类别的对应关系:
    print(train_generator.class_indices)  # 输出类似 {'cat':0, 'dog':1}
    
  • 预测时严格按照这个对应关系编写判断逻辑。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 09:04:58