TensorFlow and Keras

Callbacks

Callbacks are hooks that let code run at set moments during training — to stop early, save the best model, or log progress — without touching the training loop.

On this page 5
  1. Why they exist
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

A callback is a piece of code Keras runs for you at fixed moments during training — after each batch, after each epoch, at the start, at the end.

Think of a cricket coach on the sidelines during practice. The coach does not bowl or bat. But after every over, the coach checks the score, notes it in a diary, and can call practice off when the player stops improving. Callbacks are that coach — watching fit from outside, allowed to record things and to intervene.

Why they exist

fit is a sealed loop. You start it, it runs for the epochs you asked. But real training needs mid-run decisions:

  • Stop early. If the model stopped improving 50 epochs ago, more epochs teach it to memorise — the overfitting trap.
  • Save the best version. The last epoch is rarely the best epoch. Someone must snapshot the good one.
  • Keep a log. When training runs for hours, you want a record to look at.

Without callbacks you would rewrite the training loop for each of these. With callbacks you keep fit and attach watchers to it.

How it works

fit starts
  ├─ epoch 1 ── batches ──→  callbacks: "epoch 1 done, loss was 0.9"
  ├─ epoch 2 ── batches ──→  callbacks: "epoch 2 done, loss was 0.5"  → save!
  ├─ epoch 3 ── batches ──→  callbacks: "epoch 3 done, loss was 0.6"  → worse…
  │        ⋮
  └─ callback says STOP ──→  fit ends early, best version kept

After every epoch, Keras hands each callback the latest numbers. A callback can write them down, save the model, or raise its hand to end training.

A real example you have seen

Your phone's battery charging works this way. You plug it in and a watcher checks the level as it charges. At 100 per cent, the watcher cuts the current — you never told it when to stop, the rule did. Early stopping is the same rule pointed at a learning curve: stop when full, where "full" means "no longer improving on unseen data".

Remember this

  • Callbacks run at fixed moments in training: epoch end, batch end, train start, train end.
  • The two you will use constantly: EarlyStopping and ModelCheckpoint.
  • They watch a monitored value, usually val_loss — the score on held-out data.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install tensorflow numpy

Outputs verified with TensorFlow 2.21 (Keras 3.12), CPU, seed 0. The exact stopping epoch shifts across versions and machines — expect the pattern, not the digits.

Early stopping, checkpoints, and a custom watcher

callbacks_demo.py
import numpy as np
import keras
from keras import layers

keras.utils.set_random_seed(0)

# 40 noisy points on a line: enough to overfit if we train forever
x = np.linspace(-1, 1, 40).reshape(-1, 1).astype("float32")
y = (2 * x + 0.3 * np.random.randn(40, 1)).astype("float32")

model = keras.Sequential([
    keras.Input(shape=(1,)),
    layers.Dense(32, activation="relu"),
    layers.Dense(1),
])
model.compile(optimizer="adam", loss="mse")


class EpochTicker(keras.callbacks.Callback):
    def on_epoch_end(self, epoch, logs=None):
        if (epoch + 1) % 100 == 0:
            print(f"epoch {epoch + 1}: val_loss {logs['val_loss']:.4f}")


stopper = keras.callbacks.EarlyStopping(
    monitor="val_loss",
    patience=50,                  # give up after 50 epochs with no improvement
    restore_best_weights=True,    # roll back to the best epoch, not the last
)
saver = keras.callbacks.ModelCheckpoint(
    "best.keras", monitor="val_loss", save_best_only=True)

history = model.fit(x, y, validation_split=0.25, epochs=1000,
                    verbose=0, callbacks=[EpochTicker(), stopper, saver])

print("stopped at epoch:", stopper.stopped_epoch + 1)
print("best val_loss:", round(min(history.history["val_loss"]), 4))
Output
epoch 100: val_loss 0.9117
epoch 200: val_loss 0.2030
epoch 300: val_loss 0.1425
epoch 400: val_loss 0.1305
epoch 500: val_loss 0.1150
epoch 600: val_loss 0.1029
epoch 700: val_loss 0.0943
epoch 800: val_loss 0.0916
stopped at epoch: 834
best val_loss: 0.0914

The walkthrough

We asked for 1000 epochs and got 834. EarlyStopping watched val_loss, waited through 50 non-improving epochs (patience=50), then pulled the plug. Patience matters because validation loss wobbles — stopping at the first bad epoch would stop almost immediately.

restore_best_weights=True is the flag people forget. Without it, training stops late but the model keeps the weights of the final epoch — which is, by definition, 50 epochs past the best one.

ModelCheckpoint("best.keras", save_best_only=True) writes the model to disk each time the monitored value improves. If your machine dies at epoch 700, the best-so-far survives. The .keras extension is required in Keras 3 — details in saving and loading Keras models.

The custom callback subclasses keras.callbacks.Callback and overrides one of the hook methods: on_epoch_end, on_batch_end, on_train_begin, and friends. logs is a dict carrying the current metric values. Anything you can write in Python can live in a hook — sending yourself a message when training finishes is a five-line callback.

Order matters slightly. Callbacks run in list order. Put loggers first, deciders (EarlyStopping) later, so the log reflects what the decider saw.

Common mistakes

Monitoring loss instead of val_loss. Training loss almost never stops improving, so EarlyStopping never fires and ModelCheckpoint happily saves an ever-more-overfitted model. Monitor the validation value — the one measured on data the model does not train on.

No validation data at all. If fit gets no validation_split or validation_data, there is no val_loss, and the callback warns: Early stopping conditioned on metric 'val_loss' which is not available. Training then runs the full 1000 epochs as if the callback were absent — read warnings.

Patience of 1 or 2. Validation curves are noisy. Tiny patience stops training at the first wobble, long before the real minimum. Tens of epochs is a sane starting range for small models.

Saving every epoch to one filename without save_best_only. You end with the last model, not the best, and heavy disk churn on long runs.

Try it yourself

Set patience=5 and rerun. Compare the stopping epoch and the best val_loss against the run above. Then delete restore_best_weights=True, rerun, and evaluate — the model you hold at the end is measurably worse than the one that existed mid-run.

What to learn next

Researcher — Mathematics and papers.

Early stopping as regularisation

Early stopping is not a convenience — it is a regulariser with theory behind it. For linear models trained by gradient descent on squared loss, stopping at step $t$ is equivalent to ridge regression with penalty $\lambda \approx 1/(\eta t)$, where $\eta$ is the learning rate and $t$ the iteration count: the optimisation path visits solutions of growing effective complexity, and stopping picks a point on that path. Bishop (1995) and Sjöberg and Ljung (1995) formalised the correspondence; Yao, Rosasco and Caponnetto (2007), On early stopping in gradient descent learning, give convergence rates.

The validation-based stopping rule estimates the population risk minimum along the path. The estimate is biased optimistically: the stopping epoch is selected using the validation set, so reported validation loss at that epoch slightly understates true risk — one reason a final untouched test set remains necessary.

The hook protocol

The callback contract is the observer pattern: fit publishes events — on_(train|epoch|batch)_(begin|end) — and each callback receives (index, logs) with logs a dict of scalars. Keras's own progress bar, History object, and TensorBoard writer are themselves callbacks registered implicitly; user callbacks join the same list. The abstraction cost is one Python dispatch per event, negligible against a batch of matrix multiplies at any realistic model size, though per-batch hooks on very small models can dominate runtime — profile before logging per batch.

EarlyStopping implementation detail worth knowing: "improvement" means monitor changed by more than min_delta in the direction given by mode (min for losses, max for accuracy — auto-detected from the name). A metric whose name Keras cannot classify defaults to min, which silently inverts the logic for custom scores where bigger is better. Pass mode="max" explicitly.

References

  • Prechelt (1998), Early stopping — but when? — the empirical study behind patience-style rules.
  • Yao, Rosasco, Caponnetto (2007), On early stopping in gradient descent learning.
  • Goodfellow, Bengio, Courville (2016), Deep Learning, §7.8 — early stopping as regularisation, including the ridge equivalence.

What to learn next