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

Keras加载模型仅输出0/1,本地与Kaggle置信度结果不一致问题求助

问题:Kaggle训练的EfficientNetB4模型本地加载后predict输出仅为0或1,而非置信度

我在Kaggle上训练了一个基于EfficientNetB4的二分类模型,训练阶段执行model.predict()能正常返回0-1之间的置信度数值,但将模型导出到本地加载后,model.predict()仅返回0或1的整数数组。尝试过保存为.h5格式、仅保存权重再加载,问题依旧存在,且模型未出现过拟合情况。


模型定义代码

from tensorflow.keras.applications import EfficientNetB4
from tensorflow.keras.models import Model

base_model = EfficientNetB4(input_tensor=Input(shape=(IMG_HEIGHT, IMG_WIDTH, 3)),
                            weights='imagenet',
                            include_top=False,
                            pooling='avg'
                           )
x=base_model.output
output=Dense(1, activation='sigmoid')(x)
model=Model(inputs=base_model.input, outputs=output)
model.summary()

Kaggle模型保存代码

MODEL_DIR = "../working/tfx_model/"
version = "alpha"
export_path = os.path.join(MODEL_DIR, str(version))
print('export_path = {}\n'.format(export_path))

tf.keras.models.save_model(
    model,
    export_path,
    overwrite=True,
    include_optimizer=True,
    save_format=None,
    signatures=None,
    options=None
)

print('\nSaved model:')
!ls -l {export_path}

本地模型加载代码

model = load_model('models/tfx_model')

环境版本

  • Kaggle环境:Tensorflow 2.9.2,Keras 2.9.0
  • 本地环境:Tensorflow 2.10.0,Keras 2.10.0

可能的原因及解决方法

  1. 输入数据预处理不一致
    EfficientNet的Keras实现要求输入数据必须匹配官方预处理逻辑,如果训练时用了tf.keras.applications.efficientnet.preprocess_input,但本地预测时遗漏这一步,会导致输出异常。
    解决:确保本地预测时和训练阶段做完全一致的预处理:

    from tensorflow.keras.applications.efficientnet import preprocess_input
    
    # 加载图像后先执行预处理
    input_image = preprocess_input(input_image)
    predictions = model.predict(input_image)
    
  2. TensorFlow版本兼容性问题
    不同小版本间的SavedModel格式可能存在细微差异,导致模型加载后层行为变化。
    解决:

    • 本地降级TensorFlow到2.9.2,和Kaggle环境保持一致后再测试。
    • 在Kaggle上改用tf.saved_model.save()保存模型,再尝试本地加载。
  3. 输入数据的 dtype 或形状不匹配
    若本地输入图像是uint8整数类型,或形状缺少batch维度,可能导致sigmoid输出异常。
    解决:

    # 确保输入转换为float32类型
    input_image = input_image.astype('float32')
    # 单张图像需扩展batch维度
    if len(input_image.shape) == 3:
        input_image = tf.expand_dims(input_image, 0)
    
  4. 模型加载时激活函数异常
    加载过程中可能出现激活函数被意外替换的情况,可检查最后一层的激活设置:

    print(model.layers[-1].activation)
    

    若输出不是sigmoid,尝试在Kaggle上用save_format='h5'保存模型,本地加载时确保导入所有依赖(如EfficientNet的定义)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 02:36:30