基于tensorflow-for-poets-2构建MobileNet后无法生成全类别混淆矩阵求助
解决MobileNet多分类模型混淆矩阵生成问题
我之前在做类似的TensorFlow for Poets项目时也碰到过一模一样的问题,核心卡点就是测试集图像的真实标签和模型预测标签没法一一对应,给你拆解几个关键步骤来解决:
1. 先理清测试集的标签映射关系
你的retrained_labels.txt里的每行内容就是模型输出的类别顺序,比如第一行是"cat",对应模型预测结果的索引0;第二行是"dog",对应索引1,以此类推。首先要把测试集里每个图像的真实类别(从它所在的文件夹名称来)转换成对应的索引:
- 假设你的测试集是按类别分文件夹存放的(和训练集结构一致),比如
test_data/cat/xxx.jpg、test_data/dog/xxx.jpg,那先给每个类别名称建立到索引的映射:with open('retrained_labels.txt', 'r') as f: labels = [line.strip() for line in f.readlines()] label_to_idx = {label: idx for idx, label in enumerate(labels)}
2. 批量获取测试集的真实标签与预测标签
接下来要遍历所有测试图像,一边记录真实标签的索引,一边用训练好的模型做预测得到预测标签索引:
import tensorflow as tf import os # 加载预训练好的模型 graph = tf.Graph() with graph.as_default(): graph_def = tf.compat.v1.GraphDef() with tf.io.gfile.GFile('retrained_graph.pb', 'rb') as f: graph_def.ParseFromString(f.read()) tf.import_graph_def(graph_def, name='') true_labels = [] pred_labels = [] test_dir = "你的测试集根目录路径" # 替换成实际路径 # 遍历每个类别文件夹 for class_name in os.listdir(test_dir): class_path = os.path.join(test_dir, class_name) if not os.path.isdir(class_path): continue # 获取当前类别的真实索引 true_idx = label_to_idx[class_name] # 遍历文件夹下的所有图像 for img_file in os.listdir(class_path): img_path = os.path.join(class_path, img_file) # 读取图像并做和训练时一致的预处理(比如resize、归一化,参考retrain.py里的逻辑) with tf.compat.v1.Session(graph=graph) as sess: softmax_tensor = sess.graph.get_tensor_by_name('final_result:0') # 运行模型得到预测结果 predictions = sess.run(softmax_tensor, {'DecodeJpeg/contents:0': tf.io.gfile.GFile(img_path, 'rb').read()}) # 获取预测概率最高的类别索引 pred_idx = predictions[0].argmax() # 存入列表 true_labels.append(true_idx) pred_labels.append(pred_idx)
3. 生成并可视化混淆矩阵
有了真实标签和预测标签的列表后,不管用TensorFlow还是Scikit-learn都能轻松生成混淆矩阵:
用TensorFlow原生方法
confusion_matrix = tf.math.confusion_matrix(true_labels, pred_labels, num_classes=len(labels)) with tf.compat.v1.Session() as sess: print("混淆矩阵:") print(sess.run(confusion_matrix))
用Scikit-learn+Seaborn可视化(更直观)
from sklearn.metrics import confusion_matrix import seaborn as sns import matplotlib.pyplot as plt cm = confusion_matrix(true_labels, pred_labels) # 绘制热力图 sns.heatmap(cm, annot=True, fmt='d', xticklabels=labels, yticklabels=labels) plt.title('Confusion Matrix') plt.xlabel('Predicted Class') plt.ylabel('True Class') plt.show()
为什么之前的方法无效?
evaluate.py脚本通常需要你提前准备好测试集的标签文件(比如每个图像对应一行真实标签),如果你的测试集没有这种结构化的标签映射,脚本就没法关联真实标签和预测结果。tf.confusion_matrix本身只是计算工具,核心是你得先把成对的真实标签和预测标签喂给它,之前没成功就是因为缺少了这一步数据准备。
内容的提问来源于stack exchange,提问作者mike pythonic
相关产品推荐
相关产品推荐

