Most failed training runs are not modeling failures. They are software failures: a label column that leaked into the features, an augmentation layer still active at evaluation time, a custom layer whose get_config drops an argument, a tf.data map that silently broadcasts the wrong shape. These bugs cost hours of GPU time and, worse, sometimes produce a plausible-looking metric that nobody questions until deployment.
Training code is software. It should have tests. The useful ones are cheap: the suite below runs on CPU with tiny synthetic tensors, finishes in about two minutes, and is the first thing we add when we take over an unfamiliar Keras repository.
What not to test
Do not assert on accuracy. A test like assert val_accuracy > 0.91 will fail on a different GPU, a different seed, or a CUDA upgrade, and your team will learn to ignore red builds. Quality thresholds belong in an evaluation report that a human reads, not in a unit test. Tests should assert on contracts — shapes, dtypes, ranges, invariances, round-trips — that must hold no matter what the data looks like.
1. Shape and dtype contracts
Build the model at the configured input shape and check the output. This catches most refactor breakage in under a second.
def test_model_output_contract():
model = build_model(input_shape=(96, 96, 3), num_classes=7)
x = keras.random.uniform((2, 96, 96, 3))
y = model(x, training=False)
assert y.shape == (2, 7)
assert y.dtype == "float32"
The dtype assertion matters under mixed precision: the final layer should return float32 even when the compute policy is mixed_float16. A test is a cheaper reminder than a loss that goes to NaN at hour three.
2. Overfit one batch
The single most informative test in deep learning. Take eight examples, turn off augmentation and regularization, and train for a few dozen steps. If the loss does not collapse toward zero, something is wired wrong — a detached head, a frozen trunk, a label/prediction mismatch, a loss expecting logits while the model emits softmax probabilities.
def test_overfits_single_batch():
model = build_model(input_shape=(32, 32, 3), num_classes=4)
model.compile(optimizer=keras.optimizers.Adam(1e-3),
loss="sparse_categorical_crossentropy")
x = keras.random.uniform((8, 32, 32, 3))
y = np.array([0, 1, 2, 3, 0, 1, 2, 3])
hist = model.fit(x, y, epochs=60, batch_size=8, verbose=0)
assert hist.history["loss"][-1] < 0.05
It is a learning-capacity test, not a quality test, so the threshold is stable across machines. Keep the model small enough that it runs on a CPU runner.
3. Pipeline invariants and leakage
Test the tf.data pipeline separately from the model. Three checks cover most of what goes wrong:
- Element spec: one batch yields the expected shapes, dtypes, and value range (normalized inputs inside
[0, 1]or zero-mean, as documented). - Split disjointness: the intersection of train and validation identifiers is empty. For grouped data — patients, sites, devices, machines — assert on the group key, not the row id. This is the test that would have caught most of the leakage we have found in client code.
- Augmentation is training-only: run the same batch through the input pipeline twice in eval mode and assert the outputs are identical; in training mode, assert they differ.
keras.layers.RandomFlipand friends respect thetrainingflag, but only if your code actually passes it.
4. Serialization round-trip
Every custom layer, metric, and loss needs a round-trip test, and this is the one teams skip until the day a checkpoint will not load.
def test_custom_layer_roundtrip(tmp_path):
model = build_model(input_shape=(16,), num_classes=2)
x = keras.random.uniform((4, 16))
before = model.predict(x, verbose=0)
path = tmp_path / "m.keras"
model.save(path)
restored = keras.saving.load_model(path)
after = restored.predict(x, verbose=0)
np.testing.assert_allclose(before, after, atol=1e-6)
If get_config omits a constructor argument, this fails immediately instead of six weeks later. Add the matching model.export() check if you ship a serving artifact — the inference path deserves its own parity test.
5. A one-step smoke run of the real entry point
Finally, call the actual training script with a config override that runs one step on a few synthetic records: steps_per_epoch=2, epochs=1. This exercises argument parsing, config loading, callbacks, checkpoint paths, and logging — the plumbing that unit tests never touch and that breaks the moment a flag is renamed.
Wiring it up
Run the suite on CPU in CI on every pull request. Pin seeds inside tests, keep tensors tiny, and mark anything needing a GPU or real data with @pytest.mark.slow so it runs nightly rather than per-commit. Budget: five tests, roughly 150 lines, two minutes of runtime.
The payoff is not elegance. It is that the next failed run is a modeling question — the data is wrong, the architecture is wrong, the target is unlearnable — instead of a typo you paid twelve GPU-hours to discover.