How do you implement and interpret a confusion matrix in Python for multi-class classification?
A confusion matrix is an N×N table where entry [i][j] shows the number of samples with true class i predicted as class j. The diagonal represents correct predictions; off-diagonal entries are errors.
From the confusion matrix you can derive: precision, recall, F1 per class, and spot systematic misclassification patterns (e.g., model confuses 'cat' with 'dog').
import numpy as np
from sklearn.metrics import (
confusion_matrix, ConfusionMatrixDisplay,
classification_report
)
from sklearn.datasets import load_iris
from sklearn.ensemble import RandomForestClassifier
from sklearn.model_selection import train_test_split
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
# Multi-class classification
X, y = load_iris(return_X_y=True)
class_names = ['setosa', 'versicolor', 'virginica']
X_train, X_test, y_train, y_test = train_test_split(X, y, stratify=y, random_state=42)
model = RandomForestClassifier(n_estimators=100, random_state=42)
model.fit(X_train, y_train)
y_pred = model.predict(X_test)
# Confusion matrix
cm = confusion_matrix(y_test, y_pred)
print('Confusion Matrix:')
print(cm)
# Visual confusion matrix
fig, axes = plt.subplots(1, 2, figsize=(12, 5))
# Raw counts
disp = ConfusionMatrixDisplay(cm, display_labels=class_names)
disp.plot(ax=axes[0], colorbar=False)
axes[0].set_title('Confusion Matrix (Counts)')
# Normalized (row-normalized = recall per class)
cm_normalized = cm.astype(float) / cm.sum(axis=1, keepdims=True)
disp_norm = ConfusionMatrixDisplay(cm_normalized.round(2), display_labels=class_names)
disp_norm.plot(ax=axes[1], colorbar=False)
axes[1].set_title('Confusion Matrix (Recall-normalized)')
plt.tight_layout()
plt.savefig('confusion_matrix.png')
# Full classification report
print(classification_report(y_test, y_pred, target_names=class_names))
# Manual extraction from confusion matrix
for i, cls in enumerate(class_names):
tp = cm[i, i]
fp = cm[:, i].sum() - tp
fn = cm[i, :].sum() - tp
precision = tp / (tp + fp + 1e-10)
recall = tp / (tp + fn + 1e-10)
print(f'{cls}: Precision={precision:.2f}, Recall={recall:.2f}')