Which parts of a mixed-precision training step must stay in single precision, and why?
answer
- One rule generates the whole list
- Small quantity meeting a large one
- The optimizer steps on which copy?
- Anything that sums many terms
- Update is a no-op, weights freeze
basics
~20 sA mixed-precision step keeps the weights the optimizer updates, large reductions, the loss and normalization statistics in 32-bit. Those are the places where a tiny quantity is added to or accumulated with a much larger one, which 16-bit arithmetic destroys.
solid answer
~50 sFour things stay in 32-bit. The **master weights**: the optimizer holds a 32-bit copy and steps on that, because `learning_rate * gradient` is often many orders of magnitude smaller than the weight itself and a 16-bit add would simply return the weight unchanged, stalling training. **Large reductions**: a sum over thousands of terms in 16-bit loses the small contributions once the running total grows, so softmax denominators, log-sum-exp, and the accumulation inside a matrix multiply are done in 32-bit. **The loss and its gradient**, since the loss is itself a reduction and any exponentials feeding it can leave the reduced format's range. **Normalization statistics** -- the mean and variance a normalization layer computes -- for the same reduction reason. The unifying rule is: reduced precision is fine for bulk elementwise work and dense multiplies, and unsafe wherever a small number must survive contact with a large one.
go deeper
Recall that the optimizer keeps a 32-bit copy of the weights and steps on it, and that the loss is computed in 32-bit. Being able to name those two is enough at this level.
Explain the mechanism: an addition aligns exponents, so a small update against a large weight vanishes in 16 bits. Generalize that to reductions, softmax and normalization statistics.
Be ready to diagnose the failure from a curve -- a run that plateaus exactly rather than becoming unstable -- and to say why raising the learning rate seems to help and why that is the wrong fix.
Own the tradeoff that every exclusion buys safety with memory and bandwidth, and be able to justify which exclusions you would keep if the master copy's memory became the binding constraint.
## The single rule behind the whole list Every item on the "keep this in 32-bit" list is an instance of one problem: **a small quantity has to survive being combined with a much larger one**. Floating point stores a value as a signed significand scaled by an exponent, and an addition first aligns the two operands to a common exponent. If the operands differ enough in magnitude, the smaller one shifts entirely out of the significand's width and the addition returns the larger operand unchanged. A 16-bit format has far fewer significand bits than a 32-bit one, so the gap at which this happens is much narrower. Wherever a training step deliberately adds something small to something large, that gap is the failure. That single rule generates the whole list. ## Master weights and the optimizer step This is the item interviewers care most about. The optimizer keeps a **32-bit master copy** of every parameter. Each step casts it down to 16 bits for the forward and backward passes, but the update `w <- w - lr * g` is applied to the 32-bit master, and the next step's 16-bit weights are cast fresh from it. Why: late in training a weight might be of order 1 while a single update is of order 1e-7. In 32-bit that update lands. In 16-bit the two magnitudes are far enough apart that the add is a no-op, and the weight stops moving entirely — not noisily, but *exactly*. The model appears to plateau for no visible reason. Keeping the master copy costs memory but is non-negotiable, and the optimizer's own state (momentum buffers and second-moment estimates, for optimizers that keep them) is kept in 32-bit for the same reason: those accumulators are long-running sums that must not be truncated. A useful way to say it in an interview: **the forward and backward passes are where you spend precision; the update is where you must not**. ## Reductions A reduction is any operation that collapses many values into one by summing. The danger is the same alignment problem, applied repeatedly: once the running total has grown, each new small term contributes nothing, and the error compounds over the length of the sum. The reductions that matter in practice: - **The accumulator inside a matrix multiply.** The matrix units are built to take reduced-precision inputs and accumulate into a 32-bit register precisely so a long dot product does not decay. This one is handled by the hardware, not by you, but knowing it is what separates a real answer from a memorized list. - **Softmax and log-sum-exp.** The denominator is a sum of exponentials. Beyond precision, the exponential itself can leave the reduced format's range on a large logit, so this is a range hazard as well as a precision one. - **Normalization statistics.** A normalization layer computes a mean and a variance over some axis — both reductions — and then divides by the square root of the variance. A variance computed badly can come out near zero or negative from cancellation, and dividing by it amplifies the damage into every downstream activation. ## The loss The loss is a reduction over the batch, usually fed by a softmax or a log. It is a single scalar computed once per step, so keeping it in 32-bit costs essentially nothing and removes a whole class of range and cancellation problems at the point where the entire backward pass begins. A wrong loss corrupts every gradient behind it, so this is the cheapest safety in the entire arrangement. In a translation model trained in half precision, for example, the usual arrangement is exactly this: the bulk projections run reduced, while the softmax, the log-sum-exp reduction behind the loss, the cross-entropy itself and the normalization statistics all stay in 32-bit. The reduced-precision fraction of the step is still the overwhelming majority of the arithmetic, because those excluded operations are cheap. ## What is safe in 16 bits Stating the complement is what shows you understand the rule rather than a list. Safe: the dense multiplies (with 32-bit accumulation), the activations flowing between layers, elementwise nonlinearities, and the gradients as they propagate — all of these are values of comparable magnitude being multiplied or passed along, not small values being folded into large ones. This is the overwhelming majority of the arithmetic and the memory traffic, which is why the technique is worth doing at all despite the exclusions. ## The failure signature If the master copy is dropped and the optimizer steps on 16-bit weights directly, the run does not crash. It trains for a while and then flattens, with a loss curve that looks like a learning-rate problem. Raising the learning rate appears to help briefly — because larger updates clear the rounding threshold again — which sends people down the wrong debugging path. Recognizing that signature is the senior half of this question. ## Interview framing Name the rule first, then the list. Candidates who only recite "keep master weights in fp32" get partial credit; candidates who explain that the update is small relative to the weight, and generalize that to reductions, the loss and normalization statistics, get the question.
- What does a run look like if the optimizer steps on 16-bit weights with no master copy?It trains normally at first and then flattens. Once updates become small relative to the weights, the 16-bit addition returns the weight unchanged and the parameter freezes exactly. The loss curve reads like a learning-rate or scheduling problem, and raising the learning rate temporarily unfreezes it, which is why people misdiagnose it. The tell is that the plateau is precision-dependent: the same recipe in full precision keeps improving.
- Why keep the optimizer's momentum and second-moment buffers in 32 bits too?They are long-running accumulators. A momentum buffer is a decaying sum over many steps and a second-moment estimate is a decaying sum of squared gradients; both fold small new terms into a larger running value on every step, which is the exact operation reduced precision handles worst. Truncating them corrupts the effective step size in a way that is very hard to observe directly.
- Is keeping softmax in 32 bits a precision problem or a range problem?Both. The denominator is a sum of many exponentials, so it has the usual reduction precision issue, and the exponential of a large logit can exceed the reduced format's maximum outright and become infinity. The 32-bit path removes both hazards at once, and it is cheap because the softmax is a small fraction of the step's arithmetic.
saying these in an interview costs you the question
- Names master weights but cannot say why
- Thinks the whole optimizer state can be 16-bit
- Believes reductions are safe because inputs are small
- Says normalization layers are pure elementwise work
- Claims a frozen weight would show up as instability