基于kratzert的微调AlexNet代码,TensorFlow下混淆矩阵构建求助
如何在kratzert的AlexNet微调代码中构建混淆矩阵
我很熟悉你用的这个AlexNet微调框架的结构,咱们一步步来解决混淆矩阵的问题:
核心原理
tf.confusion_matrix() 只需要两个核心输入就能生成矩阵:
- labels:样本的真实类别标签(一维张量,每个元素对应样本的真实类别索引)
- predictions:模型输出的预测类别标签(和labels维度完全一致的一维张量,每个元素是模型判断的类别索引)
具体实现步骤
1. 准备存储数据的容器
在测试/验证代码段的最开始,初始化两个空列表,用来累积所有测试样本的真实标签和预测结果:
true_labels = [] pred_labels = []
2. 在测试循环中收集数据
找到代码里遍历测试数据集的循环(通常是验证/测试阶段的while或for循环),每次迭代时做两件事:
- 提取当前批量的真实标签:如果你的标签是one-hot编码格式,需要用
tf.argmax(labels_batch, axis=1)转换成类别索引;如果本来就是索引格式,直接取原始值即可。 - 从模型输出的logits中得到预测类别:用
tf.argmax(logits_batch, axis=1)提取概率最高的类别索引。
把这两个结果追加到之前的列表里:
# 处理真实标签(假设labels_batch是当前批量的真实标签) if labels_batch.shape[-1] > 1: # 判断是否是one-hot编码 batch_true = tf.argmax(labels_batch, axis=1).numpy() else: batch_true = labels_batch.numpy().flatten() true_labels.extend(batch_true) # 处理模型预测结果 batch_pred = tf.argmax(logits_batch, axis=1).numpy() pred_labels.extend(batch_pred)
3. 生成并查看混淆矩阵
等所有测试样本遍历完成后,把两个列表转换成张量,传入tf.confusion_matrix()即可:
import tensorflow as tf # 转换为符合要求的张量格式 true_tensor = tf.convert_to_tensor(true_labels, dtype=tf.int32) pred_tensor = tf.convert_to_tensor(pred_labels, dtype=tf.int32) # 替换YOUR_CLASS_COUNT为你数据集的实际类别数 confusion_mat = tf.confusion_matrix(labels=true_tensor, predictions=pred_tensor, num_classes=YOUR_CLASS_COUNT) # 打印矩阵,也可以转成numpy数组做后续处理 print("混淆矩阵:") print(confusion_mat.numpy())
4. 避坑提示
- 确保
labels和predictions的维度、数据类型完全一致(推荐用int32) - 如果是TensorFlow 2.x环境,记得用
.numpy()把张量转成numpy数组,方便后续操作 - 如果你的标签是从数据集读取的字符串或其他格式,要先映射成连续的整数索引
可选:可视化混淆矩阵
如果需要更直观的展示,可以用seaborn绘制热力图:
import seaborn as sns import matplotlib.pyplot as plt plt.figure(figsize=(10,8)) # 替换YOUR_CLASS_NAMES为你的类别名称列表,比如["cat", "dog", "bird"] sns.heatmap(confusion_mat.numpy(), annot=True, fmt='d', cmap='Blues', xticklabels=YOUR_CLASS_NAMES, yticklabels=YOUR_CLASS_NAMES) plt.xlabel('预测类别') plt.ylabel('真实类别') plt.show()
内容的提问来源于stack exchange,提问作者user2975921
相关产品推荐
相关产品推荐

