使用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
相关产品推荐
相关产品推荐

