Plotting with Matplotlib

Learn how to turn raw numbers, model predictions, and training logs into clear visual charts—from beginner line plots to advanced multi-panel AI dashboards.

25 minBeginnerCode Examples

The Core Thesis: You cannot fix what you cannot see. In Machine Learning, staring at thousands of raw numbers in a terminal will hide broken data, extreme outliers, and models that are memorizing instead of learning. Matplotlib is the visual microscope of Python—turning raw arrays into scatter plots, loss curves, and image grids so you can see exactly what your data and your AI model are doing.

  1. Beginner Foundations: Figure vs. Axes (The Canvas and the Frame)

When beginners first use Matplotlib, they often get confused because there are two ways to write code: the quick way (plt.plot()) and the professional way (fig, ax = plt.subplots()). Understanding one simple picture makes everything click:

1. The Figure (fig) — The Blank PageOuter Container

Think of fig as the entire picture frame or blank window. It controls the overall image size (e.g., figsize=(10, 5)), background color, and saving the final image to a file (fig.savefig("chart.png")).

Controls: Window size, DPI resolution, and saving files.

2. The Axes (ax) — The Actual ChartDrawing Area

Think of ax as an individual chart drawn inside the frame (with its own X-axis, Y-axis, title, and lines). One Figure can hold 11 chart or a grid of 66 charts side by side!

Controls: ax.plot(), ax.scatter(), ax.set_title(), ax.legend()

Why You Should Always Use fig, ax = plt.subplots(): Calling plt.plot() guesses which chart you want to draw on, which gets messy as soon as you have two charts on the same screen. Starting every plot with fig, ax = plt.subplots() gives you direct, predictable control over every chart!

  1. Beginner: The Four Essential Charts in Machine Learning

You do not need to memorize 5050 chart types. In Machine Learning and Data Science, four core plots handle 90%90\% of your daily work:

Chart TypeMatplotlib CommandWhat It ShowsPrimary AI Use Case
1. Line Plotax.plot(x, y)Continuous trends over time or stepsPlotting Training vs. Validation Loss across epochs
2. Scatter Plotax.scatter(x, y, c=labels)Relationship between 2 variables as dotsVisualizing data clusters, outliers, and regression lines
3. Histogramax.hist(data, bins=30)Shape and spread of a single featureChecking if features follow a Bell Curve or are skewed
4. Bar Chartax.bar(names, scores)Comparing categories side by sideRanking Feature Importance or comparing model accuracies
⚡ Knowledge Check

Before training a house-price model, you want to check if the price column has a normal bell-curve shape or if most houses are cheap with a few extreme multi-million-dollar outliers. Which Matplotlib plot should you use?

A) A Histogram: ax.hist(prices, bins=30)▼
✓ Correct!A histogram groups a single continuous column into value buckets (bins) and counts how many rows fall into each bucket, immediately revealing skewness and outliers.
B) A Line Plot: ax.plot(prices)▼
✕ Incorrect.A line plot connects row 00 to row 11 to row 22 in order, which creates a noisy, meaningless zigzag unless the rows are ordered over time.

  1. Intermediate: Multi-Panel Subplots and Diagnosing Model Training

When training a neural network, looking at a single chart is rarely enough. You want to see the Loss Curve on the left and the Accuracy Curve on the right at the same time. We create side-by-side grids using plt.subplots(nrows, ncols):

1. Side-by-Side Plots (1 Row, 2 Columns)

Returns axes as a 1D list of 22 charts (axes[0] is the left chart, axes[1] is the right chart):

fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))

2. 2D Grid of Plots (2 Rows, 2 Columns)

Returns axes as a 2D grid (like axes[0, 1] for top-right). Always add plt.tight_layout() so titles never overlap!

fig, axes = plt.subplots(2, 2, figsize=(10, 8))

How to Read a Training vs. Validation Loss Plot

Every AI engineer plots Training Loss and Validation Loss on the same axes to spot three classic model behaviors:

1. Underfitting

Both Training Loss and Validation Loss stay high and flat. The model is too simple or the learning rate is too small.

2. Sweet Spot (Good Fit)

Both Training Loss and Validation Loss decrease smoothly and level off close together at a low error value.

3. Overfitting

Training Loss keeps dropping toward 00, but Validation Loss curves back UP! Stop training at the lowest validation point (Early Stopping).

  1. Intermediate: Plotting Images, Heatmaps, and Colorbars (ax.imshow)

How do we visualize a 2D matrix—such as a grayscale image in Computer Vision, a Confusion Matrix in Classification, or a Transformer Attention Map? We use ax.imshow(matrix), which paints every number in a 2D array as a colored pixel:

1. Heatmaps & Confusion Matrices

Pass any 2D matrix and choose a perceptually uniform colormap like cmap="viridis" or cmap="Blues", then attach fig.colorbar(im) to show the number scale.

im = ax.imshow(attn_matrix, cmap="viridis")

2. The PyTorch Image Shape Rule!

PyTorch stores color images as Channels-First (3, H, W), but Matplotlib requires Channels-Last (H, W, 3)! Always permute axes before plotting:

img_hwc = tensor_chw.permute(1, 2, 0).cpu().numpy()

⚡ Knowledge Check

You have a single RGB image tensor in PyTorch on the GPU with shape (3, 224, 224). If you pass it directly into ax.imshow(img), Matplotlib crashes with TypeError: Invalid shape (3, 224, 224) for image data. How do you fix it?

A) Move to CPU and reorder axes to (224, 224, 3) using .permute(1, 2, 0)▼
✓ Correct!Matplotlib operates on CPU NumPy arrays and expects Height and Width first, followed by the 3 RGB color channels: (H, W, C) = (224, 224, 3).
B) Flatten the image into a 1D vector of length 150,528▼
✕ Incorrect.ax.imshow() requires a 2D grayscale grid (H, W) or a 3D color grid (H, W, 3) so it knows the spatial height and width of the image.

  1. Advanced: Decision Boundaries, Confidence Bands, and Vector Contours

At the Advanced ML level, we use Matplotlib not just to plot raw data, but to visualize the internal geometry and uncertainty of trained models:

1. Uncertainty Bands (ax.fill_between)

In Reinforcement Learning and deep learning experiments, we run 55 random seeds and plot the mean curve μ\mu with a shaded ±1σ\pm 1\sigma standard deviation band:

ax.fill_between(epochs, mean - std, mean + std, alpha=0.25)

2. 2D Decision Boundaries (ax.contourf)

By generating a 2D coordinate grid with np.meshgrid(x, y) and predicting class probabilities across every grid point, ax.contourf() colors the exact regions where a neural network switches from Class 0 to Class 1!

ax.contourf(XX, YY, Z_probs, levels=20, cmap="coolwarm", alpha=0.3)

Production Memory Rule (Preventing RAM Leaks in Training Loops): If your training script saves a chart at the end of every epoch inside a 1,000-epoch loop, Matplotlib keeps every figure open in RAM by default until the server crashes! Always call plt.close(fig) immediately after fig.savefig() to free the memory.

⚡ Knowledge Check

You plot your neural network's Training Loss and Validation Loss across 5050 epochs. From Epoch 11 to 1515, both curves drop smoothly. After Epoch 1515, Training Loss continues dropping toward 0.0010.001, while Validation Loss turns upward and climbs steeply. What is your chart telling you?

A) The model started overfitting at Epoch 15; use the Epoch 15 checkpoint▼
✓ Correct!Diverging loss curves—where training error falls while validation error rises—are the classic visual signature of overfitting. Marking Epoch 1515 with ax.axvline(15) highlights the exact Early Stopping point.
B) The model is underfitting and needs 100 more epochs▼
✕ Incorrect.Training for more epochs will only make the model memorize the training noise further and worsen validation performance.

  1. Visual Explanation: The Matplotlib Anatomy & AI Plotting Workflow

Look at how data flows from NumPy or PyTorch arrays into a structured Figure containing multiple Axes subplots, gets annotated with titles and uncertainty bands, and is exported cleanly without leaking memory:

  1. Python Implementation: Building a 2-Panel AI Training Dashboard

Here is a complete, runnable script that builds a professional 2-panel Machine Learning dashboard: Panel 1 plots Training vs. Validation Loss with an uncertainty band (fill_between) and an Early Stopping marker, while Panel 2 plots a labeled Confusion Matrix heatmap using imshow:

matplotlib_ai_dashboard.pyPython 3.11+ · Matplotlib & NumPy
import numpy as npimport matplotlib.pyplot as plt # 1. Simulate 20 Epochs of Training & Validation Lossepochs = np.arange(1, 21)trainLoss = 1.2 * np.exp(-0.22 * epochs) + 0.05valLoss = 1.2 * np.exp(-0.18 * epochs) + 0.02 * ((epochs - 12) ** 2) * (epochs > 12) + 0.12valStd = np.full_like(valLoss, 0.04)bestEpoch = int(epochs[np.argmin(valLoss)]) # 2. Create a 1x2 Multi-Panel Figure (Object-Oriented API)fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4.8), dpi=120) # PANEL 1: Loss Curves + Uncertainty Band + Early Stopping Lineax1.plot(epochs, trainLoss, label="Train Loss", color="#00F0FF", linewidth=2)ax1.plot(epochs, valLoss, label="Val Loss", color="#F43F5E", linewidth=2)ax1.fill_between(epochs, valLoss - valStd, valLoss + valStd, color="#F43F5E", alpha=0.18)ax1.axvline(bestEpoch, color="#10B981", linestyle="--", label=f"Best Epoch ()")ax1.set_title("Training vs. Validation Loss")ax1.set_xlabel("Epoch")ax1.set_ylabel("Cross-Entropy Loss")ax1.grid(True, alpha=0.25)ax1.legend() # PANEL 2: Confusion Matrix Heatmap with Annotated CountsconfMatrix = np.array([[92, 8], [11, 89]])im = ax2.imshow(confMatrix, cmap="Blues")fig.colorbar(im, ax=ax2, fraction=0.046, pad=0.04)ax2.set_title("Confusion Matrix")ax2.set_xticks([0, 1], labels=["Pred: Normal", "Pred: Spam"])ax2.set_yticks([0, 1], labels=["True: Normal", "True: Spam"]) # Overlay exact numbers inside each heatmap cellfor r in range(2):    for c in range(2):        ax2.text(c, r, str(confMatrix[r, c]), ha="center", va="center", fontweight="bold") # 3. Clean Layout, Save High-Res Image, and Close to Free RAMfig.tight_layout()print(f"Dashboard ready! Best validation loss at Epoch .")plt.close(fig)

Pro Tip (Exporting Sharp Charts for Papers & Dashboards): Always pass dpi=300 and bbox_inches="tight" inside fig.savefig("plot.png", dpi=300, bbox_inches="tight") so axis labels are never clipped off the edge of the image!

Key Points

✓A Matplotlib Figure (fig) is the outer window or page, while an Axes (ax) is an individual chart inside that window. Always prefer fig, ax = plt.subplots() over global plt.plot() calls.
✓Use Line Plots (ax.plot) for training curves over time, Scatter Plots (ax.scatter) for 2D feature relationships, Histograms (ax.hist) for checking data distributions, and Bar Charts (ax.bar) for category comparisons.
✓Plotting Training Loss alongside Validation Loss immediately diagnoses Underfitting (both high) vs. Overfitting (validation loss curves upward while training loss drops).
✓Use ax.imshow() to visualize 2D matrices such as images, Confusion Matrices, and Transformer Attention weights, and convert PyTorch images from (C, H, W) to (H, W, C) first.
✓Use ax.fill_between() to shade confidence/variance bands across multiple training runs and ax.contourf() to plot 2D model decision boundaries.
✓Always call fig.tight_layout() to prevent overlapping labels and plt.close(fig) inside training loops to prevent memory leaks.

Common Mistakes

✕ Passing a PyTorch GPU tensor with requires_grad=True directly into Matplotlib.

Matplotlib only understands CPU NumPy arrays. Passing a GPU or autograd tensor crashes with a RuntimeError. Always call tensor.detach().cpu().numpy() before plotting.

✕ Plotting RGB images with Channels-First shape (3, H, W) instead of (H, W, 3).

PyTorch stores color channels in axis 00, whereas ax.imshow() requires color channels in the last axis. Reorder axes with img.permute(1, 2, 0) or np.transpose(img, (1, 2, 0)).

✕ Creating hundreds of figures inside a training loop without calling plt.close(fig).

Matplotlib retains references to every created Figure in memory until explicitly closed, which will steadily eat gigabytes of RAM during long training runs.

✕ Calling fig.savefig() AFTER calling plt.show(), resulting in a blank white image file.

plt.show() displays and then clears the active figure canvas. Always call fig.savefig("chart.png") before calling plt.show().

The Big Picture

Training Blind (Terminal Numbers Only)

Print Raw Loss Numbers → Miss Outliers & Overfitting → Deploy Broken Model

Visual AI Engineering (Matplotlib Diagnostics)

Inspect Feature Histograms → Track Train/Val Curves → Visualize Confusion & Attention Maps → Reliable Model

The important conceptual shift is treating visualization not as a pretty slide-deck decoration at the end of a project, but as your primary debugging tool at every stage—before training (to spot bad data), during training (to catch overfitting), and after training (to see where the model makes mistakes).

Remember: A single 5-line Matplotlib chart will often reveal in two seconds a bug that would take two days to find by staring at printed arrays in a terminal.