skip to content

In Keras 3, how do you override train_step(), and what must it return?

level: seniorimportance: should knowfreq 45%

answer

  1. one batch, not the whole loop
  2. fit() keeps its callbacks and progress bar
  3. use compute_loss, not the raw function
  4. the return value becomes the logs
  5. the gradient part is backend-specific

basics

~20 s

Subclass keras.Model and define train_step(self, data): unpack the batch, run the forward pass, compute the loss with self.compute_loss(), apply the gradients through self.optimizer, update self.metrics, and return a dict mapping metric names to their current results.

solid answer

~50 s

`train_step` is the body of one training batch, and `fit()` calls it for you — so overriding it keeps callbacks, the progress bar, validation, and every other `fit()` service intact while you change only the arithmetic. The contract is: take `data` (whatever your input pipeline yields, typically `(x, y)` or `(x, y, sample_weight)`), produce `y_pred` by calling `self(x, training=True)`, get the scalar via `self.compute_loss(x=x, y=y, y_pred=y_pred, sample_weight=sample_weight)` so the compiled loss is respected, apply gradients with `self.optimizer`, call `update_state` on the metrics in `self.metrics`, and **return a dict of name → `result()`**. That returned dict is what the progress bar prints and what `History` records. Override `test_step` alongside it, or `evaluate()` and the validation pass will still use the stock logic. Note the Keras 2 helpers `self.compiled_loss` and `self.compiled_metrics` no longer exist in Keras 3, and the gradient part is backend-specific.

code

python · 18 lines
python
import tensorflow as tf  # KERAS_BACKEND=tensorflow
import keras

class CustomModel(keras.Model):
    def train_step(self, data):
        x, y = data
        with tf.GradientTape() as tape:
            y_pred = self(x, training=True)
            loss = self.compute_loss(x=x, y=y, y_pred=y_pred)
        trainable_vars = self.trainable_variables
        gradients = tape.gradient(loss, trainable_vars)
        self.optimizer.apply_gradients(zip(gradients, trainable_vars))
        for metric in self.metrics:
            if metric.name == "loss":
                metric.update_state(loss)
            else:
                metric.update_state(y, y_pred)
        return {m.name: m.result() for m in self.metrics}

go deeper

for a junior

Know that fit() calls a per-batch method you can override, and that overriding it is how Keras lets you customise training without writing your own loop.

for a middle

Be able to recite the four steps — forward pass with training=True, compute_loss, apply gradients, update metrics — and say that the returned dict becomes the progress-bar and History logs.

for a senior

Show that you know compute_loss picks up regularisation and loss_weights, that test_step must be overridden alongside, and that the gradient portion is backend-specific so a GradientTape step pins you to TensorFlow.

for a principal

Judge whether the customisation belongs in train_step at all: a custom loss, a custom Metric or a callback covers most first attempts with far less to maintain, and a bespoke step is code your team owns across every Keras upgrade and backend change.

## The seam Keras gives you `fit()` is a loop that does a great deal you do not want to rewrite: epoch iteration, batching, shuffling, the progress bar, callback dispatch, the validation pass at each epoch end, `History` bookkeeping, and distribution. All of the per-batch *arithmetic* lives in one overridable method, `Model.train_step(self, data)`. Overriding it is the supported way to change what a training batch does — a custom gradient manipulation, an extra regulariser, a GAN's two-optimizer dance — without giving up the machinery around it. `test_step` is its counterpart for `evaluate()` and for the validation pass inside `fit()`, and `predict_step` for `predict()`. ## The contract ``` class CustomModel(keras.Model): def train_step(self, data): x, y = data # 1. forward pass in training mode # 2. scalar loss # 3. gradients -> optimizer # 4. metric.update_state(...) return {m.name: m.result() for m in self.metrics} ``` Four obligations: 1. **Unpack `data` yourself.** It is exactly what your input yields — a `(x, y)` tuple from arrays, a `(x, y, sample_weight)` triple, or a single tensor for an unsupervised model. Handle the arity you actually use. 2. **Call the model with `training=True`.** `self(x, training=True)` is what puts `Dropout` and `BatchNormalization` into training behaviour. Forget the flag and dropout is off during training — a silent quality bug. 3. **Use `self.compute_loss(x=..., y=..., y_pred=..., sample_weight=...)`** rather than calling the loss function directly. It applies the loss you passed to `compile()`, handles multi-output losses and `loss_weights`, and includes any regularisation losses layers added. Calling `keras.losses.mse(y, y_pred)` by hand bypasses all of that. 4. **Return a dict of metric names to scalar results.** The returned mapping is the logs — the progress bar prints it, callbacks receive it as `logs`, and `History.history` accumulates it. Returning the raw loss tensor instead of a dict breaks all three. Metrics themselves are stateful objects: you call `update_state()` inside the step, and `fit()` resets them at each epoch boundary, so the printed value is the epoch-so-far accumulation. ## The backend-specific part Keras 3 runs on TensorFlow, JAX and PyTorch, and *gradient computation is the one part of `train_step` that cannot be written backend-agnostically*, because each framework's autodiff has a different shape: - **TensorFlow backend** — wrap the forward pass in a `tf.GradientTape`, get `tape.gradient(loss, self.trainable_variables)`, then `self.optimizer.apply_gradients(zip(gradients, self.trainable_variables))`. - **PyTorch backend** — clear existing gradients, call `loss.backward()`, and hand the resulting `.grad` values to the optimizer. - **JAX backend** — the step is functional and stateless: you override `compute_loss_and_updates` and a `train_step(self, state, data)` that receives and returns the state, because JAX cannot mutate variables in place inside a transformed function. So a `train_step` written against `tf.GradientTape` pins the model to the TensorFlow backend. If portability matters, that is a real cost to weigh, and everything *else* in the step — the forward call, `compute_loss`, metric updates — can be written with `keras.ops` and stays portable. ## What changed from Keras 2 Keras 2 examples call `self.compiled_loss(y, y_pred, regularization_losses=self.losses)` and `self.compiled_metrics.update_state(y, y_pred)`. **Neither attribute exists in Keras 3.** They were replaced by the public `self.compute_loss(...)` and `self.compute_metrics(...)`, plus explicit iteration over `self.metrics`. Pasting a Keras 2 custom-step snippet into a Keras 3 project fails with an attribute error, and recognising that instantly is a good signal in an interview. ## When it is the right tool — and when it is not Override `train_step` when the *batch computation* is non-standard: multiple optimizers, gradient surgery, a teacher-student distillation loss, an adversarial inner step, or a loss that needs the inputs as well as the targets. Do not reach for it when a simpler seam exists — a custom loss function, a custom `keras.metrics.Metric`, or a callback covers a large share of what people first attempt here, and each is far less code to maintain. ## The forgotten half The most common defect after a successful `train_step` override is leaving `test_step` alone. Validation inside `fit()` and every call to `evaluate()` go through `test_step`, so if your training loss is a custom composite and your evaluation is not, `loss` and `val_loss` are measuring different quantities — and the two curves will never be comparable, which usually gets misread as overfitting.

  • Why call self.compute_loss() instead of the loss function directly?
    `compute_loss()` applies the loss you configured in `compile()`, handles multi-output models and `loss_weights`, threads `sample_weight` through, and folds in regularisation losses that layers registered. Calling `keras.losses.mse(y, y_pred)` by hand silently drops all of that — most visibly the regularisation term, so weight decay you configured stops being applied and nothing warns you.
  • You override train_step only, then compare loss and val_loss — what is wrong?
    Validation runs through `test_step`, which you did not override, so it still computes the stock compiled loss while training computes your custom one. The two curves measure different quantities and are not comparable; the gap is usually misread as overfitting. Override `test_step` with the same loss computation, minus the gradient work.
  • Does a train_step written with tf.GradientTape run on the JAX backend?
    No. Gradient computation is the one backend-specific part of the step: TensorFlow uses `GradientTape`, PyTorch uses `loss.backward()`, and JAX needs a stateless `train_step(self, state, data)` because it cannot mutate variables inside a transformed function. Writing the tape version pins the model to the TensorFlow backend; the rest of the step can stay portable via `keras.ops`.
  • What breaks if train_step returns the loss tensor instead of a dict?
    The return value is the logs for that batch. The progress bar, the `logs` argument every callback receives, and `History.history` all expect a mapping of metric name to scalar. Returning a bare tensor leaves callbacks with nothing to monitor — an `EarlyStopping` watching a name that never appears simply never fires.

saying these in an interview costs you the question

  • Rewriting the epoch loop instead of overriding one batch
  • Using self.compiled_loss, which Keras 3 removed
  • Calling the model without training=True inside the step
  • Returning the loss tensor rather than a metrics dict
  • Overriding train_step and leaving test_step untouched

context