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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 06:48:24