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

MNIST数据集上实现MLP后无法打印混淆矩阵求助

解决MLP模型的混淆矩阵打印问题

首先咱们得明确:混淆矩阵的核心是真实标签和模型预测的类别标签这两组数据。你的模型返回的logits是未经过softmax转换的原始得分,所以得先把它转成预测类别,再和真实标签计算混淆矩阵。下面一步步来:

步骤1:把Logits转成预测类别

用tf.argmax()就能从logits里提取每个样本的预测类别索引,这里要注意axis参数对应类别维度(MNIST是10类,logits形状应该是[样本数,10],所以axis=1):

# 假设你有测试集输入x_test,以及真实标签y_test
predictions = tf.argmax(logits, axis=1)
# 如果你的y_test是one-hot编码(比如形状是[10000,10]),也要转成类别索引
true_labels = tf.argmax(y_test, axis=1)

步骤2:计算混淆矩阵

TensorFlow自带的tf.math.confusion_matrix()可以直接帮你计算,只要传入真实标签和预测标签就行:

confusion_mat = tf.math.confusion_matrix(labels=true_labels, predictions=predictions)

步骤3:打印或可视化混淆矩阵

如果只是看数值,直接转成numpy数组打印就好:

print(confusion_mat.numpy())

要是想更直观看到分类效果,用matplotlib画热力图会清晰很多:

import matplotlib.pyplot as plt
import seaborn as sns

plt.figure(figsize=(10,8))
sns.heatmap(confusion_mat.numpy(), annot=True, fmt='d', cmap='Blues', 
            xticklabels=range(10), yticklabels=range(10))
plt.xlabel('Predicted Label')
plt.ylabel('True Label')
plt.title('MNIST Confusion Matrix')
plt.show()

容易踩坑的几个点

  • 标签格式不匹配:如果你的真实标签是one-hot编码,一定要转成类别索引,否则计算混淆矩阵时会报错。
  • 测试集预处理一致:测试集的归一化、形状调整必须和训练集完全一样,不然预测结果会出错,混淆矩阵也失去参考价值。
  • 大测试集分批处理:如果测试集数据量很大,别一次性喂入模型,分batch计算预测结果再合并,避免内存溢出:
predictions = []
# 按32个样本为一批处理测试集
for batch_x in tf.data.Dataset.from_tensor_slices(x_test).batch(32):
    batch_logits = layers(batch_x, weights, biases)
    batch_pred = tf.argmax(batch_logits, axis=1)
    predictions.append(batch_pred)
# 把所有batch的预测结果合并
predictions = tf.concat(predictions, axis=0)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 08:51:30