PythonMastery
advanced 28 min read · lesson 2 of 9 in AI & Deep Learning

TensorFlow & Keras: Production Engine

1 · The lesson

read

Keras'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.

python
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.

python
@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.

python
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.

python
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. Without num_parallel_calls, it runs serially — single-threaded preprocessing is often the bottleneck.
  • batch before prefetch — 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():

python
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.

python
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():

python
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:

python
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:

python
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).

python
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:

python
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():

python
@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:

python
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.

python
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:

python
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:

python
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:

python
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.

FormatFile / DirUse case
.kerassingle ZIP fileDefault. Architecture + weights + optimiser state. Use for checkpointing and most reloading.
SavedModeldirectory treeTensorFlow Serving, TF.js conversion, TFLite conversion. Contains the graph, weights, and signature defs.
TFLite.tflite fileMobile and edge (Android, iOS, embedded). Quantised, optimised for small runtime.
TF.jsJSON + binary shardsBrowser inference. Convert from SavedModel via tensorflowjs_converter.
python
# .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

DimensionKeras (TF)PyTorch
Default userProduction engineers, applied ML, fast prototypesResearchers, paper authors
API styleDeclarative (.compile, .fit) plus subclassingImperative, more "just Python"
Production deploymentTF Serving, TFLite, TF.js — best-in-classTorchServe, ONNX, mobile via Torch Mobile
Distributed trainingtf.distribute strategies — clean abstractiontorch.distributed / DDP — more boilerplate
Research ecosystemSmaller — most new papers ship PyTorch firstDominant — HuggingFace, timm, MMDetection
Static graph optimisationtf.function — opt-intorch.compile — Torch 2.0+ caught up
TPU supportFirst-classNewer, 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

python
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

python
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

python
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) and labels (list of ints).
  • Each step in the pipeline:
- Read JPEG bytes from disk. - Decode to a 3-channel 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=AUTOTUNE on the decoding map.
python
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 Use tf.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
python
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=AUTOTUNE runs 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=True keeps the batch size static — important for @tf.function tracing (variable batch size triggers re-tracing) and required for TPU.
  • shuffle before map before batch — 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 with tf.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") and tf.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.function traces a function into a graph for speed. Wrap hot paths.
  • tf.Tensor ↔ np.ndarray conversion costs a device sync. Stay in tensor-land inside training loops; .numpy() only at boundaries.
  • tf.data pipelines with shuffle → map(parallel) → batch → prefetch are 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, override build and call. get_config() enables round-trip serialisation.
  • tf.GradientTape for 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 to float32.
  • MirroredStrategy for multi-GPU — wrap model construction in strategy.scope(). Scale LR with replica count.
  • Save as .keras for 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.