Streamlit果蔬营养识别应用上传图片触发IndexError错误求助
Streamlit果蔬营养查询应用上传图片触发IndexError问题排查与解决
我开发了一个基于Streamlit的果蔬营养信息查询应用,通过CNN模型识别上传图片对应的果蔬种类并返回营养数据,但上传图片时触发IndexError: list index out of range错误。
报错信息
IndexError: list index out of range
报错堆栈
File "E:\pythonsoft\lib\site-packages\streamlit\runtime\scriptrunner\script_runner.py", line 565, in _run_script exec(code, module.__dict__) File "D:\Python_Projects\Fruit_vegetable_nutrients\main.py", line 97, in <module> main() File "D:\Python_Projects\Fruit_vegetable_nutrients\main.py", line 82, in main label = predict_label(image) File "D:\Python_Projects\Fruit_vegetable_nutrients\main.py", line 53, in predict_label prediction = model.predict(image_array) File "E:\pythonsoft\lib\site-packages\keras\engine\training.py", line 1720, in predict data_handler = data_adapter.get_data_handler( File "E:\pythonsoft\lib\site-packages\keras\engine\data_adapter.py", line 1383, in get_data_handler return DataHandler(*args, **kwargs) File "E:\pythonsoft\lib\site-packages\keras\engine\data_adapter.py", line 1138, in __init__ self._adapter = adapter_cls( File "E:\pythonsoft\lib\site-packages\keras\engine\data_adapter.py", line 658, in __init__ self._internal_adapter = TensorLikeDataAdapter( File "E:\pythonsoft\lib\site-packages\keras\engine\data_adapter.py", line 240, in __init__ num_samples = set(int(i.shape[0]) for i in tf.nest.flatten(inputs)).pop() File "E:\pythonsoft\lib\site-packages\keras\engine\data_adapter.py", line 240, in <genexpr> num_samples = set(int(i.shape[0]) for i in tf.nest.flatten(inputs)).pop() File "E:\pythonsoft\lib\site-packages\tensorflow\python\framework\tensor_shape.py", line 896, in __getitem__ return self._dims[key].value
问题代码
import numpy as np import streamlit as st from keras.models import load_model from keras.preprocessing.image import load_img, img_to_array from PIL import Image # Load the trained CNN model model = load_model('FV.h5') # Dictionary of labels and corresponding fruit/vegetable names labels = {0: 'apple', 1: 'banana', 2: 'beetroot', 3: 'bell pepper', 4: 'cabbage', 5: 'capsicum', 6: 'carrot', 7: 'cauliflower', 8: 'chilli pepper', 9: 'corn', 10: 'cucumber', 11: 'eggplant', 12: 'garlic', 13: 'ginger', 14: 'grapes', 15: 'jalepeno', 16: 'kiwi', 17: 'lemon', 18: 'lettuce', 19: 'mango', 20: 'onion', 21: 'orange', 22: 'paprika', 23: 'pear', 24: 'peas', 25: 'pineapple', 26: 'pomegranate', 27: 'potato', 28: 'raddish', 29: 'soy beans', 30: 'spinach', 31: 'sweetcorn', 32: 'sweetpotato', 33: 'tomato', 34: 'turnip', 35: 'watermelon'} fruits = ['Apple', 'Banana', 'Bello Pepper', 'Chilli Pepper', 'Grapes', 'Jalepeno', 'Kiwi', 'Lemon', 'Mango', 'Orange', 'Paprika', 'Pear', 'Pineapple', 'Pomegranate', 'Watermelon'] vegetables = ['Beetroot', 'Cabbage', 'Capsicum', 'Carrot', 'Cauliflower', 'Corn', 'Cucumber', 'Eggplant', 'Ginger', 'Lettuce', 'Onion', 'Peas', 'Potato', 'Raddish', 'Soy Beans', 'Spinach', 'Sweetcorn', 'Sweetpotato', 'Tomato', 'Turnip'] # Dictionary of nutritional information for different types of fruits and vegetables nutrition_info = { 'apple': {'calories': 95, 'vitamins': ['Vitamin C', 'Vitamin K'], 'fat': 0.3, 'protein': 1}, 'banana': {'calories': 105, 'vitamins': ['Vitamin C', 'Vitamin B6'], 'fat': 0.4, 'protein': 1.3}, 'bell pepper': {'calories': 20, 'vitamins': ['Vitamin C', 'Vitamin A'], 'fat': 0.2, 'protein': 0.9}, # ... } def preprocess_image(image_file): # Read the image file and resize it to the required input size for the model image = load_img(image_file, target_size=(224, 224)) # Convert the image to a numpy array image = img_to_array(image) # Normalize the pixel values image = image / 255 image = np.expand_dims(image, [0]) # Make a prediction using the CNN model prediction = model.predict(image) # Get the index of the highest prediction label_index = np.argmax(prediction) # Get the corresponding label from the labels list label = labels[label_index] return label def predict_label(image_array): # Make a prediction using the CNN model prediction = model.predict(image_array) # Get the index of the highest prediction label_index = np.argmax(prediction) # Get the corresponding label from the labels list label = labels[label_index] return label def get_nutrition_info(label): # Get the nutrition information for the label info = nutrition_info[label] return info # Define the main function def main(): # Use Streamlit to create a file uploader widget uploaded_file = st.file_uploader("Choose an image...", type="jpg") if uploaded_file is not None: # Preprocess the image image = Image.open(uploaded_file).resize((250, 250)) st.image(image, use_column_width=False) uploaded_file_path = './upload_images/' + uploaded_file.name with open(uploaded_file_path, "wb") as f: f.write(uploaded_file.getbuffer()) image = preprocess_image(uploaded_file_path) # Get the predicted label label = predict_label(image) if image in labels: st.info('**Category: Vegetables**') else: st.info('**Category : Fruits**') # Get the nutrition information for the label info = get_nutrition_info(label) # Display the predicted label and nutrition information st.write("Predicted label: ", label) st.write("Calories: ", info['calories']) st.write("Vitamins: ", info['vitamins']) st.write("Fat: ", info['fat']) st.write("Protein: ", info['protein']) main()
我希望实现上传图片后获取果蔬营养值的功能,目前上传图片时触发上述错误,请求排查并解决。
问题根源
preprocess_image函数返回值错误:该函数内部已完成预测并返回字符串类型的标签(如'apple'),但后续将这个标签传入predict_label函数,而predict_label期望接收预处理后的图像数组,导致模型无法处理字符串输入,触发维度错误。- 分类判断逻辑错误:
if image in labels中,image是字符串标签,而labels的键是数字索引,永远无法匹配,分类判断完全失效。 - 冗余预测步骤:
preprocess_image已完成一次预测,后续调用predict_label属于重复操作,浪费资源。
修复后的代码
import numpy as np import streamlit as st from keras.models import load_model from keras.preprocessing.image import load_img, img_to_array from PIL import Image # Load the trained CNN model model = load_model('FV.h5') # Dictionary of labels and corresponding fruit/vegetable names labels = {0: 'apple', 1: 'banana', 2: 'beetroot', 3: 'bell pepper', 4: 'cabbage', 5: 'capsicum', 6: 'carrot', 7: 'cauliflower', 8: 'chilli pepper', 9: 'corn', 10: 'cucumber', 11: 'eggplant', 12: 'garlic', 13: 'ginger', 14: 'grapes', 15: 'jalepeno', 16: 'kiwi', 17: 'lemon', 18: 'lettuce', 19: 'mango', 20: 'onion', 21: 'orange', 22: 'paprika', 23: 'pear', 24: 'peas', 25: 'pineapple', 26: 'pomegranate', 27: 'potato', 28: 'raddish', 29: 'soy beans', 30: 'spinach', 31: 'sweetcorn', 32: 'sweetpotato', 33: 'tomato', 34: 'turnip', 35: 'watermelon'} # 转换为小写,方便后续匹配 fruits = [f.lower() for f in ['Apple', 'Banana', 'Bello Pepper', 'Chilli Pepper', 'Grapes', 'Jalepeno', 'Kiwi', 'Lemon', 'Mango', 'Orange', 'Paprika', 'Pear', 'Pineapple', 'Pomegranate', 'Watermelon']] vegetables = [v.lower() for v in ['Beetroot', 'Cabbage', 'Capsicum', 'Carrot', 'Cauliflower', 'Corn', 'Cucumber', 'Eggplant', 'Ginger', 'Lettuce', 'Onion', 'Peas', 'Potato', 'Raddish', 'Soy Beans', 'Spinach', 'Sweetcorn', 'Sweetpotato', 'Tomato', 'Turnip']] # Dictionary of nutritional information for different types of fruits and vegetables nutrition_info = { 'apple': {'calories': 95, 'vitamins': ['Vitamin C', 'Vitamin K'], 'fat': 0.3, 'protein': 1}, 'banana': {'calories': 105, 'vitamins': ['Vitamin C', 'Vitamin B6'], 'fat': 0.4, 'protein': 1.3}, 'bell pepper': {'calories': 20, 'vitamins': ['Vitamin C', 'Vitamin A'], 'fat': 0.2, 'protein': 0.9}, # 补充其余果蔬的营养信息 } def preprocess_image(image_file): # Read the image file and resize it to the required input size for the model image = load_img(image_file, target_size=(224, 224)) # Convert the image to a numpy array image = img_to_array(image) # Normalize the pixel values image = image / 255 # 添加batch维度 image = np.expand_dims(image, axis=0) return image def predict_label(image_array): # Make a prediction using the CNN model prediction = model.predict(image_array) # Get the index of the highest prediction label_index = np.argmax(prediction) # Get the corresponding label from the labels list label = labels[label_index] return label def get_nutrition_info(label): # Get the nutrition information for the label info = nutrition_info[label] return info # Define the main function def main(): # Use Streamlit to create a file uploader widget uploaded_file = st.file_uploader("Choose an image...", type="jpg") if uploaded_file is not None: # 显示上传的图片 display_image = Image.open(uploaded_file).resize((250, 250)) st.image(display_image, use_column_width=False) # 保存图片到本地(可选,若模型不需要本地路径可跳过) uploaded_file_path = './upload_images/' + uploaded_file.name with open(uploaded_file_path, "wb") as f: f.write(uploaded_file.getbuffer()) # 预处理图像 processed_image = preprocess_image(uploaded_file_path) # 获取预测标签 label = predict_label(processed_image) # 判断果蔬类别 if label in vegetables: st.info('**Category: Vegetables**') elif label in fruits: st.info('**Category: Fruits**') else: st.warning('**Category: Unknown**') # 获取营养信息并展示 if label in nutrition_info: info = get_nutrition_info(label) st.write("Predicted label: ", label) st.write("Calories: ", info['calories']) st.write("Vitamins: ", ', '.join(info['vitamins'])) st.write("Fat: ", info['fat'], "g") st.write("Protein: ", info['protein'], "g") else: st.error(f"No nutrition information found for {label}") main()
额外优化建议
- 可直接使用
uploaded_file对象进行预处理,无需保存到本地,减少IO操作:def preprocess_image(image_file): # 直接从文件对象加载图片 image = load_img(image_file, target_size=(224, 224)) image = img_to_array(image) image = image / 255 image = np.expand_dims(image, axis=0) return image # 在main函数中替换为: processed_image = preprocess_image(uploaded_file) - 确保
nutrition_info字典包含所有labels中的键,避免KeyError。
内容的提问来源于stack exchange,提问作者Pemula Saketh
相关产品推荐
相关产品推荐

