深度学习代码如何在IDE运行控制台输出预测错误的样本
实现方法
你的代码已经拿到了测试集真实标签y_true和预测标签y_pred,核心逻辑是先定位两个标签数组不一致的索引,再把索引映射到对应样本路径、类别名输出即可,不需要依赖可读性差的大尺寸混淆矩阵。
修改步骤
1. 补充测试集样本路径记录
你当前用的image_dataset_loader.load仅返回图像数组和标签,没有保留样本对应文件路径,需要在加载数据后手动生成和x_test索引一一对应的测试文件路径列表,保证后续输出能定位到具体错判的图片:
# 放在 (x_train, y_train), (x_test, y_test) = load(path, [imgPath, testPath]) 代码之后 test_file_paths = [] # 类别排序规则和后续CLASS_NAMES保持一致 sorted_test_classes = sorted([entry.name for entry in os.scandir(testPath) if entry.is_dir()]) for cls_name in sorted_test_classes: cls_full_path = os.path.join(testPath, cls_name) # 文件名排序规则和loader读取逻辑对齐,保证索引匹配 cls_img_list = sorted([ os.path.join(cls_full_path, filename) for filename in os.listdir(cls_full_path) if filename.lower().endswith(('.jpg', '.jpeg', '.png', '.bmp')) ]) test_file_paths.extend(cls_img_list)
注意:如果运行后发现路径和错判样本不匹配,说明
image_dataset_loader内部读取排序逻辑和手动遍历的规则不一致,建议替换为手动遍历读取测试集的实现,从根源保证图像、标签、路径三者索引完全对齐。
2. 修正新版Keras接口兼容问题
如果你用的是TensorFlow 2.6+版本,predict_classes接口已经被移除,把预测代码替换为以下写法避免报错:
# 替换原来的 y_pred = (AlexNet.predict_classes(x_test)) y_pred = np.argmax(AlexNet.predict(x_test, verbose=0), axis=1) y_true = np.argmax(y_test, axis=1)
3. 定位所有预测错误的样本
在class_names定义完成后,通过数组比较直接筛选所有预测值和真实值不相等的样本索引:
# 筛选错判样本索引 error_sample_indices = np.where(y_pred != y_true)[0] print(f"\n===== 预测结果统计 =====") print(f"测试集总样本数:{len(x_test)}") print(f"预测正确样本数:{len(x_test) - len(error_sample_indices)}") print(f"预测错误样本数:{len(error_sample_indices)}") print(f"整体错误率:{round(len(error_sample_indices)/len(x_test)*100, 2)}%\n")
4. 终端输出所有错例详情
遍历错例索引,逐行输出样本路径、真实类别、预测类别、FP/FN属性:
print("===== 预测错误样本详情 =====") for idx in error_sample_indices: true_cls = class_names[y_true[idx]] pred_cls = class_names[y_pred[idx]] img_path = test_file_paths[idx] print(f"样本路径:{img_path}") print(f"真实类别:{true_cls} | 预测类别:{pred_cls}") print(f"错误属性:对类别[{true_cls}]为假阴性(FN),对类别[{pred_cls}]为假阳性(FP)\n")
可选扩展功能
如果需要快速看各类别的错判分布,不需要画混淆矩阵,可以加一段分类别FP/FN统计代码:
from collections import defaultdict fn_stat = defaultdict(int) fp_stat = defaultdict(int) # 遍历错例统计 for idx in error_sample_indices: true_cls = class_names[y_true[idx]] pred_cls = class_names[y_pred[idx]] fn_stat[true_cls] += 1 fp_stat[pred_cls] += 1 print("===== 各类别错判统计 =====") for cls in class_names: print(f"类别:{cls:20} | 假阴性(FN)数:{fn_stat.get(cls, 0):3} | 假阳性(FP)数:{fp_stat.get(cls, 0):3}")
如果测试集错例太多导致终端输出刷屏,可以加个输出条数限制,或者把所有错例信息写入本地txt文件留存。
内容的提问来源于stack exchange,提问作者Josh
相关产品推荐
相关产品推荐

