TensorFlow & Keras: Production Engine
1 · The lesson
readKeras's .fit() is the perfect quick-start: three lines and you're training. It's also where most production projects hit a wall. Multi-input models with mixed loss weights. Streaming data from S3 with on-the-fly augmentation. Mixed-precision training to fit a model on a single GPU. Custom training loops that interleave generator and discriminator updates. None of that is model.fit(X, y).
This lesson is the bridge from "I trained a Keras toy" to "I shipped a TF/Keras model that handles real traffic." We'll cover the TF 2.x execution model, tf.data pipelines that don't bottleneck the GPU, the three Keras APIs and when each is right, custom training with tf.GradientTape, callbacks worth knowing, mixed precision, distribution strategies, and the saving formats you'll actually deploy.
Run in Colab or locally with
pip install tensorflow. Expected outputs in comments. TensorFlow has no WebAssembly build, so unlike most of this site it genuinely cannot run in your browser — if you want to see a network train here, build one from scratch in NumPy instead; it reaches 97% on hand-written digits in about a second.
1. TF 2.x — Eager by Default, Graphs on Demand
TF 1.x was graph-first: you defined a static computation graph, then ran it in a Session. Debugging meant staring at placeholders. TF 2.x flipped that: eager execution is the default, just like NumPy.
import tensorflow as tf a = tf.constant([1.0, 2.0, 3.0]) b = tf.constant([4.0, 5.0, 6.0]) print(a + b) # tf.Tensor([5. 7. 9.], shape=(3,), dtype=float32) print((a + b).numpy()) # [5. 7. 9.] — back to NumPy
You can print(), set breakpoints, use pdb. Each op runs immediately. The cost: Python is slow, and per-op kernel launches don't fuse well on GPU.
tf.function is the escape hatch. Decorate a Python function and TF traces it into a graph the first time it's called, then replays the compiled graph on every subsequent call.
@tf.function def fast_step(x, y): return tf.reduce_mean(tf.square(x - y)) # First call: traces the graph (slow). # Subsequent calls: graph executes (fast, fused, no Python overhead). fast_step(a, b)
setup added so this can run · defines tf, a, b
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) tf = _AutoMock('tf') a = _AutoMock('a') b = _AutoMock('b')
Inside a tf.function-decorated function, you write Python; TF traces it once and compiles. Conditionals become tf.cond, loops become tf.while_loop, prints become tf.print. Hot training loops should be inside @tf.function. Cold setup code can stay eager.
Caveat: tracing is per-input-signature. If you call fast_step with different shapes or dtypes, TF re-traces. Use tf.function(input_signature=[...]) to lock the signature in production.
2. Tensors vs NumPy — The Boundary
tf.Tensor and np.ndarray look similar but live on different sides of the device boundary.
import numpy as np import tensorflow as tf np_arr = np.array([[1, 2], [3, 4]], dtype=np.float32) tf_t = tf.convert_to_tensor(np_arr) # CPU<->device copy print(tf_t.device) # /job:localhost/replica:0/task:0/device:GPU:0 (if available) print(tf_t.numpy()) # back to NumPy, blocking copy
Conversion costs:
- NumPy → tf.Tensor copies data to the device (GPU/TPU). Not free.
- tf.Tensor → NumPy with
.numpy()blocks until the op completes, then copies back to CPU. Blocking — kills pipelining.
Rule: stay in tensor-land for the whole training step. Only call .numpy() for logging, plotting, or returning final results. Calling .numpy() inside a training loop is the most common reason "GPU utilisation is 30%."
3. tf.data — The Pipeline That Feeds the GPU
A GPU costs more per second than your engineer time. If the GPU is waiting for data, you're burning money. tf.data is the official pipeline API: it builds an asynchronous, parallel, prefetched stream of batches that overlaps preprocessing with model execution.
import tensorflow as tf # 1. Source: where does data come from? ds = tf.data.Dataset.from_tensor_slices((X_train, y_train)) # in-memory # Or from files: # ds = tf.data.Dataset.list_files("s3://bucket/images/*.jpg") # 2. Shuffle (buffer large enough for true randomness) ds = ds.shuffle(buffer_size=10_000, seed=42, reshuffle_each_iteration=True) # 3. Map: preprocess in parallel across CPU threads def preprocess(x, y): x = tf.cast(x, tf.float32) / 255.0 return x, y ds = ds.map(preprocess, num_parallel_calls=tf.data.AUTOTUNE) # 4. Batch ds = ds.batch(64, drop_remainder=True) # 5. Prefetch: GPU works on batch N while CPU prepares batch N+1 ds = ds.prefetch(tf.data.AUTOTUNE) # 6. Hand to model.fit model.fit(ds, epochs=10)
setup added so this can run · defines X_train, y_train, model
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) X_train = _AutoMock('X_train') y_train = _AutoMock('y_train') model = _AutoMock('model')
Each step matters:
shuffle(buffer_size)must be larger than your batch size, ideally close to dataset size. Too-small buffer = correlated batches = noisy gradients.map(... , num_parallel_calls=AUTOTUNE)preprocesses in parallel. Withoutnum_parallel_calls, it runs serially — single-threaded preprocessing is often the bottleneck.batchbeforeprefetch— prefetching at the batch granularity hides batch construction cost too.prefetch(AUTOTUNE)at the end lets TF decide how many batches to buffer. The single most important line for GPU utilisation.
Order matters. shuffle().batch() shuffles examples then groups them — what you usually want. batch().shuffle() shuffles batches but keeps each batch's contents fixed — almost never what you want.
For image pipelines, decoding and resizing belong inside .map():
def load_image(path): raw = tf.io.read_file(path) img = tf.io.decode_jpeg(raw, channels=3) img = tf.image.resize(img, [224, 224]) return tf.cast(img, tf.float32) / 255.0 ds = tf.data.Dataset.list_files("data/*.jpg") ds = ds.map(load_image, num_parallel_calls=tf.data.AUTOTUNE) ds = ds.batch(32).prefetch(tf.data.AUTOTUNE)
setup added so this can run · defines tf
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) tf = _AutoMock('tf')
tf.io.decode_jpeg runs on CPU, in parallel, while the GPU trains on the previous batch. With a properly tuned pipeline, GPU utilisation should be 90%+.
4. The Functional API — Beyond Sequential
Sequential is fine for linear stacks. The moment you have two inputs, a branch, a shared layer, or multiple outputs, you need the Functional API.
from tensorflow import keras from tensorflow.keras import layers # Two-input multi-task model: image + tabular features → classification + regression image_in = keras.Input(shape=(224, 224, 3), name="image") tab_in = keras.Input(shape=(10,), name="tabular") # Image branch x = layers.Conv2D(32, 3, activation="relu")(image_in) x = layers.GlobalAveragePooling2D()(x) x = layers.Dense(64, activation="relu")(x) # Tabular branch t = layers.Dense(32, activation="relu")(tab_in) # Merge merged = layers.concatenate([x, t]) merged = layers.Dense(64, activation="relu")(merged) # Two heads (multi-task) class_out = layers.Dense(5, activation="softmax", name="class")(merged) price_out = layers.Dense(1, name="price")(merged) model = keras.Model( inputs=[image_in, tab_in], outputs={"class": class_out, "price": price_out}, ) model.compile( optimizer="adamw", loss={"class": "sparse_categorical_crossentropy", "price": "mse"}, loss_weights={"class": 1.0, "price": 0.1}, # weight one loss against the other metrics={"class": ["accuracy"], "price": ["mae"]}, )
Named inputs and outputs let you pass dicts to .fit():
model.fit( {"image": img_array, "tabular": tab_array}, {"class": class_labels, "price": prices}, epochs=10, )
setup added so this can run · defines model, img_array, tab_array, class_labels, prices
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) model = _AutoMock('model') img_array = _AutoMock('img_array') tab_array = _AutoMock('tab_array') class_labels = _AutoMock('class_labels') prices = _AutoMock('prices')
This is the API you'll use 80% of the time in production. Shared layers work by reusing a layer object:
shared_embed = layers.Embedding(vocab_size, 64) left = shared_embed(left_in) right = shared_embed(right_in) # SAME weights — siamese network
setup added so this can run · defines vocab_size, left_in, right_in, layers
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) vocab_size = _AutoMock('vocab_size') left_in = _AutoMock('left_in') right_in = _AutoMock('right_in') layers = _AutoMock('layers')
5. Subclassing — Maximum Flexibility
When the model needs dynamic control flow, custom training behaviour, or research-grade weirdness, subclass keras.Model:
class ResidualMLP(keras.Model): def __init__(self, hidden=128, num_classes=10): super().__init__() self.dense1 = layers.Dense(hidden, activation="relu") self.dense2 = layers.Dense(hidden, activation="relu") self.proj = layers.Dense(hidden) self.head = layers.Dense(num_classes, activation="softmax") def call(self, x, training=False): residual = self.proj(x) h = self.dense1(x) if training: h = tf.nn.dropout(h, rate=0.3) # dynamic logic h = self.dense2(h) return self.head(h + residual) model = ResidualMLP() model.build(input_shape=(None, 50)) # define input shape model.summary()
setup added so this can run · defines keras, layers, tf
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) keras = _AutoMock('keras') layers = _AutoMock('layers') tf = _AutoMock('tf')
The training flag flows through automatically — Keras passes training=True during .fit() and training=False during .evaluate()/.predict(). Layers like Dropout and BatchNorm rely on it.
Subclassing is the right choice for:
- Research code where the architecture evolves between runs
- Models with input-dependent control flow (e.g. recursive networks)
- Anything you want to inherit from and override
Trade-off: Subclassed models can't be serialised to the Keras JSON config format without get_config() boilerplate. Use the Functional API by default; subclass when you need it.
6. Custom Layers
When you need a layer that doesn't exist, subclass keras.layers.Layer. The pattern is __init__ (config), build (lazy weight creation when input shape is known), call (forward).
class ScaledDotAttention(layers.Layer): def __init__(self, dim, **kwargs): super().__init__(**kwargs) self.dim = dim self.scale = dim ** -0.5 def build(self, input_shape): self.wq = self.add_weight("wq", shape=(input_shape[-1], self.dim)) self.wk = self.add_weight("wk", shape=(input_shape[-1], self.dim)) self.wv = self.add_weight("wv", shape=(input_shape[-1], self.dim)) def call(self, x): q = x @ self.wq k = x @ self.wk v = x @ self.wv attn = tf.nn.softmax(q @ tf.transpose(k, [0, 2, 1]) * self.scale, axis=-1) return attn @ v def get_config(self): return {**super().get_config(), "dim": self.dim}
setup added so this can run · defines layers, kwargs, tf
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) layers = _AutoMock('layers') kwargs = _AutoMock('kwargs') tf = _AutoMock('tf')
get_config() is what lets .keras serialisation round-trip your custom layer. Without it, saved models can be loaded only with custom_objects={"ScaledDotAttention": ScaledDotAttention} passed manually.
7. Custom Training Loops with tf.GradientTape
.fit() is great until you need to do something it can't express: GAN training (alternate generator/discriminator), reinforcement learning, custom loss with auxiliary heads, gradient accumulation, gradient penalties. Drop down to a manual loop with tf.GradientTape:
optimizer = keras.optimizers.AdamW(learning_rate=1e-3) loss_fn = keras.losses.SparseCategoricalCrossentropy(from_logits=True) @tf.function # compile the step def train_step(x, y): with tf.GradientTape() as tape: logits = model(x, training=True) loss = loss_fn(y, logits) # Add weight decay or regularisation here if you want grads = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables)) return loss for epoch in range(EPOCHS): for x, y in train_ds: loss = train_step(x, y) print(f"epoch {epoch} loss={loss.numpy():.4f}")
setup added so this can run · defines tf, EPOCHS, train_ds, keras, model
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) tf = _AutoMock('tf') EPOCHS = _AutoMock('EPOCHS') train_ds = [("alpha", 1), ("beta", 2), ("gamma", 3)] keras = _AutoMock('keras') model = _AutoMock('model')
The tape records every op on a watched variable; tape.gradient(loss, vars) returns the gradient of loss with respect to each variable. The @tf.function decorator compiles train_step — without it, each step pays Python's overhead.
GAN training is the canonical example where you can't use .fit():
@tf.function def gan_step(real_x): noise = tf.random.normal([BATCH, LATENT]) with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape: fake_x = generator(noise, training=True) real_out = discriminator(real_x, training=True) fake_out = discriminator(fake_x, training=True) g_loss = generator_loss(fake_out) d_loss = discriminator_loss(real_out, fake_out) g_grads = gen_tape.gradient(g_loss, generator.trainable_variables) d_grads = disc_tape.gradient(d_loss, discriminator.trainable_variables) g_opt.apply_gradients(zip(g_grads, generator.trainable_variables)) d_opt.apply_gradients(zip(d_grads, discriminator.trainable_variables)) return g_loss, d_loss
setup added so this can run · defines tf, BATCH, LATENT, generator, discriminator, generator_loss, discriminator_loss, g_opt, d_opt
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) tf = _AutoMock('tf') BATCH = _AutoMock('BATCH') LATENT = _AutoMock('LATENT') generator = _AutoMock('generator') discriminator = _AutoMock('discriminator') def generator_loss(*_a, **_kw): print('-> generator_loss() called') return _AutoMock('generator_loss()') def discriminator_loss(*_a, **_kw): print('-> discriminator_loss() called') return _AutoMock('discriminator_loss()') g_opt = _AutoMock('g_opt') d_opt = _AutoMock('d_opt')
Two tapes, two optimisers, one step. .fit() simply can't express this. Custom loops give you total control at the cost of writing the loop yourself.
8. Callbacks — Hooks Into .fit()
Callbacks let you customise .fit() without dropping to a manual loop. The ones you'll actually use:
from tensorflow.keras import callbacks cbs = [ callbacks.ModelCheckpoint( "models/best.keras", monitor="val_loss", save_best_only=True, # only overwrite if val_loss improves save_weights_only=False, # save full model ), callbacks.EarlyStopping( monitor="val_loss", patience=5, # stop if no improvement for 5 epochs restore_best_weights=True, # roll back to best, not last ), callbacks.ReduceLROnPlateau( monitor="val_loss", factor=0.5, patience=2, # drop LR by 0.5x if no improvement min_lr=1e-6, ), callbacks.TensorBoard(log_dir="logs/run01"), callbacks.CSVLogger("logs/train.csv"), ] model.fit(train_ds, validation_data=val_ds, epochs=100, callbacks=cbs)
setup added so this can run · defines train_ds, model, val_ds
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) train_ds = _AutoMock('train_ds') model = _AutoMock('model') val_ds = _AutoMock('val_ds')
restore_best_weights=True on EarlyStopping is non-optional in production. Without it, the final weights are whatever epoch you stopped on — often worse than the best epoch you trained through.
Custom callbacks are a subclass with the methods you need. Available hooks include on_train_begin, on_epoch_begin/end, on_batch_begin/end, etc.
class GradientNormLogger(callbacks.Callback): def on_train_batch_end(self, batch, logs=None): # Caveat: this requires a custom training loop that exposes grad norms, # or you can probe model.optimizer.iterations etc. via on_epoch_end. if batch % 100 == 0: print(f"batch {batch} loss={logs['loss']:.4f}")
setup added so this can run · defines callbacks
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) callbacks = _AutoMock('callbacks')
9. Mixed Precision Training
Modern NVIDIA GPUs (Ampere and later) have tensor cores that run float16 (or bfloat16) matmuls ~2× faster than float32, with half the memory. Mixed precision keeps weights in float32 (for stable updates) but does compute in float16. One line to enable:
from tensorflow.keras import mixed_precision mixed_precision.set_global_policy("mixed_float16") # weights fp32, compute fp16 # Build the model AFTER setting the policy. model = build_model()
setup added so this can run · defines build_model
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) def build_model(*_a, **_kw): print('-> build_model() called') return _AutoMock('build_model()')
One subtle requirement: the output layer should be float32 to avoid numerical issues in the loss:
outputs = layers.Dense(num_classes, dtype="float32")(x) # force fp32 output
setup added so this can run · defines x, num_classes, layers
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) x = _AutoMock('x') num_classes = _AutoMock('num_classes') layers = _AutoMock('layers')
Mixed precision gives roughly a 2× speedup and lets you fit a larger batch in the same memory — for free, on the right hardware. On older GPUs (Pascal and earlier) without tensor cores, it gives little or no speedup.
bfloat16 (Brain Float 16) is similar but has the dynamic range of float32 with the precision of float16 — fewer numerical issues. Available on TPUs and Ampere+ GPUs. Use "mixed_bfloat16" on TPUs.
10. Distributed Training
Single-GPU isn't enough? tf.distribute.MirroredStrategy replicates the model across all local GPUs and averages gradients:
strategy = tf.distribute.MirroredStrategy() print(f"Replicas: {strategy.num_replicas_in_sync}") # e.g. 4 with strategy.scope(): model = build_model() # weights created on each GPU model.compile(optimizer="adamw", loss="...") model.fit(train_ds, epochs=10) # data sharded automatically
setup added so this can run · defines train_ds, build_model, tf
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) train_ds = _AutoMock('train_ds') def build_model(*_a, **_kw): print('-> build_model() called') return _AutoMock('build_model()') tf = _AutoMock('tf')
The with strategy.scope() block is mandatory — model creation must happen inside so weights are mirrored. Multi-host: MultiWorkerMirroredStrategy. TPUs: TPUStrategy. The Keras API is identical; only the strategy object changes.
Effective batch size scales with replica count. If you used batch=64 on one GPU and now run on 4 GPUs, the effective batch is 256 — you typically scale the LR linearly (lr * 4) to compensate, then add a warmup.
TPUs (Google Cloud) are great for very large transformer training but have rough edges: bfloat16 is mandatory for performance, dynamic shapes are slow, and debugging is harder. Use them when you actually have the workload to justify them.
11. Saving Formats — How You'll Actually Deploy
Keras has four save formats. Pick the right one for the deployment target.
| Format | File / Dir | Use case |
|---|---|---|
.keras | single ZIP file | Default. Architecture + weights + optimiser state. Use for checkpointing and most reloading. |
| SavedModel | directory tree | TensorFlow Serving, TF.js conversion, TFLite conversion. Contains the graph, weights, and signature defs. |
| TFLite | .tflite file | Mobile and edge (Android, iOS, embedded). Quantised, optimised for small runtime. |
| TF.js | JSON + binary shards | Browser inference. Convert from SavedModel via tensorflowjs_converter. |
# .keras — preferred for Python reloading model.save("models/sentiment.keras") loaded = keras.models.load_model("models/sentiment.keras") # SavedModel — for TensorFlow Serving model.export("models/sentiment_savedmodel") # Keras 3 API # TFLite — for mobile converter = tf.lite.TFLiteConverter.from_saved_model("models/sentiment_savedmodel") tflite_bytes = converter.convert() with open("sentiment.tflite", "wb") as f: f.write(tflite_bytes)
setup added so this can run · defines model, keras, tf
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) model = _AutoMock('model') keras = _AutoMock('keras') tf = _AutoMock('tf')
save_weights_only=True in ModelCheckpoint saves only weights, not architecture — fine if you can rebuild the model identically (same code, same seed), risky otherwise. For a model you'll reload after the codebase has moved on, save the full model.
12. Keras vs PyTorch — When Each Is Right
| Dimension | Keras (TF) | PyTorch |
|---|---|---|
| Default user | Production engineers, applied ML, fast prototypes | Researchers, paper authors |
| API style | Declarative (.compile, .fit) plus subclassing | Imperative, more "just Python" |
| Production deployment | TF Serving, TFLite, TF.js — best-in-class | TorchServe, ONNX, mobile via Torch Mobile |
| Distributed training | tf.distribute strategies — clean abstraction | torch.distributed / DDP — more boilerplate |
| Research ecosystem | Smaller — most new papers ship PyTorch first | Dominant — HuggingFace, timm, MMDetection |
| Static graph optimisation | tf.function — opt-in | torch.compile — Torch 2.0+ caught up |
| TPU support | First-class | Newer, less mature |
In 2026 the lines are blurrier than they were in 2020 — both frameworks have closed the gap on each other's strengths. Practical rule: if you're building a product to deploy, Keras/TF. If you're iterating on research and porting from papers, PyTorch. Most teams I've seen use one in production and let researchers prototype in either, with a conversion step (often via ONNX) at the handoff.
Common Mistakes
1. Looping in Python over a Dataset
for x, y in train_ds: # fine in a Python loop? grads = compute(x, y) # no — defeats prefetch parallelism
setup added so this can run · defines train_ds, compute
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) train_ds = [("alpha", 1), ("beta", 2), ("gamma", 3)] def compute(*_a, **_kw): print('-> compute() called') return _AutoMock('compute()')
Iterating a tf.data.Dataset in Python is correct but disables much of the prefetch parallelism. Wrap the inner loop in @tf.function, or hand the dataset directly to .fit(), and TF runs the iteration as a graph op overlapped with compute.
2. Not using tf.function on hot paths
Your train_step runs 100,000 times. Without @tf.function, each call pays Python's overhead — easily 5-10× slower than the compiled version. Profile before assuming; but on hot paths, the decorator is free performance.
3. Mixing eager and graph code
Inside a @tf.function-traced function, Python print() runs at trace time only — not on subsequent calls. Use tf.print() for runtime output. Same trap with if over tensor values (use tf.cond), Python for over tensors (use tf.while_loop or vectorise), and list appending (use tf.TensorArray).
4. Saving weights when you needed architecture too
model.save_weights("ckpt.h5") # weights only — can't load without rebuilding # 3 months later: which version of the code defined this architecture?
setup added so this can run · defines model
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) model = _AutoMock('model')
Save the full model (.keras) unless you have a specific reason for weights-only. The extra disk space is worth the ability to reload without code archaeology.
5. Forgetting to set_global_policy before building the model
model = build_model() # weights created in fp32 mixed_precision.set_global_policy("mixed_float16") # too late
setup added so this can run · defines build_model, mixed_precision
# Lightweight mock for objects whose attributes/methods aren't critical class _AutoMock: def __init__(self, name='mock'): self._name = name def __getattr__(self, k): return _AutoMock(self._name + '.' + k) def __call__(self, *a, **kw): print('-> ' + self._name + '() called') return _AutoMock(self._name + '()') def __repr__(self): return '<mock ' + self._name + '>' def __str__(self): return '<mock ' + self._name + '>' def __bool__(self): return True def __iter__(self): return iter([]) def __len__(self): return 0 def __getitem__(self, k): return _AutoMock(self._name + '[...]') def __setitem__(self, k, v): pass def __enter__(self): return self def __exit__(self, *a): return False async def __aenter__(self): return self async def __aexit__(self, *a): return False def __add__(self, o): return self def __radd__(self, o): return self def __sub__(self, o): return self def __mul__(self, o): return self def __rmul__(self, o): return self def __truediv__(self, o): return self def __eq__(self, o): return isinstance(o, _AutoMock) def __hash__(self): return hash(self._name) def __lt__(self, o): return True def __le__(self, o): return True def __gt__(self, o): return False def __ge__(self, o): return False def __mro_entries__(self, bases): return (object,) def build_model(*_a, **_kw): print('-> build_model() called') return _AutoMock('build_model()') mixed_precision = _AutoMock('mixed_precision')
The policy is read at layer construction. Set it first, then build. Same trap with distribution strategies — with strategy.scope(): must wrap the model construction.
6. Calling .numpy() inside the training loop
Even for logging, calling .numpy() every step forces a host-device sync. Either log per-N-steps, or accumulate with a tf.keras.metrics.Mean() object inside the graph and read it back once per epoch.
🎯 Your Turn — A Full tf.data Pipeline
Build a tf.data pipeline that takes a list of image file paths and corresponding integer labels, and produces shuffled, batched, prefetched (image, label) pairs ready for .fit().
Requirements:
- Accept
file_paths(list of strings) andlabels(list of ints). - Each step in the pipeline:
uint8 tensor.
- Resize to (224, 224).
- Cast to float32 and scale to [0, 1].
- Shuffle with a buffer of 1000.
- Batch into 32, dropping the remainder.
- Prefetch with
AUTOTUNE. - Use
num_parallel_calls=AUTOTUNEon the decoding map.
import tensorflow as tf def build_pipeline(file_paths, labels, batch_size=32): # TODO 1: create a Dataset from the (paths, labels) pair # TODO 2: shuffle with buffer 1000 # TODO 3: map a function that loads + decodes + resizes + normalises # use num_parallel_calls=tf.data.AUTOTUNE # TODO 4: batch with drop_remainder=True # TODO 5: prefetch with AUTOTUNE ... # Usage paths = ["data/cat1.jpg", "data/dog1.jpg", "data/cat2.jpg", ...] labels = [0, 1, 0, ...] ds = build_pipeline(paths, labels, batch_size=32) for images, lbls in ds.take(1): print(images.shape, images.dtype) # (32, 224, 224, 3) float32 print(lbls.shape, lbls.dtype) # (32,) int32
Hint 1 — A function that takes a path and label, returns image and label
tf.data.Dataset.map calls your function on each element. If the dataset yields (path, label) pairs, the mapped function signature is def load(path, label):. Decode the image inside; return (image, label) unchanged. The label passes through untouched.
Hint 2 — Decoding JPEGs in graph mode
Usetf.io.read_file(path) to read bytes from disk (works inside tf.data), then tf.io.decode_jpeg(bytes, channels=3) to get a uint8 tensor of shape (H, W, 3). Then tf.image.resize for the resize (returns float32 automatically) and divide by 255.0 for the normalisation. Don't use PIL or OpenCV inside the map — they break graph tracing.
Show full solution
import tensorflow as tf AUTOTUNE = tf.data.AUTOTUNE def load_and_preprocess(path, label): """Read JPEG bytes, decode, resize, normalise.""" raw = tf.io.read_file(path) image = tf.io.decode_jpeg(raw, channels=3) # uint8 (H, W, 3) image = tf.image.resize(image, [224, 224]) # float32, may need cast image = tf.cast(image, tf.float32) / 255.0 # explicit cast + scale return image, label def build_pipeline(file_paths, labels, batch_size=32): ds = tf.data.Dataset.from_tensor_slices((file_paths, labels)) ds = ds.shuffle(buffer_size=1000, seed=42, reshuffle_each_iteration=True) ds = ds.map(load_and_preprocess, num_parallel_calls=AUTOTUNE) ds = ds.batch(batch_size, drop_remainder=True) ds = ds.prefetch(AUTOTUNE) return ds # Demo with synthetic paths/labels (requires real JPEGs on disk to actually iterate) paths = tf.constant(["/path/to/img1.jpg", "/path/to/img2.jpg"]) labels = tf.constant([0, 1], dtype=tf.int32) ds = build_pipeline(paths, labels, batch_size=2) # Expected output on real data: # (2, 224, 224, 3) float32 # (2,) int32
What makes this pipeline fast:
num_parallel_calls=AUTOTUNEruns JPEG decoding on multiple CPU threads in parallel — image decode is your main CPU cost, and it parallelises perfectly across files.prefetch(AUTOTUNE)lets TF buffer N batches ahead while the GPU is busy training on the current one. AUTOTUNE picks N adaptively based on observed throughput.drop_remainder=Truekeeps the batch size static — important for@tf.functiontracing (variable batch size triggers re-tracing) and required for TPU.shufflebeforemapbeforebatch— shuffling raw paths is cheap (just integers); shuffling decoded images would be wasteful (moving big tensors around).
Real production additions you'd often layer on top:
- Caching:
.cache()after.map()if the preprocessed data fits in memory. Next epoch reads from RAM instead of disk. - Augmentation: a second
.map()after batching withtf.image.random_flip_left_right,tf.image.random_crop, etc. — runs on CPU in the pipeline, or move into a Keras preprocessing layer that runs on GPU. - Sharded reading from cloud:
tf.data.Dataset.list_files("gs://bucket/*.tfrecord")andtf.data.TFRecordDataset(...)for cloud-scale training.
If your GPU utilisation is below ~85% during training, the pipeline is the first place to look — usually a missing num_parallel_calls, a missing prefetch, or a .map step that's secretly running Python (PIL, OpenCV, slow custom ops).
What You Learned
- TF 2.x is eager by default;
@tf.functiontraces a function into a graph for speed. Wrap hot paths. tf.Tensor↔np.ndarrayconversion costs a device sync. Stay in tensor-land inside training loops;.numpy()only at boundaries.tf.datapipelines withshuffle → map(parallel) → batch → prefetchare the difference between 30% and 90% GPU utilisation.- Functional API for multi-input, multi-output, shared, and branched models. Sequential only for linear stacks. Subclassing for research-grade dynamic control flow.
- Custom layers subclass
keras.layers.Layer, overridebuildandcall.get_config()enables round-trip serialisation. tf.GradientTapefor custom training loops — required for GANs, RL, gradient accumulation, anything.fit()can't express.- Callbacks:
ModelCheckpoint(save_best_only=True),EarlyStopping(restore_best_weights=True),ReduceLROnPlateau,TensorBoard. Always restore best weights. - Mixed precision (
mixed_float16) is a free 2× speedup on modern GPUs. Force the output layer tofloat32. MirroredStrategyfor multi-GPU — wrap model construction instrategy.scope(). Scale LR with replica count.- Save as
.kerasfor Python reload, SavedModel for serving, TFLite for mobile, TF.js for browser. - Keras vs PyTorch: Keras for production and fast prototypes, PyTorch for research. Both are excellent in 2026.
Next: Convolutional Neural Networks — the architecture family that powers every modern vision system, with the maths, the architectures, and the modern tricks.