Lesson 3 of 4 · 35 min

Mixed precision changes arithmetic, so test the update

Explain loss scaling and the correct order for gradient clipping.

Mixed precision uses lower-precision arithmetic where it is suitable while retaining higher precision where needed. It can improve throughput or memory use, but it changes numerical behavior. The correct question is whether the resulting training procedure preserves acceptable stability and quality under the measured configuration.
Small gradients can underflow in some low-precision formats. Loss scaling multiplies the loss before backpropagation so gradients are scaled upward, then removes that scale before the optimizer consumes them. Scaling does not change the intended mathematical gradient when the scaling and unscaling are handled correctly. It does not repair every instability, such as invalid inputs or an excessively large learning rate.
Gradient clipping must operate on unscaled gradients if the clipping threshold is defined in the original gradient units. Clipping scaled gradients at the ordinary threshold shrinks updates too aggressively. A framework's scaler can also detect non-finite gradients and skip an update. Record this behavior when interpreting scheduler progress and training throughput.
Precision choices should be validated with a small reference and a realistic training comparison. Exact equality to float32 may be unreasonable after different operation ordering, but large or systematic discrepancies need investigation. Compare loss trends, gradient norms, skipped updates, and final evaluation under a fixed protocol. A faster step that requires more steps to reach the target may not reduce time to useful quality.

Worked example

A teaching gradient is g = 0.2 and the loss scale is 1,024. Backpropagation produces a scaled gradient of 204.8. Unscaling recovers 0.2. If the clipping limit is 1.0, the original gradient should remain unchanged.
A buggy path clips 204.8 to 1.0 before unscaling. After division by 1,024, the optimizer receives about 0.000977. The update is roughly 205 times smaller than intended. The run may still decrease loss slowly, which makes the bug easy to mistake for a learning-rate issue. The correct order is backward with scaling, unscale, inspect or clip, optimizer step, then scaler update under the framework's documented contract.

Exercise and solution

The true gradient norm is 2, scale is 100, and clipping threshold is 1. What norm should reach the optimizer after correct unscaling and clipping? What does clip-before-unscale produce?
Correct processing gives norm 1. The buggy order clips scaled norm 200 to 1 and then divides by 100, producing 0.01. Award one point each for both results, operation order, and a test using a known gradient. State that this simplified calculation omits vector direction changes and framework details, which must follow the actual optimizer implementation.

Lab artifact: inspect a vector gradient

Norm clipping preserves a vector's direction under the usual uniform rescaling rule when the norm exceeds the threshold. Consider true gradient [3,4], whose Euclidean norm is five. With loss scale one hundred, the stored gradient is [300,400], norm five hundred. After unscaling, clipping to norm two gives [1.2,1.6]. If the scaled vector is clipped to norm two first, it becomes [1.2,1.6] before unscaling and then [0.012,0.016], norm 0.02.
StageCorrect orderClip-before-unscale
Stored scaled gradient[300,400][300,400]
First operationUnscale to [3,4]Clip to [1.2,1.6]
Second operationClip to [1.2,1.6]Unscale to [0.012,0.016]
Optimizer norm20.02
This is idealized arithmetic ignoring numerical epsilon and rounding. It shows a factor-of-one-hundred error from operation order. Real clipping utilities may include an epsilon for stability; test the implementation within a justified tolerance. The exercise's scalar case also uses idealized arithmetic, and clipping here does not change direction.

A second failure case: accumulation with inconsistent scales

Suppose two microbatches should contribute gradients g1 and g2 to one optimizer update. If their accumulated stored gradients use different scales S1 and S2, a single final division by S2 produces (S1/S2)g1 + g2 rather than g1 + g2. The weighting has changed. The framework's documented accumulation pattern keeps the scale consistent through the effective batch and unscales once before the optimizer step.
A conceptual trace makes the intended boundary visible:
code
1effective update 12:2  microbatch A: scaled backward using scale 10243  microbatch B: scaled backward using scale 10244  accumulation complete5  unscale once6  inspect finite gradients and clip in original units7  conditional optimizer step8  update scale for a later effective update
If the intended objective averages the two microbatches, its loss normalization must also be correct. Equal division by two assumes equal desired weight; unequal token counts need a separate denominator derivation. Loss scaling and statistical weighting are different operations and should not be used interchangeably.

Exercise: account for skipped updates

An experiment processes one hundred batches. Ten steps are skipped after non-finite gradient detection. Another run processes one hundred batches and completes one hundred optimizer updates. Can the comparison claim equal completed-update budgets?
No. The first completed ninety optimizer updates, although both performed one hundred batch attempts. Record attempts, completed updates, examples processed, skipped updates, and scheduler behavior. If the scheduler advances on every attempted batch while the optimizer sometimes skips, the effective schedule differs from one defined per completed update. Whether that is acceptable depends on the intended protocol; do not assume the clocks are equivalent.
Award one point for the ninety-update count, one for the budget distinction, one for the scheduler question, and two for the measurement record and bounded comparison. A high skipped-step rate is also evidence to investigate input validity, precision range, and optimization stability. It should not be hidden inside an average throughput number.

Misconceptions to correct

“Loss scaling repairs any NaN” fails when the source is invalid inputs, an undefined operation, or forward overflow outside the scale's purpose. “Mixed precision must match float32 bit for bit to be valid” sets an unsuitable universal requirement. Use small references to detect semantic errors and a declared quality/stability criterion for the full training procedure.
A final performance recommendation needs more than milliseconds per attempted step. Suppose a float32 run reaches the target in one thousand completed updates, while a faster mixed-precision run needs fifteen hundred because of different numerical behavior. Compare total time to the same predefined target, including retries and failed runs as the protocol requires. The local speed result remains useful, but the end-to-end value must be measured separately.

Interview probe

Original practice: Why did mixed-precision training become much slower to converge after clipping was added? A strong answer checks whether clipping sees scaled gradients and inspects skipped updates. Follow up with time-to-quality versus steps-per-second. A weak answer blames lower precision without checking operation order.

Sources

docsPyTorch: automatic mixed precision examplesdocs.pytorch.orgdocsPyTorch: profiler guidedocs.pytorch.org

Checkpoint

True gradient 0.2, scale 1024, clip limit 1. Which order preserves the intended update?

AClip scaled gradient, then divide by 1024.BDivide the clipping limit by 1024 before clipping the scaled gradient.CUnscale, then clip in original units.DSkip unscaling because the optimizer will infer the scale.
Sign up free to answer and see why

Checkpoint

True vector [3,4] is uniformly norm-clipped to limit 2. Ideal output?

A[1.2,1.6]B[2,2]C[0.6,0.8]D[3,4]
Sign up free to answer and see why

Checkpoint

Two accumulated contributions use scales S1 and S2, then the sum is divided by S2. What is the first contribution's effective factor?

AAlways one.BS2/S1.CS1+S2.DS1/S2.
Sign up free to answer and see why

Checkpoint

100 batch attempts include 10 skipped optimizer steps. Which count supports a completed-update budget?

A100, because backward was attempted.B90, with attempts and scheduler behavior also recorded.C110, counting skip handling as updates.D100 if the final loss is finite.
Sign up free to answer and see why

Checkpoint

Why compare time to a fixed quality target after a mixed-precision change?

AEqual attempt counts establish equal completed updates even with scaler skips.BHigher steps per second establish the same convergence rate per update.CFaster attempted steps may have different convergence or skipped updates.DMatching the last training loss proves matching held-out quality.
Sign up free to answer and see why

Can you trace scaling, clipping and skipped updates through one effective batch? Rate confidence from 1 to 5 and explain the clock used for scheduler progress.

Not yetGetting thereConfident

Wrap-up

  • Validate numerical operation order. Measure time to a quality target as well as step throughput.

Sources

Free to read · better with Enzo

Learn it with Enzo

Save your progress, answer the checkpoints, and let Enzo quiz you on what you just read.