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

深度学习代码如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 18:36:51