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

使用Streamlit/Python做图像分类时遇SparseCategoricalCrossentropy参数错误

问题:加载Keras模型时出现ignore_class参数错误

在使用Streamlit和Python开发图像分类应用时,调用tf.keras.models.load_model加载预训练模型时触发错误:
SparseCategoricalCrossentropy.__init__() got an unexpected keyword argument 'ignore_class'

报错栈信息

TypeError: SparseCategoricalCrossentropy.__init__() got an unexpected keyword argument 'ignore_class'
Traceback:
File "C:\Users\DELL\AppData\Local\Programs\Python\Python310\lib\site-packages\streamlit\runtime\scriptrunner\script_runner.py", line 565, in _run_script
    exec(code, module.__dict__)
File "C:\Users\DELL\Desktop\Gradio\flask and htm\project-folder\backend\org.py", line 18, in <module>
    model = load_model()
File "C:\Users\DELL\AppData\Local\Programs\Python\Python310\lib\site-packages\streamlit\runtime\legacy_caching\caching.py", line 625, in wrapped_func
    return get_or_create_cached_value()
File "C:\Users\DELL\AppData\Local\Programs\Python\Python310\lib\site-packages\streamlit\runtime\legacy_caching\caching.py", line 609, in get_or_create_cached_value
    return_value = non_optional_func(*args, **kwargs)
File "C:\Users\DELL\Desktop\Gradio\flask and htm\project-folder\backend\org.py", line 8, in load_model
    model = tf.keras.models.load_model('model/my_model.h5')
File "C:\Users\DELL\AppData\Local\Programs\Python\Python310\lib\site-packages\keras\utils\traceback_utils.py", line 67, in error_handler
    raise e.with_traceback(filtered_tb) from None
File "C:\Users\DELL\AppData\Local\Programs\Python\Python310\lib\site-packages\keras\losses.py", line 153, in from_config
    return cls(**config)

应用代码

import streamlit as st
import tensorflow as tf

st.set_option('deprecation.showfileUploaderEncoding', False)
@st.cache(allow_output_mutation=True)
def load_model():
    model = tf.keras.models.load_model('model/my_model.h5')
    return model

model = load_model()
st.write("""
        # Image Classification App
        """
    )

file = st.file_uploader("Please upload an image", type=["jpg", "png"])
import cv2
from PIL import Image, ImageOps
import numpy as np

def import_and_predict(image_data, model):
    size = (180,180)
    image =ImageOps.fit(image_data, size, Image.LANCZOS)
    img =np.asarray(image)
    img_reshape = img[np.newaxis,...]
    prediction = model.predict(img_reshape)
    return prediction

if file is None:
    st.text("Please upload an image file")
else:
    image = Image.open(file)
    st.image(image, use_column_width=True)
    predictions = import_and_predict(image, model)
    class_names = ['dog', 'cat', 'horse']
    string = "This image most likely is a :" +class_names[np.argmax(predictions)]
    st.success(string)

解决方案

1. 升级TensorFlow版本

ignore_class是Keras 2.13及以上版本(对应TensorFlow 2.13+)新增的参数,若当前环境TensorFlow版本低于2.13,就会出现该错误。执行以下命令升级:

pip install --upgrade tensorflow

2. 自定义损失函数加载模型

如果无法升级版本,可在加载模型时自定义兼容的损失函数,忽略ignore_class参数:

@st.cache(allow_output_mutation=True)
def load_model():
    # 定义无ignore_class参数的损失函数
    custom_loss = tf.keras.losses.SparseCategoricalCrossentropy()
    model = tf.keras.models.load_model('model/my_model.h5', 
                                       custom_objects={'SparseCategoricalCrossentropy': custom_loss})
    return model

若业务逻辑需要忽略特定类别,确保版本升级后,可明确指定参数加载:

@st.cache(allow_output_mutation=True)
def load_model():
    # 替换为实际需要忽略的类别ID
    custom_loss = tf.keras.losses.SparseCategoricalCrossentropy(ignore_class=0)
    model = tf.keras.models.load_model('model/my_model.h5', 
                                       custom_objects={'SparseCategoricalCrossentropy': custom_loss})
    return model

3. 重新导出模型

若模型是在高版本TensorFlow环境下训练保存的,低版本环境加载会存在参数兼容问题。建议在与训练环境相同版本的TensorFlow中重新导出模型,再在当前环境加载。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 09:40:28