Deep Learning with TensorFlow and PyTorch

Training Loops and Batch Processing


The textbook version of gradient descent says: compute the gradient of the loss over the whole dataset, take one step, repeat. So do that. You have 50,000 training images at 224×224 pixels in colour.

Count the bytes. Each image is 224×224×3=150,528224 \times 224 \times 3 = 150{,}528 numbers. In float32 that is 602 KB. Fifty thousand of them is 30 gigabytes, and that is just the raw input. The forward pass has to keep every intermediate activation for the backward pass, which for a moderate convolutional network multiplies the requirement by another factor of five or more. You are asking a 24 GB GPU to hold something like 150 GB.

Suppose memory were free. You would still be in trouble, because full-batch descent gives you one weight update per pass over the data. Thirty epochs means thirty updates. Thirty updates will not train anything. Meanwhile a mini-batch of 32 gives you 1,562 updates per epoch — nearly fifty thousand updates over the same thirty passes, at a fraction of the memory.

Mini-batching is not an approximation you tolerate. It is faster and it uses less memory and it usually generalises better. It is one of the rare cases where the practical compromise beats the theoretically pure version on every axis.

Four lines, and the four ways they go wrongzero_grad — elsegradients accumulateforward — getthe predictionsloss.backward— fill the.grad fieldsoptimizer.step— apply themForgetting zero_grad does not crash; it silently trains on a sum of every past batch.
The order is not stylistic — swap step and backward and the model updates on last batch's gradients forever.

The three-way trade-off

Batch sizeGradient estimateMemoryHardware useGeneralisation
1Very noisyMinimalTerrible — GPU mostly idleOften good, but training is slow and unstable
32–128Noisy but usableModestGoodUsually best
512–2048SmoothLargeExcellentOften slightly worse without extra tuning
Whole datasetExactImpossibleN/ATends towards sharp minima

The generalisation column surprises people. Larger batches give a better gradient estimate and often a worse final model. The working explanation is that the noise in a small-batch gradient acts as a mild regulariser: it prevents the optimiser from settling into a narrow, sharp minimum that fits the training data precisely and collapses on anything slightly different. A noisy optimiser gets shaken out of those and ends up in wider, flatter basins that tolerate small shifts in the data.

If you do increase the batch size for throughput reasons, increase the learning rate too. The common heuristic is linear scaling: double the batch, double the learning rate, and add a few hundred steps of learning-rate warm-up at the start so the first large steps do not destabilise things.

A default of 32 is defensible without thought. Change it because you measured something, not because a bigger number felt better.

The anatomy of a PyTorch training loop

Here is a complete, correct loop with every line's purpose stated. This is the shape you will write hundreds of times.

Python
import torchimport torch.nn as nndevice = torch.device("cuda" if torch.cuda.is_available() else "cpu")model = MyModel().to(device)opt = torch.optim.AdamW(model.parameters(), lr=1e-3)loss_fn = nn.CrossEntropyLoss()for epoch in range(num_epochs):    # ---------- TRAIN ----------    model.train()                       # dropout ON, batchnorm updates its stats    running, seen = 0.0, 0    for xb, yb in train_loader:        xb, yb = xb.to(device), yb.to(device)        opt.zero_grad()                 # 1. clear gradients from the last step        logits = model(xb)              # 2. forward pass        loss = loss_fn(logits, yb)      # 3. how wrong were we        loss.backward()                 # 4. compute d(loss)/d(every parameter)        opt.step()                      # 5. move each parameter downhill        running += loss.item() * xb.size(0)        seen += xb.size(0)    train_loss = running / seen    # ---------- VALIDATE ----------    model.eval()                        # dropout OFF, batchnorm uses running stats    correct, total, vloss = 0, 0, 0.0    with torch.no_grad():               # do not build a graph; saves memory and time        for xb, yb in val_loader:            xb, yb = xb.to(device), yb.to(device)            logits = model(xb)            vloss += loss_fn(logits, yb).item() * xb.size(0)            correct += (logits.argmax(1) == yb).sum().item()            total += xb.size(0)    print(f"epoch {epoch:3d}  train {train_loss:.4f}  "          f"val {vloss/total:.4f}  acc {correct/total:.2%}")

Four lines, four ways to get it wrong

MistakeWhat happensHow it looks
Omit opt.zero_grad()Gradients accumulate across stepsLoss diverges after a few dozen steps as if the LR were growing
Call opt.step() before loss.backward()Steps using stale or zero gradientsModel barely improves; looks like a too-small learning rate
Accumulate loss instead of loss.item()Keeps the whole computation graph alive for every batchMemory grows every iteration until CUDA out-of-memory
Forget model.eval()Dropout stays active; batchnorm keeps updating on validation dataValidation numbers are noisy, pessimistic, and irreproducible

The third one deserves emphasis because the symptom appears far from the cause. loss is a tensor attached to the graph that produced it. Adding it to a running total keeps a reference, so the graph — activations and all — cannot be freed. After 200 batches you are holding 200 complete forward passes in GPU memory. .item() extracts a plain Python float and lets everything go.

train() and eval() are not cosmetic

Two layer types behave differently in the two modes:

  • Dropout zeroes a random fraction of activations during training and does nothing at inference. Leave it on during validation and your reported accuracy is measured on a randomly crippled model — typically a point or two below the truth, and different every time you run it.
  • Batch normalisation normalises using the current batch's statistics during training, and using running averages accumulated over training at inference. Leave it in training mode during evaluation and your predictions depend on which other examples happen to share the batch — so the same input gets different predictions depending on its neighbours. With a batch size of 1 it breaks outright, because the variance of a single example is zero.

torch.no_grad() is separate and equally important. It tells autograd not to record operations, which cuts memory roughly in half and speeds up evaluation. It does not disable dropout — that is what eval() is for. You need both.

Feeding the loop: data pipelines

Your GPU can process batches faster than a naive Python loop can produce them. If loading is single-threaded and synchronous, the GPU spends most of its time waiting.

Python
from torch.utils.data import Dataset, DataLoaderclass TabularDataset(Dataset):    def __init__(self, X, y):        self.X = torch.tensor(X, dtype=torch.float32)        self.y = torch.tensor(y, dtype=torch.long)    def __len__(self):        return len(self.X)                 # how many examples exist    def __getitem__(self, idx):        return self.X[idx], self.y[idx]    # produce ONE exampletrain_loader = DataLoader(    TabularDataset(X_train, y_train),    batch_size=32,    shuffle=True,        # reshuffle every epoch -- essential for training    num_workers=4,       # 4 background processes preparing batches    pin_memory=True,     # faster host-to-GPU copies    drop_last=True,      # discard a final partial batch (helps batchnorm))val_loader = DataLoader(TabularDataset(X_val, y_val),                        batch_size=256, shuffle=False)   # never shuffle validation

A Dataset only needs to answer two questions: how many items are there, and give me item ii. The DataLoader handles batching, shuffling, and parallelism on top of that. Note shuffle=True for training only — without it, if your data happens to be sorted by class, every batch contains one class and the gradient points somewhere useless.

The TensorFlow equivalent chains transformations:

Python
import tensorflow as tftrain_ds = (tf.data.Dataset.from_tensor_slices((X_train, y_train))            .shuffle(buffer_size=10_000)          # bigger buffer = better mixing            .batch(32, drop_remainder=True)            .prefetch(tf.data.AUTOTUNE))          # overlap loading with computeval_ds = (tf.data.Dataset.from_tensor_slices((X_val, y_val))          .batch(256)          .prefetch(tf.data.AUTOTUNE))

The order of those calls matters. .shuffle() before .batch() shuffles examples, which is what you want; after .batch() it shuffles the order of fixed batches, so the same examples travel together every epoch. And .prefetch() should always be last — it tells the pipeline to prepare batch n+1n+1 while the GPU is still working on batch nn, which is frequently a 30–50% throughput win for free.

Batch shapes, and the axis that trips people

Data typePyTorch shapeTensorFlow shape
Tabular(batch, features)(batch, features)
Images(batch, channels, height, width)(batch, height, width, channels)
Sequences(batch, seq_len, features)(batch, seq_len, features)

Images are the trap. PyTorch puts channels second (NCHW); TensorFlow puts them last (NHWC). Move a preprocessing function between frameworks without transposing and you get a model that trains on a tensor where the "height" axis is actually the three colour channels. It will not error. It will just perform terribly, and the reason is invisible in the loss curve.

Python
x_tf = torch_tensor.permute(0, 2, 3, 1)     # NCHW -> NHWCx_pt = tf.transpose(tf_tensor, perm=[0, 3, 1, 2])   # NHWC -> NCHW

Gradient accumulation: a large batch on a small GPU

Suppose a paper's results need a batch of 128 and your GPU fits 32. You can simulate the large batch by running four small ones and only stepping once:

Python
accum_steps = 4                       # 32 x 4 = effective batch of 128opt.zero_grad()for i, (xb, yb) in enumerate(train_loader):    xb, yb = xb.to(device), yb.to(device)    loss = loss_fn(model(xb), yb) / accum_steps   # scale so gradients AVERAGE    loss.backward()                               # gradients accumulate in .grad    if (i + 1) % accum_steps == 0:        opt.step()        opt.zero_grad()

The division by accum_steps is the part people leave out. Without it, four accumulated gradients sum rather than average, so the effective step is four times too large — the same explosion you would get from quadrupling the learning rate. And note that this simulates a large batch for the optimiser but not for batch normalisation, which still sees only 32 examples at a time when computing its statistics.

Mixed precision: nearly free speed

Modern GPUs compute in 16-bit floating point far faster than in 32-bit. Mixed precision runs most operations in 16-bit while keeping a 32-bit master copy of the weights and scaling the loss to stop small gradients underflowing to zero.

Python
from torch.amp import autocast, GradScalerscaler = GradScaler("cuda")for xb, yb in train_loader:    xb, yb = xb.to(device), yb.to(device)    opt.zero_grad()    with autocast("cuda", dtype=torch.float16):        loss = loss_fn(model(xb), yb)     # forward runs in fp16 where safe    scaler.scale(loss).backward()         # scale up so small grads survive fp16    scaler.step(opt)                      # unscale, then step (skips if inf/nan)    scaler.update()                       # adapt the scale factor

On GPUs that support bfloat16 (NVIDIA Ampere and newer), autocast("cuda", dtype=torch.bfloat16) is usually the simpler choice: bfloat16 has the same exponent range as float32, so small gradients do not underflow and you can drop the GradScaler entirely.

Python
import tensorflow as tftf.keras.mixed_precision.set_global_policy("mixed_float16")# Keras handles loss scaling internally when you use model.fit()

Expect roughly a 1.5–3× speed-up and about half the activation memory, with accuracy usually indistinguishable. It is one of the highest-value four-line changes available.

Keras: fit() and when to abandon it

Most of the loop above is written for you:

Python
history = model.fit(    train_ds,    validation_data=val_ds,    epochs=50,    callbacks=[        tf.keras.callbacks.EarlyStopping(patience=5, restore_best_weights=True),        tf.keras.callbacks.ReduceLROnPlateau(factor=0.5, patience=3),    ],)

When you need something fit() cannot express — two optimisers, an adversarial objective, a custom gradient manipulation — you do not have to abandon Keras. Override one method:

Python
class CustomTrainStep(tf.keras.Model):    def train_step(self, data):        x, y = data        with tf.GradientTape() as tape:            y_pred = self(x, training=True)            loss = self.compute_loss(x, y, y_pred)        grads = tape.gradient(loss, self.trainable_variables)        grads = [tf.clip_by_norm(g, 1.0) for g in grads]     # custom step        self.optimizer.apply_gradients(zip(grads, self.trainable_variables))        for metric in self.metrics:                          # includes the loss tracker            if metric.name == "loss":                metric.update_state(loss)            else:                metric.update_state(y, y_pred)        return {m.name: m.result() for m in self.metrics}

You keep callbacks, progress bars and metric tracking, and you replace only the part you needed to change.

Reading a training run that has gone wrong

Most training failures announce themselves clearly if you know the vocabulary.

Loss becomes nan. Reduce the learning rate by a factor of ten first — that fixes it most of the time. If not, look for a log⁡(0)\log(0) (a manual softmax followed by a log, or a division by a count that can be zero), an unnormalised input feature in the thousands, or exploding gradients. Add torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) between backward() and step().

Loss does not move at all. Confirm a parameter actually changes: print next(model.parameters()).flatten()[:3] before and after one opt.step(). If it does not move, the optimiser is not connected to the model — usually because it was built for a different model object (for example, one created before you rebuilt the model), or because layers live in a plain Python list and were never registered.

Training is far slower than expected. Watch GPU utilisation with nvidia-smi -l 1. Utilisation bouncing between 0% and 100% means the data pipeline is the bottleneck: raise num_workers, add pin_memory=True, or add .prefetch(). Steady high utilisation means the model itself is the cost, and mixed precision or a larger batch is the lever.

Validation loss is worse than training loss from the very first epoch. That is normal and expected — a small gap is healthy. A large gap that widens every epoch is overfitting. A validation loss that is better than training loss usually means dropout is inflating the training figure, which is also normal since training loss is measured with dropout active.

The habit that repays itself fastest is to run a deliberate overfitting test before any real training: take eight examples, turn off shuffling and regularisation, and train until the loss is essentially zero. If a model cannot memorise eight examples, it will not learn 50,000, and you have just found a bug in twenty seconds instead of an hour.