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

如何在自有数据上运行GoogleNet/Inception5H并获取列联表形式的模型准确率

实现GoogleNet/Inception5H模型准确率计算(列联表或其他形式)

一、明确脚本局限性

你运行的ace_run.py是TCAV框架下的概念激活分析工具,核心目标是验证概念对模型分类决策的影响,默认不包含分类准确率计算逻辑,所以需要手动添加代码实现该功能。

二、具体实现步骤

1. 加载模型与自有数据集

先实现模型和数据的加载逻辑,适配Inception5H的输入要求:

import tensorflow as tf
import numpy as np
import pandas as pd
import os
from PIL import Image
from sklearn.metrics import confusion_matrix, classification_report

# 加载Inception5H模型
def load_inception_model(model_path):
    with tf.io.gfile.GFile(model_path, 'rb') as f:
        graph_def = tf.compat.v1.GraphDef()
        graph_def.ParseFromString(f.read())
    tf.import_graph_def(graph_def, name='')
    sess = tf.compat.v1.Session()
    # 获取模型输入输出张量
    input_tensor = sess.graph.get_tensor_by_name('input:0')
    output_tensor = sess.graph.get_tensor_by_name('softmax:0')
    return sess, input_tensor, output_tensor

# 加载自有数据集(按文件夹分类存储的场景)
def load_custom_data(source_dir, labels_path):
    # 读取标签映射表
    with open(labels_path, 'r') as f:
        labels = [line.strip() for line in f.readlines()]
    label_to_idx = {label:i for i, label in enumerate(labels)}
    
    # 遍历文件夹收集图片路径与对应标签
    data_list = []
    for class_name in os.listdir(source_dir):
        class_dir = os.path.join(source_dir, class_name)
        if not os.path.isdir(class_dir):
            continue
        label_idx = label_to_idx.get(class_name, -1)
        if label_idx == -1:
            continue
        for img_name in os.listdir(class_dir):
            img_path = os.path.join(class_dir, img_name)
            data_list.append((img_path, label_idx))
    return data_list, labels

2. 预处理图片并获取模型预测结果

Inception5H要求输入为224×224的RGB图片,且需归一化到[-1,1]区间:

def preprocess_image(img_path):
    img = Image.open(img_path).resize((224, 224))
    img_array = np.array(img)
    # 归一化处理
    img_array = (img_array / 255.0) * 2 - 1
    return img_array.reshape(1, 224, 224, 3)

# 批量获取预测结果与真实标签
def get_model_predictions(sess, input_tensor, output_tensor, data_list):
    y_true = []
    y_pred = []
    for img_path, true_label in data_list:
        img_input = preprocess_image(img_path)
        pred_probs = sess.run(output_tensor, feed_dict={input_tensor: img_input})
        pred_label = np.argmax(pred_probs)
        y_true.append(true_label)
        y_pred.append(pred_label)
    return y_true, y_pred

3. 生成准确率结果(列联表/分类报告)

利用sklearn和pandas生成所需的结构化结果:

# 主执行逻辑
if __name__ == '__main__':
    # 替换为你的实际路径参数
    model_path = 'tcav/tcav_examples/image_models/imagenet/YOUR_FOLDER/inception5h/tensorflow_inception_graph.pb'
    source_dir = 'tcav/tcav_examples/image_models/imagenet/YOUR_FOLDER/'
    labels_path = './imagenet_labels.txt'

    # 加载模型与数据
    sess, input_tensor, output_tensor = load_inception_model(model_path)
    data_list, labels = load_custom_data(source_dir, labels_path)

    # 获取预测结果
    y_true, y_pred = get_model_predictions(sess, input_tensor, output_tensor, data_list)

    # 1. 生成列联表(混淆矩阵)
    conf_matrix = confusion_matrix(y_true, y_pred)
    conf_df = pd.DataFrame(conf_matrix, index=labels, columns=labels)
    print("分类列联表(混淆矩阵):")
    print(conf_df)
    # 保存到本地文件
    conf_df.to_csv('classification_confusion_matrix.csv')

    # 2. 生成详细分类报告(含准确率、召回率、F1值)
    class_report = classification_report(y_true, y_pred, target_names=labels)
    print("\n详细分类报告:")
    print(class_report)
    # 保存到本地文件
    with open('classification_report.txt', 'w') as f:
        f.write(class_report)

    # 3. 计算整体准确率
    overall_acc = np.mean(np.array(y_true) == np.array(y_pred))
    print(f"\n整体分类准确率:{overall_acc:.4f}")

4. 运行方式

  • 可以将上述代码保存为单独的calculate_accuracy.py,替换路径参数后直接运行;
  • 也可以将代码块整合到ace_run.py的末尾,在概念分析完成后自动执行准确率计算。

三、注意事项

  • 确保数据集的分类文件夹名称与imagenet_labels.txt中的标签完全匹配,否则会出现标签映射错误;
  • 若数据集规模较大,可修改代码实现批量图片处理,避免内存占用过高;
  • 若使用TensorFlow 2.x版本,需保持代码中的tf.compat.v1兼容模式,或转换模型为TF2格式。

内容的提问来源于stack exchange,提问作者Archer

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 14:10:34