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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.27 03:29:37