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

如何绘制混淆矩阵与分类报告?无y_test时的解决方法

解决混淆矩阵与分类报告的绘制问题

先修正你的数据集划分代码

你当前的代码里,test_ds和val_ds都是直接复用train_ds做处理,等于三个数据集完全是同一批数据,根本没做划分。先把数据集拆分的逻辑改对:

# 假设你有原始的完整数据集 full_ds
full_ds = full_ds.shuffle(10000, seed=42)  # 先全局打乱数据

# 按比例划分:训练集80%、验证集10%、测试集10%
train_size = int(0.8 * len(full_ds))
val_size = int(0.1 * len(full_ds))
test_size = len(full_ds) - train_size - val_size

train_ds = full_ds.take(train_size)
remaining_ds = full_ds.skip(train_size)
val_ds = remaining_ds.take(val_size)
test_ds = remaining_ds.skip(val_size)

# 再做缓存、预取优化
train_ds = train_ds.cache().prefetch(buffer_size=tf.data.experimental.AUTOTUNE)
val_ds = val_ds.cache().prefetch(buffer_size=tf.data.experimental.AUTOTUNE)
test_ds = test_ds.cache().prefetch(buffer_size=tf.data.experimental.AUTOTUNE)

提取真实标签与预测标签

不管是混淆矩阵还是分类报告,都需要测试集的真实标签和模型的预测标签:

1. 提取测试集真实标签

import numpy as np

y_true = []
for _, labels in test_ds:
    y_true.extend(labels.numpy())
y_true = np.array(y_true)

2. 获取模型预测标签

# 假设你已经训练好模型 model
y_pred_probs = model.predict(test_ds)
y_pred = np.argmax(y_pred_probs, axis=1)  # 多分类取概率最大的类别索引,二分类可以用(y_pred_probs > 0.5).astype(int)

绘制混淆矩阵

用sklearn+matplotlib/seaborn可视化:

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

# 计算混淆矩阵
cm = confusion_matrix(y_true, y_pred)

# 可视化
plt.figure(figsize=(8, 6))
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', 
            xticklabels=class_names, yticklabels=class_names)  # class_names是你的类别名称列表,比如["猫", "狗"]
plt.xlabel('预测类别')
plt.ylabel('真实类别')
plt.title('混淆矩阵')
plt.show()

生成分类报告

直接用sklearn的工具生成:

from sklearn.metrics import classification_report

print(classification_report(y_true, y_pred, target_names=class_names))

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 22:31:18