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

如何在TensorFlow构建的Python模型中生成混淆矩阵?

生成预测值vs真实值的混淆矩阵实现方案

嘿,看了你的代码,你已经训练完二分类RNN模型并得到了预测结果yhat,接下来要生成混淆矩阵其实很简单,主要是处理好预测值的格式转换和真实值的维度对齐这两个关键点,我一步步给你讲:

关键步骤说明

你的模型用了sigmoid输出层和binary_crossentropy损失,属于二分类任务,所以:

  • 模型输出的yhat是0到1之间的概率值,需要先转换成0/1的类别标签
  • 真实值y_test目前是形状为(-1, 1)的二维数组,需要展平成一维数组,和预测值的维度匹配

用TensorFlow实现混淆矩阵

直接在你现有的代码后面追加以下代码即可:

import tensorflow as tf

# 1. 把预测概率转换成类别标签(以0.5为阈值)
# 方法1:用tf.round自动四舍五入
y_pred = tf.round(yhat).numpy().flatten()
# 方法2:手动设置阈值(更灵活,比如可以调整为0.6)
# y_pred = (yhat >= 0.5).astype(int).flatten()

# 2. 处理真实值:展平成一维数组
y_true = y_test.flatten()

# 3. 生成混淆矩阵
confusion_matrix = tf.math.confusion_matrix(
    labels=y_true,
    predictions=y_pred,
    num_classes=2  # 二分类任务,类别数设为2
)

# 打印混淆矩阵
print("混淆矩阵:")
print(confusion_matrix.numpy())

代码解释

  • tf.round(yhat):把0.5以上的概率转为1,以下转为0,刚好符合二分类的类别划分
  • .flatten():把二维数组转成一维,解决y_test和yhat的维度不匹配问题
  • tf.math.confusion_matrix的输出是一个num_classes × num_classes的矩阵,行代表真实类别,列代表预测类别:
    • 左上角是真阴性(TN):真实0,预测0
    • 右上角是假阳性(FP):真实0,预测1
    • 左下角是假阴性(FN):真实1,预测0
    • 右下角是真阳性(TP):真实1,预测1

可选:用Scikit-learn生成更直观的混淆矩阵

如果你想要更易读的输出(比如带原始标签和可视化效果),也可以用sklearn.metrics.confusion_matrix,代码如下:

from sklearn.metrics import confusion_matrix
import seaborn as sns
import matplotlib.pyplot as plt

# 生成混淆矩阵
cm = confusion_matrix(y_true, y_pred)

# 可视化混淆矩阵(还原原始类别标签)
sns.heatmap(cm, annot=True, fmt="d", cmap="Blues", 
            xticklabels=le.inverse_transform([0,1]),
            yticklabels=le.inverse_transform([0,1]))
plt.xlabel("预测类别")
plt.ylabel("真实类别")
plt.title("混淆矩阵")
plt.show()

这里用le.inverse_transform把编码后的0/1还原成你原始的标签文本,可视化后更直观。

常见坑点提醒

  • 不要直接把yhat(概率值)传入tf.math.confusion_matrix,必须先转成类别标签
  • 确保y_true和y_pred的维度一致,都是一维数组

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 07:43:49