Autograd
PyTorch autograd automatically computes gradients of scalar outputs with respect to tensor inputs. It powers all neural network training through reverse-mode automatic differentiation.
Search across all documentation pages
PyTorch autograd automatically computes gradients of scalar outputs with respect to tensor inputs. It powers all neural network training through reverse-mode automatic differentiation.
Quick-reference recipe card - copy-paste ready.
import torch
x = torch.tensor(2.0, requires_grad=True)
y = x ** 2 + 3 * x
y.backward()
print(x.grad) # dy/dx at x=2: 2*2+3 = 7
# Inference: no gradient tracking
with torch.no_grad():
pred = model(x)When to reach for this:
loss.backward() and optimizer.step() work."""autograd.py - gradient computation and gradient checking."""
import torch
import torch.nn as nn
# Simple computation graph
w = torch.tensor([1.0, 2.0], requires_grad=True)
b = torch.tensor(0.5, requires_grad=True)
x = torch.tensor([3.0, 4.0])
# Forward: y = w . x + b, loss = y^2
y = (w * x).sum() + b
loss = y ** 2
loss.backward()
print("w.grad:", w.grad) # d(loss)/dw
print("b.grad:", b.grad) # d(loss)/db
# Model with autograd
model = nn.Linear(3, 1)
x_batch = torch.randn(16, 3)
target = torch.randn(16, 1)
model.zero_grad(set_to_none=True)
pred = model(x_batch)
loss = nn.functional.mse_loss(pred, target)
loss.backward()
for name, param in model.named_parameters():
print(f"{name}: grad norm = {param.grad.norm().item():.4f}")What this demonstrates:
w, b) accumulate gradients in .grad.loss.backward() traverses the graph in reverse.zero_grad(set_to_none=True) clears gradients efficiently.backward() applies the chain rule from the output back to leaves.requires_grad=True accumulate .grad.retain_graph=True.torch.inference_mode() is stricter and faster than no_grad() for inference.| Context | Effect | Use |
|---|---|---|
requires_grad=True | Track operations | Model parameters, inputs needing grad |
torch.no_grad() | Disable tracking | Inference, metric computation |
torch.inference_mode() | Stricter no-grad | Production inference |
.detach() | Break from graph | Stop gradients through a tensor |
retain_graph=True | Keep graph after backward | Multiple backward calls |
# Custom autograd function
class MyReLU(torch.autograd.Function):
@staticmethod
def forward(ctx, input):
ctx.save_for_backward(input)
return input.clamp(min=0)
@staticmethod
def backward(ctx, grad_output):
input, = ctx.saved_tensors
grad_input = grad_output.clone()
grad_input[input < 0] = 0
return grad_inputzero_grad() - gradients accumulate across steps. Fix: call optimizer.zero_grad(set_to_none=True) before each backward.loss.backward(retain_graph=True) or separate forward passes.x.add_(1) on tracked tensors; use x = x + 1..item() for logging; del loss after backward..detach().cpu().numpy() for export.| Alternative | Use When | Don't Use When |
|---|---|---|
| PyTorch autograd | Default for neural nets | You need symbolic math (use JAX) |
torch.func (functorch) | Per-sample gradients, vmap | Simple model training |
| Manual gradients | Teaching, verification | Production model code |
| JAX autograd | Functional transforms, TPU | Existing PyTorch ecosystem |
backward() computes gradients of a scalar output.y.backward(gradient=torch.ones_like(y)).inference_mode disables more autograd infrastructure - faster.inference_mode for production inference.no_grad is fine for validation loops during training.for name, p in model.named_parameters():
if p.grad is None:
print(f"no grad: {name}")param.requires_grad = False or model.freeze() patterns.torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)optimizer.step().loss.backward(create_graph=True)
grad_grad = torch.autograd.grad(loss, model.parameters(), create_graph=True)create_graph=True keeps the graph for higher-order gradients.optimizer.zero_grad(set_to_none=True) sets .grad to None instead of zeroing..grad populated by default.retain_grad() on non-leaf tensors if needed.with torch.no_grad():
features = frozen_encoder(x)
output = trainable_head(features)Stack versions: This page was written for Python 3.14.0 (stable 3.14, maintenance 3.13), FastAPI 0.115+, Django 5.2, Flask 3.1, Pydantic 2, PyTorch 2.6+, pandas 2.2+, Polars 1.x, ruff 0.9+, and uv 0.6+.
Reviewed by Chris St. John·Last updated Jul 16, 2026