如何在自有数据上运行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
相关产品推荐
相关产品推荐

