Course Content
Deep Learning with TensorFlow and PyTorch
4 sections · 15 lessons
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,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.
The three-way trade-off
| Batch size | Gradient estimate | Memory | Hardware use | Generalisation |
|---|---|---|---|---|
| 1 | Very noisy | Minimal | Terrible — GPU mostly idle | Often good, but training is slow and unstable |
| 32–128 | Noisy but usable | Modest | Good | Usually best |
| 512–2048 | Smooth | Large | Excellent | Often slightly worse without extra tuning |
| Whole dataset | Exact | Impossible | N/A | Tends 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.
1import torch2import torch.nn as nn34device = torch.device("cuda" if torch.cuda.is_available() else "cpu")5model = MyModel().to(device)6opt = torch.optim.AdamW(model.parameters(), lr=1e-3)7loss_fn = nn.CrossEntropyLoss()89for epoch in range(num_epochs):10 # ---------- TRAIN ----------11 model.train() # dropout ON, batchnorm updates its stats12 running, seen = 0.0, 013 for xb, yb in train_loader:14 xb, yb = xb.to(device), yb.to(device)1516 opt.zero_grad() # 1. clear gradients from the last step17 logits = model(xb) # 2. forward pass18 loss = loss_fn(logits, yb) # 3. how wrong were we19 loss.backward() # 4. compute d(loss)/d(every parameter)20 opt.step() # 5. move each parameter downhill2122 running += loss.item() * xb.size(0)23 seen += xb.size(0)24 train_loss = running / seen2526 # ---------- VALIDATE ----------27 model.eval() # dropout OFF, batchnorm uses running stats28 correct, total, vloss = 0, 0, 0.029 with torch.no_grad(): # do not build a graph; saves memory and time30 for xb, yb in val_loader:31 xb, yb = xb.to(device), yb.to(device)32 logits = model(xb)33 vloss += loss_fn(logits, yb).item() * xb.size(0)34 correct += (logits.argmax(1) == yb).sum().item()35 total += xb.size(0)3637 print(f"epoch {epoch:3d} train {train_loss:.4f} "38 f"val {vloss/total:.4f} acc {correct/total:.2%}")Four lines, four ways to get it wrong
| Mistake | What happens | How it looks |
|---|---|---|
Omit opt.zero_grad() | Gradients accumulate across steps | Loss diverges after a few dozen steps as if the LR were growing |
Call opt.step() before loss.backward() | Steps using stale or zero gradients | Model barely improves; looks like a too-small learning rate |
Accumulate loss instead of loss.item() | Keeps the whole computation graph alive for every batch | Memory grows every iteration until CUDA out-of-memory |
Forget model.eval() | Dropout stays active; batchnorm keeps updating on validation data | Validation 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.
1from torch.utils.data import Dataset, DataLoader23class TabularDataset(Dataset):4 def __init__(self, X, y):5 self.X = torch.tensor(X, dtype=torch.float32)6 self.y = torch.tensor(y, dtype=torch.long)78 def __len__(self):9 return len(self.X) # how many examples exist1011 def __getitem__(self, idx):12 return self.X[idx], self.y[idx] # produce ONE example1314train_loader = DataLoader(15 TabularDataset(X_train, y_train),16 batch_size=32,17 shuffle=True, # reshuffle every epoch -- essential for training18 num_workers=4, # 4 background processes preparing batches19 pin_memory=True, # faster host-to-GPU copies20 drop_last=True, # discard a final partial batch (helps batchnorm)21)22val_loader = DataLoader(TabularDataset(X_val, y_val),23 batch_size=256, shuffle=False) # never shuffle validationA Dataset only needs to answer two questions: how many items are there, and give me item i. 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:
1import tensorflow as tf23train_ds = (tf.data.Dataset.from_tensor_slices((X_train, y_train))4 .shuffle(buffer_size=10_000) # bigger buffer = better mixing5 .batch(32, drop_remainder=True)6 .prefetch(tf.data.AUTOTUNE)) # overlap loading with compute78val_ds = (tf.data.Dataset.from_tensor_slices((X_val, y_val))9 .batch(256)10 .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+1 while the GPU is still working on batch n, which is frequently a 30–50% throughput win for free.
Batch shapes, and the axis that trips people
| Data type | PyTorch shape | TensorFlow 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.
x_tf = torch_tensor.permute(0, 2, 3, 1) # NCHW -> NHWCx_pt = tf.transpose(tf_tensor, perm=[0, 3, 1, 2]) # NHWC -> NCHWGradient 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:
1accum_steps = 4 # 32 x 4 = effective batch of 1282opt.zero_grad()34for i, (xb, yb) in enumerate(train_loader):5 xb, yb = xb.to(device), yb.to(device)6 loss = loss_fn(model(xb), yb) / accum_steps # scale so gradients AVERAGE7 loss.backward() # gradients accumulate in .grad89 if (i + 1) % accum_steps == 0:10 opt.step()11 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.
1from torch.amp import autocast, GradScaler23scaler = GradScaler("cuda")45for xb, yb in train_loader:6 xb, yb = xb.to(device), yb.to(device)7 opt.zero_grad()8 with autocast("cuda", dtype=torch.float16):9 loss = loss_fn(model(xb), yb) # forward runs in fp16 where safe10 scaler.scale(loss).backward() # scale up so small grads survive fp1611 scaler.step(opt) # unscale, then step (skips if inf/nan)12 scaler.update() # adapt the scale factorOn 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.
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:
1history = model.fit(2 train_ds,3 validation_data=val_ds,4 epochs=50,5 callbacks=[6 tf.keras.callbacks.EarlyStopping(patience=5, restore_best_weights=True),7 tf.keras.callbacks.ReduceLROnPlateau(factor=0.5, patience=3),8 ],9)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:
1class CustomTrainStep(tf.keras.Model):2 def train_step(self, data):3 x, y = data4 with tf.GradientTape() as tape:5 y_pred = self(x, training=True)6 loss = self.compute_loss(x, y, y_pred)7 grads = tape.gradient(loss, self.trainable_variables)8 grads = [tf.clip_by_norm(g, 1.0) for g in grads] # custom step9 self.optimizer.apply_gradients(zip(grads, self.trainable_variables))10 for metric in self.metrics: # includes the loss tracker11 if metric.name == "loss":12 metric.update_state(loss)13 else:14 metric.update_state(y, y_pred)15 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) (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.