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

CNN苹果分类模型训练精度高但测试精度低的问题求助

苹果品种CNN分类:训练精度高但测试精度极低问题排查

问题概述

使用CNN进行苹果4品种(braeburn、red_apples、red_delicious、rotten)分类时,训练数据精度很高,但测试数据精度仅约29%,数据按80:20划分,怀疑存在过拟合或代码/数据层面的问题。

数据集结构:

  • 两个根文件夹:TrainingData、TestData
  • 每个根文件夹下各有4个子文件夹,对应4个类别,存放对应苹果图片

代码与关键问题分析

以下是用户提供的代码及核心问题点:

# 错误点1:训练和测试数据集路径指向同一文件夹,导致训练测试数据重叠
TRAIN_DIR = 'apple_fruit'
TEST_DIR = 'apple_fruit'
classes = ['braeburn','red_apples','red_delicious','rotten'] 

train_datagen = ImageDataGenerator(rescale = 1./255, shear_range=0.2, zoom_range=0.2, horizontal_flip=True, fill_mode='nearest')
test_datagen = ImageDataGenerator(rescale = 1./255) 

training_set = train_datagen.flow_from_directory(TRAIN_DIR,
shuffle=True,
target_size = (100,100),
batch_size = 25,
classes =['braeburn','red_apples','red_delicious','rotten'])

test_set= test_datagen.flow_from_directory(TEST_DIR,
target_size = (100, 100),
shuffle=True,
 batch_size = 25,classes = classes)

model =Sequential()
# 错误点2:卷积层过滤器数量断崖式下降,特征提取能力被削弱
model.add(Conv2D(filters=128, kernel_size=(3,3),input_shape=(100,100,3), activation='relu', padding='same'))
model.add(MaxPooling2D(pool_size=(2, 2)))
model.add(Conv2D(filters=16, kernel_size=(3,3), activation='relu', padding = 'same'))
model.add(MaxPooling2D(pool_size=(2, 2)))
model.add(Flatten())
model.add(Dense(256))
model.add(Activation('relu'))
model.add(Dropout(0.6))
model.add(Dense(4,activation='softmax'))

model.compile(optimizer ='adam', loss = 'categorical_crossentropy', metrics = ['accuracy'])

# 错误点3:未加入验证集监控,无法判断过拟合趋势
history = model.fit(x=training_set,
steps_per_epoch=len(training_set),
epochs =10)

model.save('Ripe2_model6.h5')

loaded_model = keras.models.load_model("Ripe2_model6.h5")
predictions = model.predict(x=test_set, steps=len(test_set), verbose=True)
# 错误点4:对softmax输出使用round处理,破坏概率分布,导致预测错误
pred = np.round(predictions)

y_true=test_set.classes
y_pred=np.argmax(pred, axis=-1)
cm = confusion_matrix(y_true=test_set.classes, y_pred=y_pred)

# 混淆矩阵绘图函数(略)
plot_confusion_matrix(cm=cm, classes=classes, title='Confusion Matrix')

print(accuracy_score(y_true, y_pred))
print(recall_score(y_true, y_pred, average=None))
print(precision_score(y_true, y_pred, average=None))

核心错误总结

  1. 数据集路径错误:训练和测试目录均指向apple_fruit,导致模型在同一批数据上训练和测试,训练精度高但测试结果无意义;当切换到真正的独立测试集时,模型因未见过全新数据,精度暴跌。
  2. 预测结果处理错误:softmax输出是类别概率分布,使用np.round()会将概率低于0.5的值置0、高于0.5置1,可能出现多个1或全0的情况,导致argmax无法正确识别类别。
  3. 模型结构不合理:卷积层过滤器数量从128骤降到16,特征提取能力断层,无法有效提取层次化的图像特征。
  4. 缺乏过拟合监控:训练时未加入验证集,无法观察训练/验证精度的变化趋势,无法判断是否真的过拟合。

当前模型评估结果

  • 测试集准确率:0.2909
  • 召回率:[0.23484848 0.32319392 0.15151515 0.36213992]
  • 精确率:[0.23308271 0.32319392 0.15151515 0.36363636]

修复与优化建议

  1. 修正数据集路径
    将TRAIN_DIR改为'TrainingData',TEST_DIR改为'TestData',确保训练和测试数据完全独立。

  2. 修正预测逻辑
    移除np.round(predictions),直接对原始softmax输出取argmax:

    y_pred = np.argmax(predictions, axis=-1)
    
  3. 优化模型结构
    调整卷积层过滤器数量,保持逐步递减,增加卷积层数增强特征提取:

    model = Sequential()
    model.add(Conv2D(filters=64, kernel_size=(3,3), input_shape=(100,100,3), activation='relu', padding='same'))
    model.add(MaxPooling2D(pool_size=(2, 2)))
    model.add(Conv2D(filters=32, kernel_size=(3,3), activation='relu', padding='same'))
    model.add(MaxPooling2D(pool_size=(2, 2)))
    model.add(Conv2D(filters=16, kernel_size=(3,3), activation='relu', padding='same'))
    model.add(MaxPooling2D(pool_size=(2, 2)))
    model.add(Flatten())
    model.add(Dense(128, activation='relu'))
    model.add(Dropout(0.4))  # 调整dropout比例,平衡特征保留与过拟合抑制
    model.add(Dense(4, activation='softmax'))
    
  4. 加入过拟合监控
    训练时加入验证集,并绘制训练/验证曲线,直观判断过拟合情况:

    history = model.fit(x=training_set,
                        steps_per_epoch=len(training_set),
                        validation_data=test_set,
                        validation_steps=len(test_set),
                        epochs=20)
    # 绘制精度曲线
    plt.plot(history.history['accuracy'], label='Training Accuracy')
    plt.plot(history.history['val_accuracy'], label='Validation Accuracy')
    plt.legend()
    plt.show()
    

    若训练精度持续上升但验证精度下降,说明存在过拟合,可进一步增加dropout、加入BatchNormalization,或减少模型参数。

  5. 其他优化方向

    • 检查训练/测试集的类别分布,若存在类别不平衡,可使用类别权重(class_weight)或数据增强补充样本。
    • 使用EarlyStopping回调函数,当验证loss不再下降时自动停止训练,避免无效迭代:
      from keras.callbacks import EarlyStopping
      early_stop = EarlyStopping(monitor='val_loss', patience=3, restore_best_weights=True)
      history = model.fit(..., callbacks=[early_stop])
      
    • 尝试学习率调整,比如使用ReduceLROnPlateau回调函数动态调整学习率,帮助模型收敛到更优解。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 16:50:23