如何在Keras/TensorFlow中获取混淆矩阵各分类(TP、TN、FP、FN)的文件名
获取TP/TN/FP/FN对应的图像文件名
你不需要重新生成预测,因为你的test_generator已经设置了shuffle=False,测试集的文件顺序和标签、预测结果完全对应,直接用现有数据就能匹配出各类别的文件名。
步骤1:补全预测结果
先运行代码生成预测类别(你当前代码里的yy_pred未定义):
# 生成测试集预测概率 y_pred_probs = model.predict(test_generator) # 二分类任务,以0.5为阈值转换为类别标签 y_pred = (y_pred_probs > 0.5).astype(int).flatten()
步骤2:提取测试集文件路径
test_generator自带的filenames属性保存了所有测试图像的相对路径:
test_filenames = test_generator.filenames
步骤3:筛选各类别文件列表
结合真实标签y_true和预测标签y_pred,根据混淆矩阵的对应关系筛选:
你的混淆矩阵为[[22, 10],[9, 50]],行是真实标签、列是预测标签,对应:
- TN(真实0,预测0):22个
- FP(真实0,预测1):10个
- FN(真实1,预测0):9个
- TP(真实1,预测1):50个
执行以下代码得到各类文件列表:
# 筛选各类别文件 tn_files = [test_filenames[i] for i in range(len(y_true)) if y_true[i] == 0 and y_pred[i] == 0] fp_files = [test_filenames[i] for i in range(len(y_true)) if y_true[i] == 0 and y_pred[i] == 1] fn_files = [test_filenames[i] for i in range(len(y_true)) if y_true[i] == 1 and y_pred[i] == 0] tp_files = [test_filenames[i] for i in range(len(y_true)) if y_true[i] == 1 and y_pred[i] == 1]
步骤4:查看或保存结果
可以直接打印列表,或保存到文本文件方便后续视觉检查:
# 打印TP文件示例 print("True Positives 文件列表:") for f in tp_files: print(f) # 保存到本地文件 with open('tp_files.txt', 'w') as f: for file in tp_files: f.write(file + '\n') # 同理保存TN/FP/FN文件列表 with open('tn_files.txt', 'w') as f: for file in tn_files: f.write(file + '\n')
关键注意点
shuffle=False是核心前提,它保证了test_filenames的顺序和y_true、y_pred的顺序完全一致,不需要额外对齐操作。
内容的提问来源于stack exchange,提问作者temp
相关产品推荐
相关产品推荐

