Customizing Keras Training with Model Subclassing
Learn how to subclass tf.keras.Model to override train_step, allowing for custom training logic while retaining the benefits of the Keras compile and fit API.
09 Jun 2026, 08:11 UTC

The Problem: Rigid Training Loops
Standard Keras models using the Sequential or Functional API are efficient for most tasks, but they struggle when you need non-standard training logic. If your project requires GANs, Variational Autoencoders (VAEs), or complex multi-loss optimization, the default model.fit() behavior becomes a bottleneck.
The common reaction is to write a fully manual training loop using tf.GradientTape. However, doing this means losing the built-in convenience of Keras callbacks, metric tracking, and the compile API.
The Solution: Overriding train_step
By subclassing tf.keras.Model, you can override the train_step method. This allows you to inject custom logic into the training process while still using model.fit() to handle the orchestration of epochs, batches, and callbacks. This approach is fully supported in TensorFlow 2.8 and newer.
How train_step Works
The train_step method is called by Keras for every batch of data. To implement it correctly, you must follow a specific sequence:
- Forward Pass: Pass the input
xthrough the model. - Loss Calculation: Use
self.compiled_lossto ensure the loss function defined inmodel.compile()is used. - Gradient Application: Use a
tf.GradientTapeto calculate gradients and apply them viaself.optimizer. - Metric Updates: Update
self.compiled_metricsso that progress bars and callbacks receive accurate data.
Worked Example: Custom CNN on MNIST
The following example demonstrates a simple CNN where we override train_step to maintain full control over the gradient application while keeping the Keras API.
import tensorflow as tf
class CustomCNN(tf.keras.Model):
def __init__(self):
super().__init__()
self.conv1 = tf.keras.layers.Conv2D(32, 3, activation='relu')
self.pool1 = tf.keras.layers.MaxPool2D()
self.flatten = tf.keras.layers.Flatten()
self.dense = tf.keras.layers.Dense(10, activation='softmax')
def call(self, x, training=False):
x = self.conv1(x)
x = self.pool1(x)
x = self.flatten(x)
return self.dense(x)
def train_step(self, data):
# Unpack the data
x, y = data
with tf.GradientTape() as tape:
y_pred = self(x, training=True)
# Use compiled_loss to integrate with model.compile()
loss = self.compiled_loss(y, y_pred, regularization_losses=self.losses)
# Compute and apply gradients
gradients = tape.gradient(loss, self.trainable_variables)
self.optimizer.apply_gradients(zip(gradients, self.trainable_variables))
# Update metrics
self.compiled_metrics.update_state(y, y_pred)
# Return a dictionary mapping metric names to current values
return {m.name: m.result() for m in self.metrics}
# Setup data
(x_train, y_train), _ = tf.keras.datasets.mnist.load_data()
x_train = x_train[..., None].astype('float32') / 255.0
train_ds = tf.data.Dataset.from_tensor_slices((x_train, y_train)).batch(128)
# Execution
model = CustomCNN()
model.compile(optimizer='adam',
loss='sparse_categorical_crossentropy',
metrics=['accuracy'])
model.fit(train_ds, epochs=1)
Verification: Run this script using python script_name.py. You should see the standard Keras training progress bar, confirming that model.fit() is successfully calling your custom train_step.
Engineering Trade-offs
While subclassing provides flexibility, it introduces several responsibilities that the Functional API handles automatically:
| Feature | Functional/Sequential API | Model Subclassing |
|---|---|---|
| Input Shapes | Inferred or defined via Input layer | Must be handled in call or via first batch |
| Graph Compilation | Automatic static graph | Eager by default; requires @tf.function for speed |
| Distributed Training | Handled by Keras | Model must be instantiated within Strategy scope |
Critical Limitations
- Manual Management: You are responsible for gradient clipping and mixed-precision scaling. If you enable mixed precision, you must use
optimizer.get_scaled_lossandoptimizer.get_unscaled_gradients. - Metric Aggregation: Only metrics updated via
self.compiled_metricsare automatically aggregated across multiple GPUs in aMirroredStrategy.
Practical Next Steps
To move from a prototype to production, consider these optimizations:
- Performance: Wrap your
train_steplogic in a@tf.functiondecorator to compile the logic into a graph, significantly reducing Python overhead. - Scaling: If using multiple GPUs, wrap the model instantiation and
model.compile()insidewith tf.distribute.MirroredStrategy().scope():. - Validation: Override
test_stepsimilarly totrain_stepto ensure your validation logic matches your training logic exactly.
0 replies
A thoughtful contribution can make all the difference. Be the first to share one.