Writing a custom autograd Function
torch.autograd.Function lets you hand-write both the forward step and its backward rule — the escape hatch for operations autograd cannot, or should not, derive itself.
- 8 min read
- 3 reading levels
- Published
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
A custom autograd Function is you teaching PyTorch a new maths step: how to compute it forwards, and how to pass blame through it backwards.
Think of a new worker joining an assembly line. Two instructions make her a real part of the line. First: what to do with the parts arriving from the left. Second: when a defect report arrives from the right, how to pass it leftward — how much of the fault to forward to her suppliers. A worker who cannot pass reports breaks the whole line's ability to trace defects.
Every operation in PyTorch is such a worker, with both instructions built in. A custom Function is a worker you hire, writing both instructions yourself.
Why it exists
Autograd already knows the backward rule for every built-in step, and chains them for you automatically. Three situations fall outside that comfort:
- The step is not built from PyTorch operations at all — it calls outside code, like a hand-tuned GPU routine.
- The step is non-differentiable — like hard rounding — and you want to decide what blame should flow anyway.
- The automatic chain is wasteful or unstable, and a hand-derived rule is faster or safer.
For these, you write the two instructions explicitly.
How it works
forward: inputs -> [ your step ] -> outputs (may store notes)
|
notepad
|
backward: blame_in <- [ your rule ] <- blame_out (reads the notes)The notepad matters: your backward rule often needs values from the forward moment, and the notepad is the official place to keep them.
A real example you have seen
Compressed models on cheap phones — the ones that recognise faces offline — are trained with rounding baked into the network. Rounding gives zero useful blame (nudging a weight slightly changes nothing, then suddenly everything). The teams train them anyway, using exactly this tool: a custom step that rounds forwards but passes blame as if it had not. A white lie in the backward direction, and it works.
Remember this
- A custom Function = your forward + your backward, as one new operation.
- The notepad (
ctx) carries forward-time values to backward-time. - Main uses: outside code, non-differentiable steps, hand-optimised rules.
What to learn next
- Checking your gradients with gradcheck — the test your new Function must pass.
- Jacobians, vjp and jvp with torch.func — the transform world your Function should stay compatible with.
- detach, no_grad and inference_mode — the lighter tool that covers many "custom gradient" wishes.
Developer — Code and libraries.
Setup
pip install torchOutputs verified with torch 2.5.1, CPU.
A complete custom operation
import torch
class Cube(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
ctx.save_for_backward(x) # keep what backward will need
return x ** 3
@staticmethod
def backward(ctx, grad_out):
(x,) = ctx.saved_tensors
return grad_out * 3 * x ** 2 # incoming grad times local derivative
x = torch.tensor([2.0], requires_grad=True)
y = Cube.apply(x) # never call forward() yourself
y.backward()
print(y)
print(x.grad) # 3 * 2^2 = 12tensor([8.], grad_fn=<CubeBackward>) tensor([12.])
The walkthrough
Both methods are @staticmethods taking ctx first. ctx is the notepad — a context object autograd creates per call, connecting the forward moment to the backward moment. save_for_backward is the only sanctioned way to store tensors on it: it cooperates with the version-counter checks, so tampered saves still get caught. Non-tensor extras (shapes, flags) can ride as plain attributes, ctx.anything = value.
backward receives blame, returns blame. grad_out is the gradient of the loss with respect to your output — the incoming defect report. You multiply by your local derivative (d(x³)/dx = 3x²) and return the gradient with respect to your input. One returned value per forward argument; inputs that need no gradient get None.
Cube.apply(x), never Cube.forward(x). apply is what registers the node into the graph — notice grad_fn=<CubeBackward> on the output, your class woven into the recording like any built-in. Calling forward directly computes numbers and records nothing.
The white-lie pattern: straight-through rounding
The most copied custom Function in existence — rounding forwards, pretending identity backwards:
import torch
class RoundSTE(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
return torch.round(x) # nothing saved: backward needs nothing
@staticmethod
def backward(ctx, grad_out):
return grad_out # the lie: pass blame through unchanged
x = torch.tensor([0.4, 1.6, 2.5], requires_grad=True)
y = RoundSTE.apply(x)
y.sum().backward()
print(y)
print(x.grad) # honest rounding would give zeros heretensor([0., 2., 2.], grad_fn=<RoundSTEBackward>) tensor([1., 1., 1.])
The true derivative of rounding is zero almost everywhere — training would freeze. The pretend-identity keeps blame moving, and quantised networks train on exactly this trick.
Common mistakes
Storing tensors as ctx.x = x instead of save_for_backward. It appears to work and skips the safety net: in-place edits to a stashed tensor go undetected, memory can leak through reference cycles, and higher-order gradients misbehave. Tensors go through save_for_backward; everything else may ride as attributes.
Returning the wrong number of gradients. forward(ctx, a, b, flag) obliges backward to return three things — gradients for a and b, and None for flag. Miscounting raises a mismatch error naming your function.
Deriving the backward wrongly and never noticing. The forward runs, training runs, results quietly underperform. Wrong-but-plausible backward rules are the most expensive silent bug in custom code — which is why the very next lesson exists: checking your gradients with gradcheck. Write the Function, then check it. Always.
Reaching for a Function when detach already says it. Many "custom gradient" wishes are one stop-gradient away — x + (q − x).detach() is straight-through rounding, no class needed. A custom Function is for rules detach cannot express, or when clarity demands a named operation.
Breaking the chain with fresh tensors. Backward must compute with grad_out, not ignore it. Returning 3 * x ** 2 alone (dropping grad_out) silently discards everything downstream — the chain rule has a chain for a reason.
Try it yourself
Write Clamp01 — forward clamps to [0, 1]; backward passes gradient only where the input was inside the range (save a boolean mask on the notepad). Then feed it inputs straddling the range, backward a sum, and confirm blocked positions read zero. Keep it — the next lesson will test it properly.
What to learn next
- Checking your gradients with gradcheck — the test your new Function must pass.
- Jacobians, vjp and jvp with torch.func — the transform world your Function should stay compatible with.
- detach, no_grad and inference_mode — the lighter tool that covers many "custom gradient" wishes.
Researcher — Mathematics and papers.
What you are really defining: a VJP
A Function supplies the primitive's vector–Jacobian product: given upstream adjoint ȳ = grad_out, backward must return x̄ = ȳᵀ (∂f/∂x) per input — the same contract every built-in derivative in derivatives.yaml satisfies (see non-scalar backward for the vᵀJ framing). Autograd's chaining assumes nothing about your implementation except this contract; a wrong VJP corrupts every gradient upstream of the node with no runtime symptom — hence gradcheck as a non-optional step.
Support machinery worth knowing: ctx.needs_input_grad (tuple of bools — skip computing unneeded input grads); ctx.mark_non_differentiable(out) for integer-like outputs; once_differentiable decorator to refuse double backward honestly rather than silently mis-serve it. For double backward support, the backward itself must be composed of differentiable ops (or be its own Function); saved tensors are then repacked with graph tracking.
Since torch 2.0, the recommended signature splits construction: forward(*args) without ctx, plus setup_context(ctx, inputs, output) — required for compatibility with torch.func transforms (vmap/jacrev over your Function); the combined forward(ctx, ...) form above remains fully supported in eager code, and is what you will see in most existing codebases. Verified against 2.5.1.
The legitimate use-cases, taxonomised
- Foreign compute: custom CUDA/Triton kernels, calls into SciPy or C++ — anything opaque to the tape. The Function is the adapter making foreign code a first-class differentiable citizen.
- Surrogate gradients: straight-through estimators (Bengio et al., 2013, arXiv:1308.3432; analysed in Yin et al., 2019, Understanding Straight-Through Estimator), binary/quantised networks (Courbariaux et al., 2016, BinaryNet; quantisation-aware training per Jacob et al., 2018), spiking networks' surrogate spike derivatives (Neftci et al., 2019).
- Numerically superior hand derivatives: fusing log∘softmax or attention (FlashAttention's backward — Dao et al., 2022 — is a hand-scheduled recomputation exactly of this kind), stable implementations where naive chained derivatives cancel catastrophically.
- Memory control: recomputation-based backwards;
torch.utils.checkpointis implemented over this machinery (Chen et al., 2016, arXiv:1604.06174). - Implicit differentiation: differentiating through fixed points and optimisation solutions via the implicit function theorem rather than unrolling — Deep Equilibrium Models (Bai et al., 2019) and OptNet-style differentiable optimisation (Amos and Kolter, 2017) write backward from the IFT: at a fixed point z* = f(z*, x), one solves a linear system in the backward instead of storing the iteration.
Cost framing
A Function's forward/backward costs are whatever you write — the design freedom includes trading recompute for memory (save nothing, recompute in backward) or memory for speed (save aggressively). The tape overhead per node is constant; the interesting budget is your saved-tensor footprint, which participates in activation memory like any built-in (see graph mechanics).
References
- PyTorch docs, Extending torch.autograd — the normative guide, including setup_context and vmap rules.
- Bengio, Léonard, Courville (2013) — the straight-through estimator.
- Bai, Kolter, Koltun (2019), Deep Equilibrium Models — implicit-differentiation backward at scale.
- Dao et al. (2022), FlashAttention — a hand-written backward as a headline result.
What to learn next
- Checking your gradients with gradcheck — the test your new Function must pass.
- Jacobians, vjp and jvp with torch.func — the transform world your Function should stay compatible with.
- detach, no_grad and inference_mode — the lighter tool that covers many "custom gradient" wishes.