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
相关产品推荐
相关产品推荐

