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

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()

我希望实现上传图片后获取果蔬营养值的功能,目前上传图片时触发上述错误,请求排查并解决。


问题根源

  1. preprocess_image函数返回值错误:该函数内部已完成预测并返回字符串类型的标签(如'apple'),但后续将这个标签传入predict_label函数,而predict_label期望接收预处理后的图像数组,导致模型无法处理字符串输入,触发维度错误。
  2. 分类判断逻辑错误:if image in labels中,image是字符串标签,而labels的键是数字索引,永远无法匹配,分类判断完全失效。
  3. 冗余预测步骤: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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 03:50:29