Fixing Confusion Matrix plot lines

Viewed 1890

I am trying to plot a confusion matrix as shown below

cm  = confusion_matrix(testY.argmax(axis=1), predictions.argmax(axis=1))

disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=lb.classes_)
disp = disp.plot(include_values=True, cmap='viridis', ax=None, xticks_rotation='horizontal')

plt.show()

The result:

Confusion Matrix I get

As you can see, it's showing the axes of the boxes instead of outlining the boxes. I can't see the numbers outside the yellow boxes, because of the axes. I am not good with plots. So I can't find out what I need to change.

What I expect: Expected Matrix

FOUND SOLUTION

plt.tick_params(axis=u'both', which=u'both',length=0)
plt.grid(b=None)
4 Answers

Turn the grid off

E.g.,

import matplotlib.pyplot as plt
fig, _ = plt.subplots(nrows=1, figsize=(10,10))
ax = plt.subplot(1, 1, 1)
ax.grid(False)

...

disp = ConfusionMatrixDisplay(...)
_ = disp.plot(..., ax=ax, ...)
cm  = confusion_matrix(testY.argmax(axis=1), predictions.argmax(axis=1))

disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=lb.classes_)
disp = disp.plot(include_values=True, cmap='viridis', ax=None, xticks_rotation='horizontal')
plt.grid(False)
plt.show()

Change your cmap parameter in plot() function. It stands for colour-mapping your integer values with colors.

Check

https://matplotlib.org/3.1.0/tutorials/colors/colormaps.html

for more details.

As the answer

cm  = confusion_matrix(testY.argmax(axis=1), predictions.argmax(axis=1))

disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=lb.classes_)
disp = disp.plot(include_values=True, cmap='Blues', ax=None, xticks_rotation='horizontal')

plt.show()

The graph which you are showing as example is by sns plot. You can use sns heatmap to plot your matrix.

import seaborn as sns
categories = lb.classes_
sns.heatmap(cm, annot=True,categories =categories, cmap='Blues')
Related