In PyTorch, in what order do zero_grad, forward, backward and step go?
answer
- five calls, one per batch
- clear before you compute
- the update reads what backward wrote
- set_to_none is the modern default
- scheduler and clipping have fixed slots
basics
~10 sPer 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 sPyTorch 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 linesimport 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
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.
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.
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.
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