skip to content

Training Loops and Optimizers

In PyTorch you write the loop yourself — zero_grad, forward, loss, backward, step — which means you also own schedulers, gradient clipping, mixed precision, and checkpointing via state_dict. Writing a correct training loop from memory is a standard interview exercise.

on this pageshow

questions

7

In PyTorch, in what order do zero_grad, forward, backward and step go?

level: juniorimportance: must knowfreq 88%

answer

  1. five calls, one per batch
  2. clear before you compute
  3. the update reads what backward wrote
  4. set_to_none is the modern default
  5. scheduler and clipping have fixed slots

basics

~10 s

Per batch: optimizer.zero_grad(), then the forward pass, then compute the loss, then loss.backward() to fill each parameter's .grad, then optimizer.step() to apply the update. Zeroing must precede backward; the update must follow it.

solid answer

~40 s

PyTorch gives you no `fit()`, so the loop is yours to write. For every batch: `optimizer.zero_grad()` clears the `.grad` attribute of every parameter the optimizer owns; the forward pass `outputs = model(x)` builds the graph; `loss = loss_fn(outputs, y)` reduces it to a scalar; `loss.backward()` populates `.grad`; `optimizer.step()` reads those gradients and updates the weights. Order matters in both directions. Backward **adds into** `.grad` rather than overwriting it, so skipping the zeroing silently mixes this batch's gradient with every previous one. And `step()` only reads `.grad`, so calling it before `backward()` either updates with stale gradients or does nothing at all. Since PyTorch 2.0, `zero_grad()` defaults to `set_to_none=True`, which sets `.grad` to `None` instead of writing zeros — cheaper, and the next `backward()` allocates fresh.

code

python · 17 lines
python
import torch
from torch import nn

model = nn.Linear(4, 2)
optimizer = torch.optim.SGD(model.parameters(), lr=0.1)
loss_fn = nn.CrossEntropyLoss()
x = torch.randn(8, 4)
y = torch.randint(0, 2, (8,))

model.train()
for _ in range(3):
    optimizer.zero_grad(set_to_none=True)
    logits = model(x)
    loss = loss_fn(logits, y)
    loss.backward()
    optimizer.step()
    print(loss.item())

go deeper

for a junior

Be ready to type the five calls in order without hesitation and to say plainly which one updates the weights. If you can only remember one rule, remember that zeroing comes before backward.

for a middle

Explain why the ordering is forced: backward adds into .grad, step only reads it. Know that set_to_none=True is the default and what that makes .grad between calls.

for a senior

Show where clipping, the scheduler, mixed precision and evaluation slot into the same loop, and name the silent failures — a missing zero_grad, a loss tensor accumulated for logging — that look like bad hyperparameters.

for a principal

Frame the tradeoff: a hand-written loop is why PyTorch research code is debuggable, and also why every team eventually rewrites the same loop. Have a view on when to standardize on an in-house trainer versus keeping the loop explicit.

## The loop PyTorch does not write for you High-level frameworks hide the weight update behind a single `fit()` call. PyTorch deliberately does not: the training loop is ordinary Python that you type out, which is why "write me a training loop" is one of the most common live-coding exercises for anyone who claims PyTorch on a CV. There are five calls per batch and each one has a job. ## The five calls **1. `optimizer.zero_grad(set_to_none=True)`** — clears the `.grad` attribute of every parameter that was handed to the optimizer when it was constructed. **2. `outputs = model(x)`** — the forward pass. Calling the module runs `forward()` and, for every operation on a tensor that requires gradients, records the operation so a backward pass is possible. **3. `loss = loss_fn(outputs, y)`** — reduces the batch to a single scalar. `backward()` can only start from a scalar (or you must supply an explicit gradient argument), so this reduction is not optional. **4. `loss.backward()`** — walks the recorded graph from the loss back to the leaves and writes each parameter's gradient into `param.grad`. **5. `optimizer.step()`** — iterates over its `param_groups`, reads each `param.grad`, and applies the update rule (plain SGD, momentum, Adam's moment estimates, and so on) to `param.data`. ## Why zeroing comes first The critical detail is that step 4 **accumulates**: PyTorch adds the newly computed gradient into whatever is already in `.grad`. That behaviour exists so several backward passes can contribute to one update, but it means that if you never clear `.grad`, batch 100's update is driven by the sum of gradients from batches 1 through 100. Nothing raises. The loss simply refuses to converge, or diverges, and the bug looks like a bad learning rate. Zeroing at the *top* of the loop rather than at the bottom is the convention because it is robust to `continue` statements and early exits inside the body. ## Why the step comes last `optimizer.step()` is a pure consumer of `.grad`. If you call it before `backward()`, either `.grad` is `None` (with `set_to_none=True` the parameter is skipped entirely and nothing happens) or it holds the previous iteration's gradient, and you take a step in a stale direction. You do not need to wrap `step()` in a no-grad context — the optimizer implementations already perform their updates with gradient tracking disabled. ## `set_to_none` and what `.grad` becomes Before PyTorch 2.0 the default was to fill the existing gradient tensors with zeros; since 2.0 the default is `set_to_none=True`, which drops the reference so the tensor can be freed and the next `backward()` allocates a new one. This is faster and saves memory, but it changes two observable things: reading `param.grad` between `zero_grad()` and `backward()` gives `None` rather than a zero tensor, and optimizers that would otherwise update a parameter on a zero gradient (momentum, weight decay) skip that parameter entirely for the step. If you have code that inspects or manipulates `.grad` directly, guard for `None`. ## Where everything else slots in The canonical loop has fixed insertion points for the extras: - **Gradient clipping** goes between `backward()` and `step()` — it edits `.grad` in place, so the gradients must exist and must not have been consumed yet. - **A learning-rate scheduler's `step()`** goes *after* `optimizer.step()`; PyTorch warns if it detects the reverse order. - **Mixed precision** wraps the forward and the loss in `autocast`, and routes backward and step through the gradient scaler. - **Logging** should use `loss.item()` or `loss.detach()`, not the loss tensor itself — keeping the tensor alive keeps its whole graph alive. - **`model.train()`** belongs before the epoch's batches, and `model.eval()` before validation. ## What interviewers watch for The two answers that end the question early are forgetting `zero_grad` entirely and putting `optimizer.step()` before `loss.backward()`. A close third is claiming that `backward()` updates the weights — it does not; it only computes gradients, and the optimizer is the only thing that touches parameter values.

  • Does optimizer.step() need to be wrapped in a no-grad context?
    No. The optimizer implementations in torch.optim already perform their parameter updates with gradient tracking disabled internally, so wrapping the call yourself changes nothing. You do need a no-grad or inference context around a *validation* forward pass, but that is a separate concern from the update itself.
  • Why does zero_grad live on the optimizer when nn.Module also has one?
    They clear different sets. optimizer.zero_grad() clears only the parameters in the optimizer's param_groups; model.zero_grad() clears every parameter of that module. They coincide in the common case, but diverge when you optimize a subset — for example a frozen backbone with only the head registered — where clearing on the model touches parameters the optimizer never updates.
  • Why log loss.item() instead of accumulating the loss tensor across batches?
    The loss tensor still holds a reference to the autograd graph that produced it. Appending it to a running total keeps every batch's graph alive, so memory grows until you run out. Calling .item() extracts a Python float and drops the reference; .detach() does the same while keeping a tensor.

Think of .grad as a whiteboard the backward pass writes on: wipe it first, write on it, then read it. Skip the wipe and you read today's note on top of last week's.

saying these in an interview costs you the question

  • Claiming loss.backward() updates the weights
  • Calling optimizer.step() before loss.backward()
  • Thinking backward() overwrites .grad instead of adding to it
  • Saying zero_grad is optional because gradients are recomputed
  • Believing set_to_none=False is still the default

context

open as a page

Why does PyTorch's nn.CrossEntropyLoss expect raw logits, not softmax output?

level: middleimportance: must knowfreq 62%

basics

~10 s

nn.CrossEntropyLoss applies log_softmax internally, so it must receive unnormalized scores. Feeding it softmax probabilities applies softmax twice, which flattens the distribution, shrinks gradients and slows or stalls training. No error is raised.

open as a page

What do model.train() and model.eval() change in a PyTorch model?

level: middleimportance: must knowfreq 76%

basics

~20 s

They flip a boolean flag that certain layers read. Dropout zeroes activations only in train mode; BatchNorm normalizes with batch statistics and updates its running averages in train, and uses the stored averages in eval. Neither call affects gradient tracking.

open as a page

Your PyTorch job resumed from a checkpoint and the loss jumped — what did it omit?

level: seniorimportance: must knowfreq 52%

basics

~20 s

Almost always the optimizer state. Saving only model.state_dict() discards Adam's moment estimates and momentum buffers, so the first steps after resume behave like a freshly initialized optimizer. A resumable checkpoint also needs the scheduler, the AMP scaler and the step counter.

open as a page

How do torch.amp.autocast and GradScaler fit into a PyTorch training loop?

level: middleimportance: should knowfreq 54%

basics

~20 s

Run only the forward pass and the loss under torch.amp.autocast, which executes eligible operations in a 16-bit type while the weights stay 32-bit. With float16, route the backward pass through torch.amp.GradScaler so small gradients do not underflow to zero.

open as a page

When do you call a PyTorch lr_scheduler's step() — every batch or every epoch?

level: middleimportance: should knowfreq 48%

basics

~20 s

It depends on the schedule. Epoch-granularity schedules such as StepLR are stepped once per epoch; warmup and OneCycleLR are stepped once per optimizer step. Either way the call goes after optimizer.step(), and ReduceLROnPlateau also needs a metric argument.

open as a page

Where does torch.nn.utils.clip_grad_norm_ go in a PyTorch loop, and what does it clip?

level: seniorimportance: should knowfreq 44%

basics

~20 s

Call it after loss.backward() and before optimizer.step(). It computes one global L2 norm across all the gradients you pass it and, if that norm exceeds max_norm, scales every gradient by the same factor so the combined norm equals max_norm.

open as a page