skip to content

When subclassing keras.callbacks.Callback, what is in logs at on_epoch_end vs on_train_batch_end?

level: middleimportance: must knowfreq 58%

answer

  1. a plain dict of scalar floats
  2. validation has not run yet mid-epoch
  3. batch values are running aggregates
  4. self.model is attached before training
  5. writing to logs reaches later callbacks

basics

~20 s

At on_epoch_end, logs holds the epoch's training metrics plus validation metrics under val_-prefixed keys. At on_train_batch_end, only training metrics are present — validation has not run — and their values are running aggregates over the epoch so far, not that single batch.

solid answer

~50 s

A custom callback subclasses `keras.callbacks.Callback` and overrides hooks such as `on_epoch_end(self, epoch, logs=None)` or `on_train_batch_end(self, batch, logs=None)`. The `logs` argument is a plain dict of floats keyed by metric name. At epoch end it carries `loss`, every compiled metric, and — only if `fit()` was given validation data — the same names again prefixed with `val_`. At batch end there are no `val_` keys at all, because validation runs once per epoch after the last batch. The other batch-level subtlety: Keras metrics are stateful and reset at the start of each epoch, so the value you read at batch 50 is the aggregate over batches 1–50, not batch 50 alone. Inside any hook, `self.model` gives you the model Keras attached, so `self.model.stop_training = True` halts the run; mutating `logs` is visible to callbacks later in the list and to the History object.

code

python · 15 lines
python
import math
import keras


class StopOnBadLoss(keras.callbacks.Callback):
    def on_train_batch_end(self, batch, logs=None):
        logs = logs or {}
        loss = logs.get("loss")
        if loss is not None and not math.isfinite(loss):
            self.model.stop_training = True

    def on_epoch_end(self, epoch, logs=None):
        logs = logs or {}
        gap = logs.get("loss", 0.0) - logs.get("val_loss", 0.0)
        logs["train_val_gap"] = gap

go deeper

for a junior

Know that you subclass keras.callbacks.Callback, that hooks receive a logs dict of metric values, and that validation metrics appear under val_-prefixed keys only at the end of an epoch.

for a middle

Explain the exact hook signatures, why no val_ key exists at batch end, that batch metric values are epoch-so-far aggregates, and how self.model is attached before training begins.

for a senior

Show what you build with it: a stopping rule on a wall-clock or NaN condition via model.stop_training, keeping batch hooks cheap because they sit on the per-step critical path, and using logs mutation to feed derived metrics into the training record.

for a principal

Take a position on where instrumentation belongs — a callback inside the fit loop versus the surrounding job — and on the coupling cost of callbacks that reach into model internals across many teams and training scripts.

## The subclass contract Every Keras callback derives from `keras.callbacks.Callback` and overrides whichever hooks it needs. Keras calls them; you never do. The hook set covers the three loops: - Training: `on_train_begin`, `on_epoch_begin`, `on_train_batch_begin`, `on_train_batch_end`, `on_epoch_end`, `on_train_end`. - Evaluation: `on_test_begin`, `on_test_batch_begin`, `on_test_batch_end`, `on_test_end`. - Prediction: `on_predict_begin`, `on_predict_batch_begin`, `on_predict_batch_end`, `on_predict_end`. Epoch hooks take `(epoch, logs=None)`; batch hooks take `(batch, logs=None)`; begin/end-of-phase hooks take `(logs=None)`. The index argument is zero-based. Always keep the `logs=None` default — Keras may call a hook without one. Before training starts, Keras calls `set_model()` and `set_params()`, which populate `self.model` and `self.params`. That is how a callback reaches the model it is attached to without you passing it in: `self.model.optimizer`, `self.model.layers`, `self.model.get_weights()` are all available. `self.params` carries run-level information such as the number of epochs and steps. ## What is actually in logs `logs` is an ordinary Python dict of scalar floats — not tensors, and not backend-specific objects, which matters in Keras 3 where the backend may be TensorFlow, JAX or PyTorch. The keys are metric names. At `on_epoch_end` during `fit()`, the dict contains `loss`, one entry per compiled metric under the name that metric reports (`accuracy`, `mae`, whatever you named it), and — if and only if `fit()` received `validation_data` or `validation_split` — a `val_`-prefixed copy of each of those keys computed on the validation set. This is the dict `ModelCheckpoint` and `EarlyStopping` look their `monitor` up in, which is why `monitor="val_loss"` with no validation data finds nothing. At `on_train_batch_end` the dict has the training metrics only. There are no `val_` keys, because validation is a separate pass that runs after the final training batch of the epoch. Reaching for `logs["val_loss"]` in a batch hook is the most common mistake in a hand-written callback, and it is a `KeyError` rather than a silent zero. The second batch-level subtlety is aggregation. Keras metrics accumulate state across an epoch and reset at the epoch boundary, so the number in a batch-end log is the metric's value over all batches so far in this epoch — a running average, not an instantaneous per-batch reading. A callback that plots `logs["loss"]` per batch expecting a spiky curve gets a smoothed one, and the smoothing changes as the epoch progresses (early batches move the average far more than late ones). If you need a genuinely per-batch number you have to compute it yourself rather than read it out of the logs. ## Mutating logs The dict is passed by reference through the callback list in order, so writing `logs["my_stat"] = value` in `on_epoch_end` makes that key visible to every callback that runs after yours — including `CSVLogger`, which writes whatever columns it sees, and `History`, which copies the epoch's logs into `history.history`. That is the supported way to get a derived quantity into the training record without wrapping the model. The corollary is that the callback list order is meaningful: a callback that consumes a key must come after the one that adds it. ## Stopping and reaching into the model `self.model.stop_training = True` is the documented way to halt a run from a callback — it is what `EarlyStopping` itself does. The training loop checks the flag after the current epoch, so `fit()` returns normally rather than raising, and the `History` it returns is short. This is how you build stopping rules Keras does not ship: stop on a wall-clock budget, on a non-finite loss, on an external signal file, on a metric of your own. The learning rate is reachable the same way, via `self.model.optimizer.learning_rate`, which is exactly how the built-in scheduling callbacks work. ## Cost Batch hooks run once per training step, inside the loop, on the critical path. Anything expensive there — writing to disk, computing a full metric, pulling weights to host memory — multiplies by the step count and can dominate epoch time. Put heavy work in `on_epoch_end`, and if you need per-step data, accumulate cheaply in the batch hook and flush once per epoch. For a hook body that is genuinely a one-liner, `keras.callbacks.LambdaCallback` avoids writing a class at all.

  • How would you write a callback that stops training when the loss becomes NaN?
    Override `on_train_batch_end`, read `logs["loss"]`, and if it is not finite set `self.model.stop_training = True`. The flag is checked by the training loop, so `fit()` returns normally with a short History rather than raising. Keras also ships `keras.callbacks.TerminateOnNaN`, which does exactly this — worth naming, since reaching for a hand-rolled version of an existing callback is itself a small red flag.
  • Why does a per-batch loss curve read out of logs look smoother than you expect?
    Because compiled Keras metrics are stateful: they accumulate over the epoch and reset at the epoch boundary. The value in a batch-end log is therefore the running aggregate across batches 1..N of this epoch, so late batches barely move it. If you want true per-batch values you must compute them yourself instead of reading the logs dict.
  • You add a key to logs in on_epoch_end and want CSVLogger to write it. What has to be true?
    Your callback has to run before CSVLogger in the list passed to `fit(callbacks=[...])`, because the dict is threaded through the callbacks in order and CSVLogger writes the keys it sees when its own hook fires. Order the list producer-first. The same mechanism puts the key into `history.history`, since History copies the epoch's logs.
  • Which hooks fire during model.predict(), and what is in their logs?
    `on_predict_begin`, `on_predict_batch_begin`, `on_predict_batch_end` and `on_predict_end`. There are no metrics during prediction — nothing is compiled into a loss and no targets exist — so the logs dict carries no metric keys. Callbacks written against `fit()` and casually reused on `predict()` typically fail there for exactly that reason.

saying these in an interview costs you the question

  • Reads logs['val_loss'] inside a batch-level hook
  • Thinks batch logs hold that batch's own loss
  • Returns False from a hook to stop training
  • Expects logs to contain backend tensors, not floats
  • Does heavy disk writes in on_train_batch_end

context