基于VGG的Keras猫狗分类模型验证准确率高但预测效果差的问题排查
嘿,这绝对是Keras数据生成器的经典“坑”!你看到的97%验证准确率是真实的,但混淆矩阵和分类报告完全不对,核心原因是生成器的状态错位导致预测结果和真实标签的顺序不匹配。
为什么会这样?
你先调用了model.evaluate_generator(),这个方法会把validation_generator里的所有验证数据从头到尾遍历一遍——遍历完成后,生成器的内部指针会停在最后一条数据的位置。这时候你再调用model.predict_generator(),生成器会从头开始新一轮的遍历,而你用validation_generator.classes来做真实标签,这个列表是按文件夹原始顺序排列的,和新一轮预测的结果顺序完全对不上,自然就出现了类似随机猜测的混淆矩阵。
另外还有个容易忽略的细节:Keras的flow_from_directory默认shuffle=True,哪怕你重置生成器,shuffle后的标签顺序也和classes列表不一致,这会让问题雪上加霜。
修复步骤
1. 固定验证生成器的顺序
创建验证生成器时,一定要设置shuffle=False,保证生成器输出的数据顺序和validation_generator.classes完全一致:
validation_generator = test_datagen.flow_from_directory( validation_data_dir, target_size=(img_height, img_width), batch_size=batch_size, class_mode='categorical', shuffle=False) # 关键!验证集不要打乱顺序
2. 重置生成器状态再预测
在调用predict_generator之前,先重置生成器的指针,让它回到数据的起始位置:
# 先重置生成器,避免之前的evaluate_generator导致指针错位 validation_generator.reset() # 然后进行预测 Y_pred = model.predict_generator(validation_generator, nb_validation_samples // batch_size) y_pred = np.argmax(Y_pred, axis=1) # 现在真实标签和预测结果的顺序就完全匹配了 print('Confusion Matrix') print(confusion_matrix(validation_generator.classes, y_pred)) print('Classification Report') target_names = ['Cats', 'Dogs'] print(classification_report(validation_generator.classes, y_pred, target_names=target_names))
3. (可选)调换评估和预测的顺序
如果你不想手动重置生成器,可以把predict_generator放在evaluate_generator前面——这样预测的时候生成器是从头开始的,评估的时候再遍历一遍,顺序也不会乱:
# 先预测 validation_generator.reset() # 保险起见还是重置一下 Y_pred = model.predict_generator(validation_generator, nb_validation_samples // batch_size) y_pred = np.argmax(Y_pred, axis=1) # 再评估 scores=model.evaluate_generator(generator=validation_generator, steps=nb_validation_samples // batch_size) print("Accuracy = ", scores[1]) # 输出混淆矩阵和报告 print('Confusion Matrix') print(confusion_matrix(validation_generator.classes, y_pred)) print('Classification Report') target_names = ['Cats', 'Dogs'] print(classification_report(validation_generator.classes, y_pred, target_names=target_names))
额外小优化
你的输出层用了Dense(2, activation='sigmoid'),但class_mode='categorical'搭配softmax激活会更匹配(或者把class_mode改成binary,输出层用Dense(1, activation='sigmoid')),不过这不是当前问题的核心,只是个小建议。
按照上面的步骤修改后,你就能得到和97%准确率匹配的混淆矩阵和分类报告啦!
内容的提问来源于stack exchange,提问作者rana hd

