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

