tf.math.confusion_matrix生成混淆矩阵报predict_top_k属性错误求助
问题原因
- 报错核心是
predict_top_k并非Keras原生函数式(Functional)、序列(Sequential)模型自带的方法,该方法是你参考的教程中针对TensorFlow Hub发布的CropNet预训练模型做的专属封装,你自行搭建的二分类MobileNetV2模型没有实现该方法,因此触发属性错误。 - 原代码提取真实标签的写法存在隐患:直接对
validation_dataset.map()的结果转list,得到的是分batch的Tensor对象,没有做拼接展平,即使解决了预测方法的问题,传入tf.math.confusion_matrix()也会触发维度、类型不匹配错误。
修复代码
直接替换你原有混淆矩阵相关的代码片段即可,不需要修改模型结构和训练逻辑:
import seaborn as sns import numpy as np # 提取验证集真实标签,拼接为一维整数数组 y_true = [] for _, batch_labels in validation_dataset: y_true.append(batch_labels.numpy()) y_true = np.concatenate(y_true, axis=0).astype(np.int32) # 用模型原生predict方法推理,sigmoid输出转0/1类别标签 y_pred_prob = model2.predict(validation_dataset, batch_size=BATCH_SIZE).flatten() y_pred = (y_pred_prob >= 0.5).astype(np.int32) # 二分类阈值取0.5,可根据需求调整 def show_confusion_matrix(cm, labels): plt.figure(figsize=(10, 8)) sns.heatmap(cm, xticklabels=labels, yticklabels=labels, annot=True, fmt='g') plt.xlabel('Prediction') plt.ylabel('Label') plt.show() # 计算并绘制混淆矩阵 confusion_mtx = tf.math.confusion_matrix( y_true, y_pred, num_classes=len(class_names) ) show_confusion_matrix(confusion_mtx, class_names)
注意事项
- 由于
validation_dataset加载时设置了shuffle,上述代码遍历数据集取真实标签、传入模型做预测的顺序完全一致,不会出现预测结果和真实标签错位的问题。 - 你的模型是单神经元+sigmoid激活的二分类结构,输出值为样本属于正类的概率,默认以0.5为阈值判定类别,如果你的数据集存在类别不均衡问题,可以根据业务需求调整阈值大小。
- 不需要保留原代码中
rev_label_names相关的字典映射逻辑,image_dataset_from_directory生成的标签本身就是0、1的整数索引,和class_names的顺序一一对应,不需要额外做映射转换。
内容的提问来源于stack exchange,提问作者aguycalledankit
相关产品推荐
相关产品推荐

