Visualisation with Matplotlib
Be able to draw line, scatter and histogram charts and read a training curve.
Prerequisites
Intuition
Choose the chart according to the question, not according to what looks nice:
| The question | Chart |
|---|---|
| How does something change over time? | a line chart |
| Are two variables related? | a scatter plot |
| How are the values distributed? | a histogram |
| How do the groups differ? | a box plot or a bar chart |
| Where are the errors? | a confusion matrix as a heat map |
Four rules that make charts readable:
- Label the axes, with units.
- Base bar charts at zero — otherwise the differences are exaggerated. Line charts may be zoomed.
- Always plot the data when there is little of it, not just the means.
- Do not use colour as the only distinction — around 8 % of men have some form of colour blindness. Add different line styles or markers.
Code
import matplotlib
matplotlib.use("Agg") # draw to a file without a screen — needed on a server
import matplotlib.pyplot as plt
import numpy as np
rng = np.random.default_rng(0)
epochs = np.arange(1, 41)
train = 1.2 * np.exp(-epochs / 12) + 0.05
val = 1.2 * np.exp(-epochs / 12) + 0.05 + np.maximum(0, (epochs - 18) * 0.006)
fig, ax = plt.subplots(1, 3, figsize=(13, 3.6))
# 1. The training curve — the most important picture in all of ML
ax[0].plot(epochs, train, label="training", linestyle="-")
ax[0].plot(epochs, val, label="validation", linestyle="--")
ax[0].axvline(int(epochs[np.argmin(val)]), color="gray", linestyle=":",
label="the best epoch")
ax[0].set(xlabel="epoch", ylabel="loss", title="Training curve")
ax[0].legend()
# 2. A scatter plot — are x and y related?
x = rng.normal(size=200)
y = 0.7 * x + rng.normal(scale=0.6, size=200)
ax[1].scatter(x, y, s=12, alpha=0.6)
ax[1].set(xlabel="feature", ylabel="target variable", title=f"r = {np.corrcoef(x, y)[0,1]:.2f}")
# 3. A histogram — what does the distribution look like?
ax[2].hist(rng.gamma(2.0, 2.0, size=1000), bins=30)
ax[2].set(xlabel="value", ylabel="count", title="Distribution (right-skewed)")
fig.tight_layout()
fig.savefig("figure.png", dpi=120)
Reading a training curve — four patterns and what they mean:
| Pattern | Diagnosis | Action |
|---|---|---|
| Both fall and level off together | healthy | done |
| Training falls, validation turns upwards | overfitting | early stopping, more data, regularisation |
| Both level off high | underfitting | a larger model, longer training, better features |
| The curve jumps wildly | too high a learning rate or too small a batch | lower the lr, increase the batch |
The second row is the most common. The point where the validation curve turns is the epoch you should save the model from — everything after that is learning the training data by heart.
Mastery means
- Draws the common chart types
- Reads a training curve
- Chooses the chart type according to the question
Sign in to do the exercises and build your mastery up.
Sources
- Matplotlib — dokumentation (PSF-liknande) — Matplotlib licence (BSD-like)
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0
- Statistics Sweden — about charts — myndighetsmaterial