# plot the confusion matrix mat = confusion_matrix(test_data.target, predicted_categories) sns.heatmap(mat.T, square = True, annot=True, fmt = "d", xticklabels=train_data.target_names,yticklabels=train_data.target_names) plt.xlabel("true labels") plt.ylabel("predicted label") plt.show()