skip to content

When do you write a custom torch.autograd.Function, and how?

level: seniorimportance: nice to knowfreq 38%

answer

  1. only when autograd cannot derive it
  2. two staticmethods and a ctx
  3. call apply, never forward
  4. one gradient per forward input
  5. gradcheck in double precision

basics

~20 s

Subclass torch.autograd.Function with static forward and backward methods, save what backward needs via ctx.save_for_backward, and invoke it with .apply(). You need one only when autograd cannot derive the gradient itself: a non-differentiable step, a custom CUDA or C++ kernel, or a hand-written backward that is faster or more numerically stable.

solid answer

~40 s

You subclass `torch.autograd.Function` and define two `@staticmethod`s: `forward(ctx, *inputs)` computes the output and stashes anything backward will need with `ctx.save_for_backward(...)`; `backward(ctx, *grad_outputs)` reads them from `ctx.saved_tensors` and returns one gradient per input to `forward` (or `None` for inputs that need none). You call it as `MyFn.apply(x)`, never `MyFn.forward(x)` — `apply` is what registers the node in the graph. Reach for it in four situations: you wrapped a custom CUDA/C++ kernel autograd cannot see inside; the operation is non-differentiable and you want a defined surrogate gradient, as with a straight-through estimator for quantization; the composed backward is numerically unstable and a hand-derived form is better; or memory matters and recomputing beats storing. Verify with `torch.autograd.gradcheck` on double-precision inputs — a wrong hand-written backward trains badly rather than crashing.

code

python · 15 lines
python
import torch

class Square(torch.autograd.Function):
    @staticmethod
    def forward(ctx, x):
        ctx.save_for_backward(x)
        return x * x

    @staticmethod
    def backward(ctx, grad_output):
        (x,) = ctx.saved_tensors
        return grad_output * 2 * x

a = torch.randn(4, dtype=torch.double, requires_grad=True)
print(torch.autograd.gradcheck(Square.apply, (a,)))   # True

go deeper

for a junior

Know that this exists and that you almost never need it: composing standard PyTorch operations already gives you a correct backward pass automatically.

for a middle

Describe the shape of one — two staticmethods, ctx.save_for_backward, .apply() — and name the arity rule that backward returns one gradient per forward input.

for a senior

Justify when it earns its keep (custom kernels, straight-through estimators, stability, memory) and insist on gradcheck in double precision, since a wrong backward degrades training silently.

for a principal

Weigh the long-term cost: a custom Function is opaque to tracing, export and compile, so it becomes a boundary in every downstream pipeline. Prefer a registered custom operator when the model must ship, and require a gradient test in CI for any hand-written backward.

## The default is: don't Autograd already differentiates any composition of PyTorch operations, including control flow, indexing and broadcasting. Writing your own `Function` means taking responsibility for correctness that the framework was giving you for free. Assume you don't need one until one of these is true: 1. **You wrapped an external kernel.** A custom CUDA extension, a C++ op or a third-party library call is opaque to autograd, so somebody has to supply the derivative. 2. **The forward is non-differentiable and you want a useful surrogate.** Rounding, sign, argmax, hard thresholding, and quantization all have zero or undefined gradient almost everywhere. The straight-through estimator — forward quantizes, backward passes the gradient through unchanged — exists precisely because you can decouple the two directions in a `Function`. 3. **The composed backward is unstable or slow.** A closed-form derivative you derived by hand can avoid catastrophic cancellation or fuse several kernels into one. 4. **You are trading memory for compute.** A `Function` that recomputes rather than saves is the essence of activation checkpointing; `torch.utils.checkpoint.checkpoint` is the packaged version of that idea, and is what you should use before writing your own. ## The skeleton ``` class Square(torch.autograd.Function): @staticmethod def forward(ctx, x): ctx.save_for_backward(x) return x * x @staticmethod def backward(ctx, grad_output): (x,) = ctx.saved_tensors return grad_output * 2 * x y = Square.apply(x) ``` The rules that matter: - **Both methods are `@staticmethod`.** There is no `self`; the per-call state lives on `ctx`. - **Call `.apply()`.** Calling `forward` directly runs the numerics but creates no graph node, so the operation silently vanishes from the backward pass. - **`backward` returns one value per `forward` input**, in order. Non-tensor inputs (a flag, an int) get `None`. Returning the wrong arity raises at backward time, not at definition time. - **`grad_output` matches the shape of the forward output.** If forward broadcast or reduced, backward must un-broadcast or expand to match the input shape — a very common source of shape errors. - **`ctx.needs_input_grad`** is a tuple of booleans telling you which inputs actually require gradients; skipping the unneeded computations is a cheap optimisation. ## save_for_backward, and why not just ctx.x = x Stashing a tensor as a plain attribute works numerically but bypasses autograd's bookkeeping. `ctx.save_for_backward` registers the tensor with the version counter, so if it is modified in place before backward runs you get the explicit `RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation` instead of a silently wrong gradient. It also participates in the graph's lifetime management, letting the saved tensor be released when the graph is freed. Use `save_for_backward` for tensors that are inputs or outputs of the forward; plain attributes are fine for small non-tensor metadata such as a shape tuple or a scalar hyperparameter. ## Verifying it A hand-written backward that is subtly wrong does not crash — the model just learns worse, and you will blame the learning rate. So test it: ``` a = torch.randn(4, dtype=torch.double, requires_grad=True) torch.autograd.gradcheck(Square.apply, (a,)) ``` `gradcheck` compares your analytic gradient against a numerical finite-difference estimate and raises on mismatch. Use `dtype=torch.double`: in float32 the finite differences are dominated by rounding noise and the check produces false failures. `torch.autograd.gradgradcheck` does the same for second derivatives, which matters if your `Function` must support `create_graph=True`. For an intentionally non-exact gradient like a straight-through estimator, `gradcheck` will fail by design — that is the point of the estimator — so test those against a hand-computed expectation instead. ## Extras worth knowing - **`setup_context`.** Instead of putting `ctx` in `forward`'s signature, you may define `forward(*inputs)` plus a separate `setup_context(ctx, inputs, output)` static method. That separation is required if you want the `Function` to work with `torch.func` transforms such as vmap. The `ctx`-in-forward spelling remains supported for ordinary use. - **`@torch.autograd.function.once_differentiable`** decorates a `backward` that is not itself differentiable, giving a clear error rather than a confusing one if someone asks for second-order gradients. - **`mark_non_differentiable`** tells autograd that a particular output carries no gradient — useful when a Function returns both a value and an index tensor. - **Export and compile.** A custom `Function` is a black box to tracing and graph capture. If the model must go through `torch.export` or `torch.compile`, prefer registering a proper custom operator so the exporter knows about it, rather than assuming an autograd `Function` will survive the trip.

  • Why must you call MyFn.apply(x) rather than MyFn.forward(x)?
    apply() is what constructs the autograd node and wires it into the graph, binding your backward to the output's grad_fn. Calling forward directly just runs the numerics as plain tensor ops with no ctx and no graph node, so the operation is invisible to the backward pass. Nothing raises; gradients simply stop flowing through it, which looks like a vanishing-gradient problem rather than a wiring mistake.
  • What is a straight-through estimator and why does it need a custom Function?
    Quantization or rounding has a derivative that is zero almost everywhere, so gradients die at that op. A straight-through estimator keeps the discretizing forward but defines backward to pass the incoming gradient through unchanged, often clipped to the valid input range. Because forward and backward deliberately disagree, you cannot express it as a composition of differentiable ops — a custom Function is exactly the mechanism for decoupling the two directions.
  • Why does gradcheck insist on double-precision inputs?
    It compares your analytic gradient to a finite-difference approximation, which subtracts two nearly equal numbers. In float32 the rounding error in that subtraction is comparable to the quantity being measured, so correct implementations fail the tolerance check. Double precision pushes the noise floor far below the signal. Cast inputs with dtype=torch.double for the test even if the production op runs in float32 or bfloat16.
  • Your custom Function raises about the number of gradients returned. What happened?
    backward must return exactly one value per positional input to forward, in the same order, with None for inputs that take no gradient — non-tensor arguments such as flags, dimensions or scalars included. Forgetting the None placeholders for those trailing arguments is the usual cause. The mismatch surfaces at backward time rather than at class definition, so it appears mid-training rather than at import.

saying these in an interview costs you the question

  • Writing a custom Function when composing existing ops would differentiate fine
  • Calling forward() directly instead of apply() and losing the graph node
  • Storing tensors as ctx attributes instead of ctx.save_for_backward
  • Shipping a hand-written backward without running gradcheck
  • Returning one gradient when forward took several arguments

context