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

CNN水果分类始终预测牛油果的常见问题排查

水果分类CNN始终预测同一类别的问题分析

我搭建了一个用于10种水果分类的Convolutional Neural Network(CNN),该模型在测试集(x_test)上实现了100%的准确率,但输入任意自定义图像(如空白屏幕)时,始终预测为牛油果。我发现将图像转为数组后显示会偏蓝,且执行img_data = img_data.astype('float32')后图像完全无法识别。训练集中每种水果的图片数量一致,相关代码如下:

import os
import re
import cv2
import numpy as np
from keras.utils import np_utils
from random import shuffle

def sorted_correctly(add = list):
    try:
        add.remove(".DS_Store")
    except:
        pass
    finally:
        convert = lambda text: int(text) if text.isdigit() else text
        aplhanum_key = lambda key: [convert(c) for c in re.split("([0-9]+)",key)]
        return sorted(add ,key=aplhanum_key)

names = ["apple","avocado","banana","cherry","kiwi","mango","orange","pineapple","strawberries","watermelon"]
num_classes = len(names)

def create_data(Dir):
    global names
    img_data_list = []
    fruits = sorted_correctly(os.listdir(Dir))
    labels = []
    unexpected = []
    for fruit in fruits:
        img_list = sorted_correctly(os.listdir(Dir + "/" + fruit))
        for img in img_list:
            try:
                input_img = cv2.imread(Dir + "/" +  fruit + "/" + img)
                input_img_resize = cv2.resize(input_img,(IMG_COLS,IMG_ROWS))
            except Exception as e:
                unexpected.append(img)
            else:
                img_data_list.append(input_img_resize)
                labels.append(names.index(fruit))
       
    img_data = np.array(img_data_list)
    img_data = img_data.astype('float32')
    img_data /= 255
    labels = np.ones((len(img_data)),dtype='int64') #delete if problems arise
    y = np_utils.to_categorical(labels,len(names)) 
    return img_data,y

def create_prediction_data(Dir):
    img_list = sorted_correctly(os.listdir(Dir))
    img_data_list = []
    for image in img_list:
        try: 
            input_img = cv2.imread(Dir +"/" + image)
            input_img_resize = cv2.resize(input_img,(IMG_COLS,IMG_ROWS))
        except:
            pass
        else:
            img_data_list.append(input_img_resize)
    img_data = np.array(img_data_list)
    imgs = img_data
    img_data = img_data.astype("float32")
    img_data /= 255
    return img_data,imgs


def randomize(x,y):  
    temp = list(zip(x,y))
    shuffle(temp)
    x,y = zip(*temp)
    return x,y

def unison_shuffled_copies(a, b):
    assert len(a) == len(b)
    p = np.random.permutation(len(a))
    return a[p], b[p]

create_data遍历本地目录生成numpy数组,create_prediction_data功能类似。

from keras.models import Sequential
from keras.layers import Conv2D, MaxPooling2D, Flatten, Dense, Dropout
from PIL import Image

print("Train:")
print(x_train.shape)
print("Test:")
print(x_test.shape)
model = Sequential()
model.add(Conv2D(32,(3,3),activation='relu',input_shape=(128,128,3)))
model.add(MaxPooling2D((2,2)))
model.add(Conv2D(32,(3,3),activation="relu"))
model.add(MaxPooling2D((2,2)))
model.add(Conv2D(64,(3,3),activation='relu'))
model.add(MaxPooling2D(pool_size=(2,2)))

model.add(Flatten())
model.add(Dense(128,activation="relu"))
model.add(Dropout(0.5))
model.add(Dense(len(names),activation="softmax"))

model.compile(optimizer='adadelta',
              loss='categorical_crossentropy',
              metrics=['accuracy'])
history = model.fit(x_train,y_train,epochs=5,batch_size=8,validation_data=(x_test,y_test))

loss_train = history.history["loss"]
loss_val = history.history["val_loss"]

score = model.evaluate([x_test],y_test,batch_size=32)
print("Evaluated.")
print("Test lost:", score[0])
print("Test accuracy:", score[1])
prediction_images,m = create_prediction_data(PREDICTION_DIRECTION)
for i in range(len(prediction_images)):
    img = Image.fromarray(m[i],"RGB")
    img.show()
    img = prediction_images[i]
    img = np.reshape(img,(1,128,128,3))
    prediction = np.argmax(model.predict(img))
    print(names[prediction])

测试准确率为1。

可能的问题分析

1. 标签生成逻辑致命错误

在create_data函数中,你已经通过labels.append(names.index(fruit))生成了对应每个水果的正确标签,但后续直接用np.ones覆盖了所有标签:

labels = np.ones((len(img_data)),dtype='int64') #delete if problems arise

这会让所有样本的标签都变成1,对应names列表里的第二个元素牛油果。模型全程都在学习“所有输入都是牛油果”,测试集100%准确率只是因为测试集标签也被错误设置为全1,模型完全没有学习到任何分类特征,自然对任意输入都输出牛油果。

2. 图像通道格式不匹配

  • OpenCV的cv2.imread读取的图像是BGR通道顺序,但你用Image.fromarray(m[i],"RGB")显示时采用RGB格式,通道顺序不一致导致图像偏蓝。
  • 如果训练时模型接收的是BGR格式,而自定义测试图像未做通道转换(比如用其他工具读取为RGB),模型看到的特征会完全偏离训练时的分布,进一步导致预测异常。

3. 图像预处理的可视化误解

执行img_data = img_data.astype('float32')后图像无法识别,是因为后续执行了img_data /= 255,将像素值缩放到了[0,1]区间,而Image.fromarray需要的是0-255的uint8类型数据。你在create_prediction_data中返回的imgs是未缩放的原始uint8数组,所以能正常显示;而prediction_images是缩放后的float32数组,直接用于显示必然异常——这只是可视化问题,不是预测错误的核心原因,但会干扰问题排查。

4. 模型训练的有效性不足

你使用的adadelta优化器学习率调整相对保守,且仅训练了5个epochs。结合标签错误的问题,测试集的100%准确率完全没有参考价值。即使标签正确,这么少的训练轮次也可能导致模型欠拟合,无法泛化到自定义图像。

5. 自定义预测图像的预处理一致性问题

虽然create_prediction_data和create_data都做了resize和归一化,但需要确认:

  • 自定义图像的尺寸是否严格统一为(128,128)
  • 自定义图像的色彩空间、像素范围是否与训练集完全一致(比如是否将RGB转为BGR)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 03:25:17