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

训练完成水果新鲜度分类模型后,如何实现正确的单张图片预测?

单张图片预测Keras水果二分类模型出错问题

问题背景

我是机器学习新手,第一个项目是基于Keras训练区分新鲜/腐烂水果的二分类模型,批量测试时模型运行正常,但单张图片预测时,多次调整预处理代码都得到错误输出。以下是训练代码和两次尝试的测试代码:

训练代码

import numpy as np 
import pandas as pd
import os
import cv2
import matplotlib.pyplot as plt
from tqdm import tqdm
from random import shuffle
from keras.utils  import to_categorical
import pickle
def load_rand():
    X=[]
    dir_path='D:/dataset/train'
    for sub_dir in tqdm(os.listdir(dir_path)):
        print(sub_dir)
        path_main=os.path.join(dir_path,sub_dir)
        i=0
        for img_name in os.listdir(path_main):
            if i>=6:
                break
            img=cv2.imread(os.path.join(path_main,img_name))
            img=cv2.resize(img,(100,100))
            img=cv2.cvtColor(img,cv2.COLOR_BGR2RGB)
            X.append(img)
            i+=1
    return X
X=load_rand()
X=np.array(X)
X.shape
def show_subpot(X,title=False,Y=None):
    if X.shape[0]==36:
        f, ax= plt.subplots(6,6, figsize=(40,60))
        list_fruits=['rottenoranges', 'rottenapples', 'freshbanana', 'freshoranges', 'rottenbanana', 'freshapples']
        for i,img in enumerate(X):
            ax[i//6][i%6].imshow(img, aspect='auto')
            if title==False:
                ax[i//6][i%6].set_title(list_fruits[i//6])
            elif title and Y is not None:
                ax[i//6][i%6].set_title(Y[i])
        plt.show()
    else:
        print('Cannot plot')
show_subpot(X)
del X
def load_rottenvsfresh():
    quality=['fresh', 'rotten']
    X,Y=[],[]
    z=[]
    for cata in tqdm(os.listdir('D:/dataset/train')):
        if quality[0] in cata:
            path_main=os.path.join('D:/dataset/train',cata)
            for img_name in os.listdir(path_main):
                img=cv2.imread(os.path.join(path_main,img_name))
                img=cv2.resize(img,(100,100))
                img=cv2.cvtColor(img,cv2.COLOR_BGR2RGB)
                z.append([img,0])
        else:
            path_main=os.path.join('D:/dataset/train',cata)
            for img_name in os.listdir(path_main):
                img=cv2.imread(os.path.join(path_main,img_name))
                img=cv2.resize(img,(100,100))
                img=cv2.cvtColor(img,cv2.COLOR_BGR2RGB)
                z.append([img,1])
    print('Shuffling your data.....')
    shuffle(z)
    for images, labels in tqdm(z):
        X.append(images);Y.append(labels)
    return X,Y
X,Y=load_rottenvsfresh()
Y=np.array(Y)
X=np.array(X)
y_ser=pd.Series(Y)
y_ser.value_counts()
def load_rottenvsfresh_valset():
    quality=['fresh', 'rotten']
    X,Y=[],[]
    z=[]
    for cata in tqdm(os.listdir('D:/dataset/test')):
        if quality[0] in cata:
            path_main=os.path.join('D:/dataset/test',cata)
            for img_name in os.listdir(path_main):
                img=cv2.imread(os.path.join(path_main,img_name))
                img=cv2.resize(img,(100,100))
                img=cv2.cvtColor(img,cv2.COLOR_BGR2RGB)
                z.append([img,0])
        else:
            path_main=os.path.join('D:/dataset/test',cata)
            for img_name in os.listdir(path_main):
                img=cv2.imread(os.path.join(path_main,img_name))
                img=cv2.resize(img,(100,100))
                img=cv2.cvtColor(img,cv2.COLOR_BGR2RGB)
                z.append([img,1])
    print('Shuffling your data.....')
    shuffle(z)
    for images, labels in tqdm(z):
        X.append(images);Y.append(labels)
    return X,Y
X_val,Y_val=load_rottenvsfresh_valset()
Y_val=np.array(Y_val)
X_val=np.array(X_val)
y_ser=pd.Series(Y_val)
y_ser.value_counts()
import keras 
from keras.layers import Dense,Dropout, Conv2D,MaxPooling2D , Activation, Flatten, BatchNormalization, SeparableConv2D
from keras.models import Sequential
X=X/255.0
X_val=X_val/255.0
model.evaluate(X_val,Y_val)
model=load_model('D:/dataset/rottenvsfresh.h5')
from keras.models import Model, load_model
new_model=load_model('D:/dataset/rotten.h5')
new_model.evaluate(X_val,Y_val)
plt.imshow(X_val[0])
model.predict(X_val[0].reshape(1,100,100,3))
show_subpot(X_val[-36*11:-36*10])
#model.predict(X_val[-36*11:-36*10])
y_pred =model.predict(X_val[-36*11:-36*10]) 
np.round(y_pred).astype(int)

首次尝试的单张图片测试代码

# Load the image you want to classify
img_path = "C:/Users/d/Desktop/hjbjkh/lo.jpg"
img = load_image(img_path)

# Preprocess the image and prepare it for classification
img = np.array([img])  # Add a batch dimension

# Print a summary of the model's architecture
model.summary()

# Use the classification model to predict the class of the image
predictions = model.predict(img)

# Get the predicted class
predicted_class = np.argmax(predictions)

# If desired, display the image and the predicted class
print("Predicted class:", predicted_class)
plt.imshow(img[0])

修改后的尝试代码

# Import PyTorch
import torch

# Load the image you want to classify
img_path = "C:/Users/d/Desktop/hjbjkh/ad.jpg"
img = cv2.imread(img_path)

# Resize the image to the desired dimensions
img = cv2.resize(img, (100, 100))

# Convert the image from BGR color space to RGB color space
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)

# Convert the image from a NumPy array to a PyTorch tensor
img = torch.from_numpy(img)

# Add a batch dimension to the image
img = img.unsqueeze(0)

# Convert the image from a PyTorch tensor to a NumPy array
img = img.numpy()

# Use the classification model to predict the class of the image
predictions = model.predict(img)

# Get the predicted class
predicted_class = np.argmax(predictions)


# If desired, display the image and the predicted class
print("Predicted class:", predicted_class)
plt.imshow(img[0])

问题分析与正确解决方案

两次测试代码的核心问题

  1. 首次测试代码

    • 未定义load_image函数,直接调用会抛出未定义错误;
    • 缺少训练时的归一化步骤(训练集做了X=X/255.0,测试图未做);
    • 未确保图片尺寸与训练集一致(100x100)。
  2. 修改后代码

    • 不必要引入PyTorch转换,Keras模型仅接受NumPy数组输入,转换过程可能导致维度/数据类型异常;
    • 同样缺少归一化步骤;
    • 二分类模型用np.argmax无意义,模型输出是单个概率值(对应腐烂水果的概率)。

正确的单张图片测试代码

import cv2
import numpy as np
import matplotlib.pyplot as plt
from keras.models import load_model

# 加载训练好的模型
model = load_model('D:/dataset/rottenvsfresh.h5')

# 加载并预处理图片,完全复刻训练流程
img_path = "C:/Users/d/Desktop/hjbjkh/lo.jpg"
img = cv2.imread(img_path)
# 调整尺寸为训练时的100x100
img = cv2.resize(img, (100, 100))
# BGR转RGB,与训练时一致
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
# 归一化到0-1区间
img = img / 255.0
# 添加batch维度,转为模型要求的(1,100,100,3)格式
img = np.expand_dims(img, axis=0)

# 执行预测
predictions = model.predict(img)
# 二分类用0.5作为阈值判断,0对应新鲜,1对应腐烂
predicted_class = np.round(predictions).astype(int)[0][0]
class_name = "新鲜水果" if predicted_class == 0 else "腐烂水果"

print(f"预测类别: {class_name}")
plt.imshow(img[0])
plt.show()

关键注意事项

  • 测试图片的预处理流程必须完全匹配训练集:尺寸、颜色空间转换、归一化一个都不能少;
  • Keras模型输入需要batch维度,用np.expand_dims或reshape添加即可;
  • 二分类模型的输出是单个概率值,无需用np.argmax,直接用0.5阈值判断更准确。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 16:15:45