skip to content

For a linear -> tanh -> linear chain, which tensor must each op stash for its own backward step?

level: seniorimportance: should knowfreq 44%

answer

  1. not every op saves the same thing
  2. ask what the rule reads back
  3. tanh's derivative from its own output
  4. a linear op wants its input
  5. one bit per element suffices for ReLU

basics

~20 s

Each op stashes only what its own backward step reads: a linear op keeps its input, tanh keeps its own output, ReLU needs nothing more than the sign of its input, and a plain addition keeps no tensor at all.

solid answer

~50 s

Bookkeeping is per-op, not per-layer. Take a speech frontend that maps log-mel frames through a linear op, a `tanh`, then a second linear op. The first linear must keep the frames it received, because its weight gradient is formed from that input; the `tanh` keeps its own **output** `h`, since its derivative is `1 - h^2` and needs no second evaluation of the function; the second linear keeps `h` too, as its input. Notice `h` is saved once and referenced by two records. Other ops are cheaper than they look: a ReLU only needs to know the sign of what entered it, which is a single bit per element rather than a full-precision tensor, and adding two tensors saves nothing at all because its backward step is a pass-through that needs only shapes. Auditing a model op by op is what tells you which tensors are genuinely pinned by the graph.

code

python · 22 lines
python
import math

tape = []                      # one record per op, appended in forward order
x, w1, w2 = 0.5, 1.2, -0.8

a = w1 * x                     # linear-ish: its backward reads the INPUT x
tape.append(("mul", {"x": x}))
h = math.tanh(a)               # tanh: its backward reads its OWN OUTPUT h
tape.append(("tanh", {"h": h}))
y = w2 * h                     # second linear: its backward reads its input h
tape.append(("mul", {"h": h}))

g = 1.0                        # dL/dy, then sweep in reverse tape order
_, saved = tape.pop()
g_w2 = g * saved["h"]
g = g * w2
_, saved = tape.pop()
g = g * (1.0 - saved["h"] ** 2)    # 1 - tanh^2, straight from the saved output
_, saved = tape.pop()
g_w1 = g * saved["x"]

print(round(g_w1, 6), round(g_w2, 6))   # -0.284631 0.53705

go deeper

for a junior

Know that the backward step of an op needs some value from the forward pass, and that it differs by op — a linear op wants what came in, tanh wants what came out. You are not expected to enumerate a whole chain yet.

for a middle

Walk a short chain and name the saved tensor per op, including the ones that save nothing. Be able to justify saving tanh's output rather than its input using the form of the derivative.

for a senior

Demonstrate the audit habit: which tensors are genuinely pinned, which records share one, and what breaks when a saved tensor is mutated in place. Recognising a silently wrong gradient as a bookkeeping failure is the signal here.

for a principal

Own the guidance for custom ops in a codebase: what a forward must stash, when saving the output is worth its fragility, and how in-place mutation is prevented by convention or by checking rather than left to reviewer vigilance.

## The question every op answers When an op runs under recording, it appends a node to the graph and pins whichever tensors its backward step will read. The right way to reason about this is not "training keeps activations" but, op by op: **what does this particular gradient rule need to be evaluated?** The answer differs sharply between ops, and the differences are what a senior candidate is expected to notice. ## Walking the speech frontend chain A small speech frontend takes a window of log-mel frames `x`, applies a linear op to get `a`, applies `tanh` to get `h`, then applies a second linear op to get the output `z`. **The first linear op.** Its gradient with respect to its weights is formed from the tensor that entered it, so `x` must survive. Its gradient with respect to its input is formed from the weight matrix — but weights are long-lived parameters that exist independently of the graph, so nothing extra is pinned on their account. Note what is *not* needed: the op's own output `a` plays no part in either rule and does not have to be kept for this op. **The tanh op.** Its derivative can be written as `1 - h^2`, purely in terms of its own output. So saving `h` is sufficient and is the cheaper choice: the alternative, saving the input `a`, would force the backward step to evaluate `tanh` a second time to get there. Sigmoid behaves the same way — its derivative is `s * (1 - s)` in terms of its output — and so does softmax. This is a general pattern worth naming: for several saturating nonlinearities, the *output* is the more useful thing to keep. **The second linear op.** Its input is `h`, so `h` is what it needs. **And now the interesting bit:** `h` is referenced by two records, the `tanh` node and the second linear node. It is stored once; the graph simply holds two references to the same tensor. Bookkeeping counts distinct tensors that are pinned, not the number of records that mention them. So the surviving set for this chain is `{x, h}` plus the parameters, which were never transient in the first place. `a` need not survive. ## Ops that need less than you expect **ReLU.** Its gradient passes the incoming gradient through where the input was positive and blocks it where it was not. Nothing about the magnitude matters — only the sign of the input, one bit per element. An implementation may keep the full input tensor for simplicity, but the *information* required is a mask. **Plain addition of two tensors.** The backward step hands the same gradient to both operands. No input value appears anywhere in that rule, so nothing needs saving; at most the op records the shapes so the gradient can be routed to the right operands. **Reshape and transpose.** These move values around without changing them; the backward step is the inverse rearrangement, which needs only the original shape, not the data. **Multiplying two tensors together.** Here each operand's gradient is formed from the *other* operand, so both must be kept. Multiplying by a fixed scalar, by contrast, needs neither operand — the rule is just a scaling. ## Why this matters beyond trivia First, when you write a custom op you have to decide explicitly what its forward stashes. The decision rule is exactly the one above: save the minimum from which the local gradient can be evaluated, and prefer the output when the output determines the derivative. Second, and this is the sharp practical edge: **a saved tensor must not be mutated before the sweep reads it.** If some later op overwrites a buffer in place, and that buffer is the very tensor a node stashed, the backward step reads the wrong values. The gradient is then silently wrong rather than loudly broken — no error is raised, the numbers are simply not the derivative of anything. Ops that save their *output* are especially exposed, because an in-place activation applied afterwards is precisely the sort of thing that overwrites it. Autodiff systems guard this with version counters on tensors, but the reasoning is worth being able to reconstruct without help. Third, it corrects a common half-truth. People say "the forward pass saves all activations". It saves the ones some op's rule reads. Intermediate values that no rule mentions — like `a` in the chain above — do not have to be pinned at all. ## The whiteboard version Asked about `y = W2 * relu(W1 x + b1)`, the answer is: `x` (the first linear's input), the sign pattern of `W1 x + b1` (for the ReLU), and `relu(W1 x + b1)` (the second linear's input). The bias addition saves nothing. `W1` and `W2` are parameters, alive regardless. Being able to enumerate that set, and say which op consumes each member, is the whole skill this question tests.

  • Why would an op prefer to save its output rather than its input?
    Because for several nonlinearities the derivative is expressible in the output directly: `tanh` gives `1 - h^2` and sigmoid gives `s * (1 - s)`. Saving the output means the backward step evaluates a cheap polynomial instead of re-running the function on the saved input. The cost is fragility — an in-place write over that output later in the forward pass corrupts what the node stashed.
  • What goes wrong if a tensor an op saved is overwritten in place before the backward sweep reads it?
    The backward step evaluates its rule at the wrong values, so the gradient is silently incorrect rather than raising an error. Training then drifts for reasons no stack trace explains. It bites hardest with ops that save their output, since an in-place activation applied afterwards overwrites exactly that buffer. Version counters on tensors exist to turn this into a loud failure.
  • In a linear -> tanh -> linear chain, how many distinct tensors are pinned by the graph?
    Two transient ones: the chain's input, kept by the first linear, and the `tanh` output, which is referenced twice — once by the `tanh` node and once as the second linear's input — but stored once. The pre-activation between the linear and the `tanh` need not survive at all. Weights and biases are parameters and were never transient.

Each op leaves one sticky note for its future self. Some need the ingredient they were handed, some need the dish they produced, and some need nothing but the shape of the plate.

saying these in an interview costs you the question

  • Says every intermediate tensor is always saved
  • Claims tanh must save its input to compute its derivative
  • Thinks a linear op saves its output for the backward step
  • Believes an addition of two tensors saves both operands
  • Assumes an in-place overwrite of a saved tensor raises an error

context