Why does a hard argmax in the forward pass hand back a zero gradient to everything upstream?
answer
- piecewise constant between jumps
- exactly zero, not merely small
- the index blocks, the value does not
- forward hard, backward pretend identity
- biased surrogate of a different function
basics
~20 sA hard argmax is flat between its jumps, so its derivative is zero almost everywhere and undefined at ties. Reverse mode multiplies by that zero, and any parameter whose only route to the loss crosses it gets an exactly-zero gradient.
solid answer
~50 sArgmax is piecewise constant: nudge the scores infinitesimally and the winning index does not move, so the local derivative is zero; at a tie it jumps, so it is undefined there. Reverse mode composes vector-Jacobian products, and a zero Jacobian annihilates everything behind it — the tell is gradients that are *exactly* zero, not small, on precisely the parameters upstream of the hard op. The usual repair is a straight-through estimator: run the hard choice forward, but in the backward pass pretend the hard op was the identity and copy the incoming gradient to its input. Vector-quantized latents do exactly this: the decoder consumes the snapped codebook vector while the encoder receives the decoder's gradient as if the snap never happened, with an auxiliary term keeping the encoder output near its chosen code. The gradient you then follow is biased: it belongs to a different function from the one you evaluated forward.
go deeper
Recall that argmax, rounding and hard thresholds are flat between their jumps, so training cannot push a gradient through them.
Explain why the derivative is zero almost everywhere and undefined at ties, and why upstream gradients come out exactly zero rather than merely small.
Diagnose it from a run: find the parameters with bit-exact zero gradient, trace forward to the hard operation, and propose a surrogate while stating plainly what its bias costs.
Decide between a biased surrogate through the discrete step and restructuring so the hard choice sits outside the gradient path entirely; own the risk that the function trained and the function differentiated are not the same.
## Why the derivative is zero, not small Argmax maps a score vector to the index of its largest entry. Perturb any score by a small amount and, unless two scores are tied, the winner is unchanged: the output is *locally constant*. A locally constant function has derivative exactly zero. At a tie the output jumps discontinuously, so the derivative does not exist there at all. Zero almost everywhere, undefined on a measure-zero set — that is the whole story, and it is shared by rounding to an integer, a hard threshold that emits 0 or 1, snapping a vector to the nearest entry of a codebook, and sampling a discrete index. Reverse mode composes vector-Jacobian products along each path. A zero Jacobian in the middle of a path multiplies everything behind it by zero, so any parameter whose *only* route to the loss passes through the hard op receives a gradient that is exactly `0.0`. That exactness is the diagnostic signal: a vanishing-gradient problem produces tiny numbers, a bad learning rate produces small updates, and a non-differentiable barrier produces bit-exact zeros. ## Index versus value — the distinction people get wrong Not every operation involving a maximum blocks gradient. Max-pooling outputs the maximum **value**, and that value is a differentiable function of the inputs away from ties: the backward pass routes the entire incoming gradient to the element that won and zero to the losers. Gradient flows; it is merely sparse. The same holds for gathering values by an index — the *gathered values* carry gradient even though the index does not. The barrier appears when the **index itself**, or something derived only from it (a one-hot vector, a hard 0/1 mask, a snapped code chosen by identity rather than value), is what continues into the loss. Then the scores' only influence on the output is through a decision that is flat in them. ## Diagnosing it in a real run The symptoms: a subnetwork whose parameters never change; a loss that plateaus immediately at whatever the untrained forward pass produces; encoder outputs identical across steps. To confirm, run one backward pass from the real loss and print per-parameter gradient norms — the affected block is exactly zero while the layers after the hard op are healthy. Then walk the graph forward from those parameters and find the first operation whose output does not change under a small perturbation of its input. A quick ablation confirms it: temporarily replace the hard choice with the underlying continuous scores; if gradient appears everywhere, you have found the barrier. ## The straight-through estimator You cannot differentiate the hard op, so you substitute one you can. The straight-through estimator keeps the forward pass exact — the hard, discrete value is what the rest of the model consumes, so training and inference agree — and replaces the backward pass with that of an identity, copying the incoming gradient straight to the pre-hard input. A vector-quantized latent model is the clean example. The encoder emits a continuous vector; the forward pass snaps it to the nearest codebook entry and hands the snapped code to the decoder. In the backward pass, the gradient the decoder produces with respect to the snapped code is copied verbatim onto the encoder's continuous output, as though the snap had not happened. Because the copied gradient is only sensible while the encoder output and its assigned code are close, this is paired with a term that pulls the encoder output toward the code it selected — without that pressure, the surrogate's premise erodes and the estimator degrades. ## What the surrogate costs Be precise about what has happened: you are performing gradient descent on a function that is *not* the one you evaluate in the forward pass. The estimator is **biased** — not high-variance-but-centred, but systematically the gradient of a different map. The true gradient of the hard op is zero, so this is not an approximation error that shrinks with more samples or a smaller step; it is a deliberate substitution, justified empirically rather than derived. The bias grows with the size of the discrepancy the surrogate ignores — the distance between the pre-hard value and the post-hard value. Two consequences follow. First, anything that keeps that gap small (an auxiliary pull toward the chosen output, a well-covered set of choices) directly improves the estimator. Second, the usual optimisation intuitions weaken: convergence guarantees do not apply, and instability tends to appear as oscillation in which discrete option gets selected rather than as an exploding loss. ## The alternative worth naming Before reaching for a surrogate, ask whether the discrete decision has to be inside the gradient path at all. Often the loss can be defined on the continuous scores, with the hard choice applied only at inference or only for reporting. That keeps the whole path exactly differentiable and gives up nothing except the ability to train against the discrete outcome directly. Choosing between an exact gradient of an approximate objective and a biased gradient of the exact objective is the real decision here.
- Max-pooling also picks a maximum. Why does it not block gradients?Because it outputs the maximum value, not the index. Away from ties that value is a differentiable function of the inputs, and the backward pass routes the whole incoming gradient to the winning element and zero to the rest. Gradient flows, just sparsely. The barrier appears only when the index itself, or a one-hot derived from it, is what continues into the loss.
- What does a straight-through estimator actually estimate, and why is it biased?It returns the gradient of a surrogate in which the hard op was replaced by the identity, not the gradient of the function evaluated forward — whose true gradient is zero almost everywhere. So it is systematically the derivative of a different map, and the bias grows with the gap between the pre-hard and post-hard values. More samples or a smaller learning rate do not remove it.
- How would you confirm from a training run that a non-differentiable op is the culprit rather than vanishing gradients?Print per-parameter gradient norms after one backward pass from the real loss. A hard op gives bit-exact zeros on everything upstream while the layers after it look healthy; vanishing gradients give small but nonzero values that decay smoothly with depth. Confirm by swapping the hard choice for the underlying continuous scores and checking that gradient reappears.
saying these in an interview costs you the question
- Says the gradient is very small rather than exactly zero
- Blames vanishing gradients or the learning rate
- Claims any operation involving a maximum blocks gradient
- Believes the straight-through estimator gives the true gradient
- Cannot say what the surrogate's bias depends on