plot_confusion_matrix

keras.plot_confusion_matrix(
    y_true
    y_pred
    *
    labels=None
    normalize=None
    title='Confusion matrix'
)

Draw a classifier’s confusion matrix as a heatmap.

The true class runs down the rows and the predicted class across the columns, so the diagonal is what the model got right. A reader moves across a row to hear where one class’s samples went.

Parameters

Name Type Description Default
y_true array_like The true classes: class indices, or one-hot rows. required
y_pred array_like The predicted classes: class indices, or what model.predict returns. Rows of class probabilities are read by their largest; a single probability per sample, as a sigmoid gives, is the positive class from 0.5 up. required
labels iterable of str Each class’s name, in index order. By default the class indices. None
normalize (None, 'true', 'pred', 'all') None keeps the counts. "true" divides each row by its total, so a cell is the share of that true class predicted as the column’s; "pred" divides each column; "all" divides by every sample. None
title str The chart’s title. "Confusion matrix"

Returns

Name Type Description
matplotlib.figure.Figure The heatmap, ready for :func:maidr.show, :func:maidr.render or :func:maidr.save_html. Not managed by pyplot.

Raises

Name Type Description
ValueError If y_true and y_pred hold a different number of samples, are empty, hold a negative class index or a value that is neither an index nor a probability, or normalize is not one of the above, or labels names fewer classes than the data holds. More labels than the data holds are allowed: the classes that never occur are shown as empty rows and columns.

Notes

The tutorial way to put a confusion matrix in TensorBoard is to log it as an image, which carries no numbers for a screen reader to read. Drawing it from the predictions keeps them. NumPy alone computes it.

Examples

>>> import maidr
>>> from maidr.keras import plot_confusion_matrix
>>> figure = plot_confusion_matrix(
...     y_test, model.predict(x_test), labels=["cat", "dog"], normalize="true"
... )
>>> maidr.show(figure)