Why does quantization-aware training need a surrogate gradient for the rounding step?
answer
- the staircase has no useful slope
- zero almost everywhere kills the chain rule
- pretend the node was the identity
- pass through inside, zero outside
- gradient of a network you did not evaluate
basics
~20 sRounding is a staircase: its derivative is zero almost everywhere, so exact backpropagation would send zero gradient to every quantized weight. The surrogate pretends the rounding step is the identity and passes the incoming gradient straight back.
solid answer
~50 sThe fake-quant node contains `round()`, which is piecewise constant — slope zero between steps, undefined at them. Backpropagating it exactly multiplies every upstream gradient by zero, so nothing upstream of a quantized tensor ever learns. The standard fix is a straight-through surrogate: in the backward pass, treat the rounding as if it were the identity and pass the gradient through unchanged. The practical version is *clipped*: pass the gradient for values inside the clipping range, and zero it for values outside, because a clamped output genuinely cannot respond to that input. The cost is that the gradient you get is not the gradient of the function you evaluated — it is the gradient of an idealised unquantized network. That mismatch is why weights near a bin boundary can oscillate between two levels late in training, and why naive unclipped pass-through reports a confidently wrong gradient for saturated values.
go deeper
Remember the core fact: rounding has zero slope, so training would receive no gradient at all, and a stand-in derivative of one is substituted in the backward pass.
Explain the substitution precisely — forward computes the real rounding, backward acts as if it were the identity, and the gradient is zeroed for values pinned at the clipping boundary.
Diagnose with it. Tie noisy late-run loss, drifting out-of-range weights and stale normalization statistics back to the mismatch between the function evaluated and the function differentiated, and name the schedule fixes.
Own the risk assessment: decide how far down in bit width a surrogate-driven run is trustworthy for your model family, and what evidence justifies committing a product to a regime where the training signal is a deliberate approximation.
## The obstruction Quantization-aware training injects a rounding operation into the forward pass. That operation is a staircase function. Between any two steps it is exactly flat, so its derivative is zero; at the steps themselves it jumps, so the derivative does not exist. "Zero almost everywhere, undefined on a measure-zero set" is the worst possible shape for backpropagation: the chain rule multiplies by that derivative, and every gradient arriving from downstream is annihilated. A network with a rounding node in it, differentiated honestly, has no learning signal for anything upstream of the node — which for weight quantization means the weight itself. This is not a numerical problem you can tune away with a smaller learning rate. It is a structural property of the function. ## The surrogate The standard answer is to lie in the backward pass. In the forward pass, compute the real thing: clip, round, rescale. In the backward pass, pretend the whole node was the identity function and pass the incoming gradient back unchanged. This is straight-through estimation applied to the fake-quant node. Justification, informally: the node's output tracks its input closely — it never differs from it by more than half a step, as long as the input is in range. A function that stays within half a step of the identity is, on the scale the optimizer cares about, approximately the identity, so borrowing the identity's derivative of 1 is a defensible approximation. It is an approximation, not a derivation; there is no sense in which it is *the* gradient. ## The clipped variant, and why it matters The half-a-step argument holds only inside the clipping range. Outside it, the node's output is pinned to the boundary: increase the input by any amount and the output does not move at all. There the true derivative is genuinely zero, and it is zero for a real reason rather than an artefact of the staircase. So the widely used form is the **clipped** straight-through estimator: gradient passes through with factor 1 for values inside `[lo, hi]`, and is set to 0 for values outside. Naive unclipped pass-through instead reports, for a weight sitting far outside the range, a gradient claiming the loss will respond if you nudge that weight — when in fact the forward pass has saturated and cannot respond at all. That is not an approximation being slightly loose; it is a gradient pointing at a sensitivity that does not exist, and it lets weights drift ever further out of range while the optimizer believes it is making progress. The symmetric hazard is a range that is too tight: with clipped gradients, a value that lands outside receives no gradient and can become stuck there permanently, since nothing pulls it back in. Practitioners therefore either widen ranges, or let a small gradient leak outside the range, or learn the range itself as a parameter so it can grow to recapture stranded values. ## What the approximation costs **Gradient mismatch.** The gradient you apply is the gradient of a smooth, unquantized network; the loss you measured came from a quantized one. The two agree in direction most of the time and disagree exactly where the quantization is doing something — near bin boundaries. **Oscillation.** A weight sitting close to the midpoint between two levels rounds to one bin; the surrogate gradient, computed as though rounding never happened, may push it just past the midpoint; next step it rounds to the neighbouring bin and the gradient pushes it back. The underlying float parameter wobbles, and the *effective* quantized weight flips between two values from step to step. This shows up as noisy loss late in a low-bit run, and it corrupts BatchNorm running statistics because the statistics are collected under one bin assignment and used under another. **Bit-width sensitivity.** At 8 bits the steps are fine, the identity approximation is close, and QAT behaves almost like ordinary fine-tuning. At 4 bits and below the steps are coarse, the approximation is bad, and runs become sensitive to learning rate, initialization from a good float checkpoint, and schedule. ## The extreme case Binary weight networks (weights constrained to two values) and ternary networks (three) make the point unmissable. There the forward pass throws away nearly all information about the underlying float parameter, and the true gradient is zero everywhere. Training such a network at all is only possible because a latent full-precision weight is kept, updated using the surrogate gradient, and re-binarized on each forward pass. The surrogate is not a convenience there; it is the entire mechanism by which learning happens. ## Mitigations worth naming - Start from a trained full-precision checkpoint and fine-tune, rather than training low-bit from scratch. - Use a reduced learning rate, so the mismatch does not get amplified. - Learn the clipping threshold so the in-range region adapts instead of stranding values. - Anneal toward hard rounding, using a softer approximation of the staircase early in training and tightening it later, so the mismatch shrinks as the solution settles.
- What goes wrong if you pass the gradient through unclipped, for values outside the range?You report a gradient for a saturated value. The forward pass has pinned that value to the boundary, so the loss cannot respond to it at all, yet the optimizer is told it can. Weights then drift further and further outside the range while training appears to progress, and the effective quantized network stops improving. Zeroing the gradient outside the range removes that phantom sensitivity.
- How do binary or ternary weight networks train at all if the true gradient is zero everywhere?They keep a latent full-precision weight that the optimizer actually updates, and re-binarize it on every forward pass. The surrogate gradient is applied to the latent weight, so the accumulated small updates eventually flip a sign. Without the surrogate there is no learning signal whatsoever, which makes these the clearest case that the estimator is the mechanism, not an optimization.
- Why does a 4-bit QAT run behave much worse than an 8-bit one under the same surrogate?The identity approximation is only good when the step size is small. At 8 bits the node never moves a value by more than half a fine step, so the surrogate is close to honest. At 4 bits the steps are coarse, the gap between the evaluated function and the differentiated one is large, and the run becomes sensitive to learning rate, initialization and schedule.
- What does weight oscillation between adjacent bins do to a network with BatchNorm?It corrupts the running statistics. The mean and variance are accumulated while a weight sits in one bin, then used at evaluation when it has flipped to the neighbour, so the normalization no longer matches the weights it normalizes. Freezing BatchNorm statistics late in the QAT schedule, or damping the oscillation with a lower learning rate, is the usual remedy.
It is like steering a car whose steering wheel only clicks between fixed notches, while you plan your turns with a map that assumes smooth continuous steering. The plan is close enough to follow — until you are right at a notch boundary and keep flicking back and forth.
saying these in an interview costs you the question
- Says rounding has a small nonzero derivative
- Thinks the surrogate is the true gradient
- Passes gradient through for clamped values too
- Claims the mismatch vanishes at any bit width
- Cannot explain why weights oscillate between bins
- Believes low-bit runs need no float checkpoint to start from