Train a CIFAR-10 CNN with a Custom Keras Training Loop using tf.GradientTape
Learn how to train a CIFAR‑10 CNN with a custom Keras training loop, tf.GradientTape for gradients, manual metric logging, and best‑weight checkpointing.
23 Nov 2025, 18:15 UTC

Desired outcome
Train a small convolutional neural network on the CIFAR-10 dataset using a custom training loop that gives you full control over gradient computation, metric logging, and checkpointing. By the end of the guide you will have a script that logs training and validation metrics per epoch and saves the best model weights to disk.
Prerequisites
- Python 3.8 or newerTensorFlow 2.12+ installed (CPU or GPU build)Ability to write and execute a Python script in your working directoryEnough disk space to store a few megabytes of checkpoint files
Procedure
1. Load and preprocess the data
Use TensorFlow’s built‑in CIFAR‑10 loader, convert images to float32 in the range [0, 1], and create batched tf.data pipelines for training and validation.
import tensorflow as tf from tensorflow.keras import layers, models # Hyper‑parameters (adjust as needed) BATCH_SIZE = 64 EPOCHS = 20 AUTOTUNE = tf.data.AUTOTUNE (x_train, y_train), (x_val, y_val) = tf.keras.datasets.cifar10.load_data() # Normalize to [0, 1] x_train = x_train.astype('float32') / 255.0 x_val = x_val.astype('float32') / 255.0 train_ds = (tf.data.Dataset.from_tensor_slices((x_train, y_train)) .shuffle(10000) .batch(BATCH_SIZE) .prefetch(AUTOTUNE)) val_ds = (tf.data.Dataset.from_tensor_slices((x_val, y_val)) .batch(BATCH_SIZE) .prefetch(AUTOTUNE))2. Define the model architecture
Subclass
tf.keras.Modelto keep the custom loop simple. The example uses two convolutional blocks followed by a dense classifier.class SimpleCNN(tf.keras.Model): def __init__(self, num_classes=10): super().__init__() self.conv1 = layers.Conv2D(32, 3, activation='relu', padding='same') self.pool1 = layers.MaxPooling2D() self.conv2 = layers.Conv2D(64, 3, activation='relu', padding='same') self.pool2 = layers.MaxPooling2D() self.flatten = layers.Flatten() self.dense1 = layers.Dense(128, activation='relu') self.output_layer = layers.Dense(num_classes) def call(self, inputs, training=False): x = self.conv1(inputs) x = self.pool1(x) x = self.conv2(x) x = self.pool2(x) x = self.flatten(x) x = self.dense1(x) return self.output_layer(x) model = SimpleCNN()3. Choose optimizer, loss and metrics
optimizer = tf.keras.optimizers.Adam() loss_fn = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True) # Metrics that we will update manually train_loss_metric = tf.keras.metrics.Mean(name='train_loss') train_acc_metric = tf.keras.metrics.SparseCategoricalAccuracy(name='train_accuracy') val_loss_metric = tf.keras.metrics.Mean(name='val_loss') val_acc_metric = tf.keras.metrics.SparseCategoricalAccuracy(name='val_accuracy')4. Implement a single training step with tf.GradientTape
This function computes gradients for one batch and applies them. It also updates the training metrics.
@tf.function # Optional: trace for performance; remove if you encounter tracing issues def train_step(images, labels): with tf.GradientTape() as tape: logits = model(images, training=True) loss = loss_fn(labels, logits) gradients = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(gradients, model.trainable_variables)) train_loss_metric.update_state(loss) train_acc_metric.update_state(labels, logits) return loss5. Validation step (no gradients)
def val_step(images, labels): logits = model(images, training=False) v_loss = loss_fn(labels, logits) val_loss_metric.update_state(v_loss) val_acc_metric.update_state(labels, logits) return v_loss6. Main training loop with checkpointing
After each epoch we log metrics, compare validation accuracy to the best seen so far, and save the model weights if there is an improvement.
import os CHECKPOINT_DIR = './cifar10_checkpoints' os.makedirs(CHECKPOINT_DIR, exist_ok=True) best_val_acc = 0.0 for epoch in range(1, EPOCHS + 1): # Reset metrics at the start of each epoch train_loss_metric.reset_states() train_acc_metric.reset_states() val_loss_metric.reset_states() val_acc_metric.reset_states() # Training batches for step, (x_batch, y_batch) in enumerate(train_ds): train_step(x_batch, y_batch) # Validation batches for x_batch, y_batch in val_ds: val_step(x_batch, y_batch) # Retrieve results train_loss = train_loss_metric.result() train_acc = train_acc_metric.result() val_loss = val_loss_metric.result() val_acc = val_acc_metric.result() template = ('Epoch {:02d}: loss={:.4f}, acc={:.2f}% | ' 'val_loss={:.4f}, val_acc={:.2f}%') print(template.format(epoch, float(train_loss), float(train_acc)*100, float(val_loss), float(val_acc)*100)) # Checkpoint if validation accuracy improved if float(val_acc) > best_val_acc: best_val_acc = float(val_acc) checkpoint_path = os.path.join(CHECKPOINT_DIR, 'best_weights') model.save_weights(checkpoint_path) print(f' --> Saved new best model to {checkpoint_path}') print('Training finished. Best validation accuracy: {:.2f}%'.format(best_val_acc*100))Expected checks
- After each epoch the script prints four numbers: training loss, training accuracy, validation loss, validation accuracy. Loss should generally decrease and accuracy increase over epochs.
- The directory
./cifar10_checkpointscontains a file namedbest_weights(and an accompanyingbest_weights.data-00000-of-00001shard) when an improvement occurs. - Reloading the saved weights into a freshly instantiated
SimpleCNNmodel and running a forward pass on a few validation samples should produce identical logits to those obtained before saving (you can verify withtf.reduce_all(tf.equal(logits_before, logits_after))). - If you interrupt training (e.g., Ctrl‑C) you can resume from the latest checkpoint by loading the weights before the next epoch.
Recovery options
Loss becomes NaN or training diverges
- Reduce the learning rate (e.g.,
optimizer = tf.keras.optimizers.Adam(1e-4)) and restart from the latest checkpoint. - Alternatively, re‑initialize the optimizer’s slot variables by creating a new optimizer instance.
Out‑of‑memory (OOM) on GPU
- Decrease
BATCH_SIZE(e.g., to 32 or 16). - Enable mixed precision:
tf.keras.mixed_precision.set_global_policy('mixed_float16')before building the model.
Accidentally overwritten checkpoint
Since the script only writes when validation accuracy improves, you can safely delete the checkpoint directory and let the script recreate it from scratch.
Limitations and practical verification
- The custom loop does not automatically leverage Keras features such as built‑in callbacks (e.g., EarlyStopping, ReduceLROnPlateau). You must implement those manually if needed.
- The
@tf.functiondecorator ontrain_stepimproves speed but can hide debugging information; remove it while developing to see eager execution errors. - Saving only weights assumes the model architecture stays identical. If you change the layer order or number of filters, reloading will raise a
ValueError. - To verify that the checkpoint is usable, run a short verification script after training:
# verification.py import tensorflow as tf from your_module import SimpleCNN # replace with actual import model = SimpleCNN() latest = tf.train.latest_checkpoint('./cifar10_checkpoints') model.load_weights(latest) # Grab a batch from validation set for x_batch, _ in val_ds.take(1): preds = model(x_batch, training=False) print('Sample logits shape:', preds.shape) breakIf the script runs without error and prints a tensor shape, the checkpoint is loadable.
Summary
You now have a complete, end‑to‑end example of a Keras model trained with a custom loop using
tf.GradientTape. The procedure covers data loading, model definition, manual gradient application, metric handling, epoch‑wise validation, and selective checkpointing. Adjust the hyper‑parameters, architecture, or recovery strategies to fit your specific workload or hardware constraints.
0 replies
A thoughtful contribution can make all the difference. Be the first to share one.