Skip to content

Mixed precision, loss scaling, gradient accumulation, and activation checkpointing

ModuleL11.1 · build · Python · Pass 9 · 5 to 6 h
You buildpython/tinyllm/train/precision.py: cast, autocast, autocast_bf16, autocast_dtype, DynamicLossScaler, grad_accumulate, train_step_mixed · python/tinyllm/train/recompute.py: checkpoint, checkpoint_sequential
Contractcourse/contracts/py/tinyllm/train/precision.pyi · course/contracts/py/tinyllm/train/recompute.pyi
Testscourse/tests/L11.1/ (what they check: section 4) · your own tests in python/tests/l11-1-precision/, rung R5, graded by mutation (threshold 0.80, every pitfall mutant required)
NeedsL0.1 Tensor, from_op, grad mode · L0.3 mse (the tests’ loss) · L0.4 Module and layers · M09.1 round_to_bf16, round_to_fp16 · M10.2 SGD and M10.3 AdamW (the tests’ optimizers) · M10.4 clip_grad_norm_ · M08.4 checkpoint_schedule · reading: M05.1 the memory plan, L0.5 train_step
Used bythe capstone trainer: C1 trains with --bf16 --micro-batch 8 --accum 4 --checkpoint-activations (joins the registry with the capstone)
MilestoneMS-L11 (a Llama trains with bf16, accumulation, and checkpointing to within 2% of float32 with a lower peak memory)
Optional depthMicikevicius et al., Mixed Precision Training (2018); Kalamkar et al., A Study of BFLOAT16 for Deep Learning Training (2019); Chen et al., Training Deep Nets with Sublinear Memory Cost (2016); Korthikanti et al., Reducing Activation Recomputation in Large Transformer Models (2022)
  • Mixed precision rounds what a matmul reads and writes, and nothing else: the weights stay float32 masters, so an update smaller than bf16’s spacing still lands (test_hand_example_bf16_matmul, test_autocast_grads_are_bf16_and_masters_stay_fp32, test_train_step_mixed_bf16_tracks_fp32).
  • fp16 gradients below 2−242^{-24} flush to zero. Multiplying the loss by SS multiplies every gradient by SS (backprop is linear), which lifts them into range; dividing by SS before the update undoes it, and a step whose gradients overflowed is skipped and SS halved (test_fp16_underflow_is_rescued_by_loss_scaling, test_scaler_skips_inf_steps).
  • Weighting micro-batch ii‘s mean loss by ni/Nn_i/N makes kk micro-batches give exactly the gradient of one batch of NN rows (test_hand_example_accumulation, test_accumulation_equals_one_big_batch).
  • A checkpointed segment reruns the same operations on the same values during backward, so its gradients are equal bit for bit, provided the rerun replays the random draws and its parameters are parents of the node (test_checkpoint_grads_bitwise_equal, test_checkpoint_replays_dropout_rng).
Terminal window
ol start L11.1 # stubs precision.py and recompute.py into your repo
ol tests L11.1 # read the test catalog first
ol check L11.1 # course tests, then your tests graded by mutation
ol mutate L11.1 # the full mutation grade of your tests
ol check L11.1 --ref-deps # only if your M08.4, M09.1, M10.4, or L0.x is not passing
ol diff L11.1 # after passing: your code against the reference

The capstone (C1) trains a 10.4M-parameter Llama with a 512-token context. Your L0.5 loop runs one batch per step, all in float32, and keeps every layer’s activations until backward() returns. Three things break at that size. The batch the recipe calls for (32 sequences) does not fit: M05.1’s memory_plan puts the float32 activations of one 512-token sequence through the 8 layers at 161 MiB, so a batch of 32 needs 5.0 GiB of activations next to 159 MiB of weights, gradients, and AdamW moments. Real hardware trains in bf16 or fp16 for speed and memory, and you need to know that the run you design here still converges when every matmul reads 8-bit mantissas, before you trust a recipe that depends on it. And when a gradient underflows or overflows in fp16, a float32 loop never sees it. This module adds the four tools every large training run uses, each behind one call: autocast (emulated low precision), DynamicLossScaler (fp16 range), grad_accumulate (a big batch from small ones), and checkpoint_sequential (activations kept on M08.4’s schedule). The emulation runs on numpy and is not faster; what it gives you is the exact numerics, so the convergence and memory claims of MS-L11 are measured, not assumed.

SymbolMeaningType / shape
rd16(x)\mathrm{rd}_{16}(x)xx rounded to nearest, ties to even, into bf16 or fp16 (M09.1)elementwise
ppthe precision of a format, in significant bits: 24 (fp32), 11 (fp16), 8 (bf16)int
u=2−pu = 2^{-p}the unit roundoff: the largest relative error of one roundingfloat
wwa parameter (the float32 master copy)float32[...]
η\etathe learning ratefloat
LL, g=∂L/∂wg = \partial L/\partial wthe loss and a gradientfloat, like ww
SSthe loss scale (loss_scale)float, a power of 2
NN, nin_i, kkrows in the full batch, rows in micro-batch ii, number of micro-batchesint
LiL_ithe mean loss over micro-batch ii‘s rowsfloat
sjs_j, PPsegment sizes and peak memory of a checkpoint schedule (M08.4)int

Two 16-bit formats. Both keep a sign, an exponent, and a mantissa in 16 bits. bf16 keeps float32’s 8 exponent bits and 7 mantissa bits: the same range (up to about 3.4×10383.4 \times 10^{38}) with p=8p = 8, so u=2−8≈0.4%u = 2^{-8} \approx 0.4\%. fp16 has 5 exponent bits and 10 mantissa bits: p=11p = 11, u≈0.05%u \approx 0.05\%, but its largest finite value is 65504, its smallest normal 2−14≈6.1×10−52^{-14} \approx 6.1 \times 10^{-5}, and its smallest subnormal 2−24≈6×10−82^{-24} \approx 6 \times 10^{-8}. bf16 trades precision for range; fp16 trades range for precision. numpy has neither, so M09.1’s round_to_bf16 and round_to_fp16 emulate them: a float array whose entries are all values of the format.

What a bf16 matmul does. A tensor core reads two bf16 matrices, multiplies, accumulates the sums in float32, and writes a bf16 result. So Y=ABY = AB becomes

Y=rd16(rd16(A) rd16(B)),the product computed in float32.Y = \mathrm{rd}_{16}\big(\mathrm{rd}_{16}(A)\, \mathrm{rd}_{16}(B)\big), \qquad \text{the product computed in float32}.

Each operand entry has relative error up to uu, and the output one more rounding. Errors of different entries are independent-looking, so a dot product of length KK drifts by about K u\sqrt K\, u relative, not KuK u; for training this is noise of a few tenths of a percent, smaller than minibatch noise.

The cast is an op with a gradient. Write the emulated storage as an op cast: forward y=rd16(x)y = \mathrm{rd}_{16}(x). Its true derivative is zero almost everywhere (a step function), which would stop all learning; what hardware does is pass the gradient through the conversion, and a gradient that arrives in bf16 is a bf16 value. So the backward of a cast is the cast of the gradient: xˉ=rd16(yˉ)\bar x = \mathrm{rd}_{16}(\bar y). Built from from_op (L0.1), the autocast matmul is cast(cast(a) @ cast(b)), and its backward computes Aˉ=rd16(rd16(Yˉ) rd16(B)⊤)\bar A = \mathrm{rd}_{16}(\mathrm{rd}_{16}(\bar Y)\, \mathrm{rd}_{16}(B)^\top): a bf16 backward matmul, exactly what a GPU runs. Because every matmul in your models goes through Tensor.__matmul__ (F.matmul, Linear, attention), autocast swaps in the rounding version for the length of a with block and puts the original back on the way out, also when the block raises. That is what torch.autocast does with its dispatch keys.

Why the weights stay float32. The update w←w−ηgw \leftarrow w - \eta g is where precision matters most. The spacing of bf16 numbers near 1 is 2−7≈0.00782^{-7} \approx 0.0078: a weight of 1 updated by ηg=10−3\eta g = 10^{-3} rounds straight back to 1, and the update is lost. Training in pure bf16 stalls once updates become small. So the parameters, the optimizer state (AdamW’s moments), and the update stay float32: the “master weights”. Only the matmuls, which dominate the time and the activation memory, run on rounded operands. M05.1’s memory plan counts this: bf16 AdamW is 2 (weights) + 2 (grads) + 4 (master) + 8 (moments) = 16 bytes per parameter, the same as float32, and the saving is all in activations.

Loss scaling. In fp16 the problem is range, not precision. Gradients of a mean loss over many tokens are small: a cross-entropy gradient on the logits is (softmax−onehot)/N(\mathrm{softmax} - \mathrm{onehot})/N, and backpropagated through a few layers many entries fall below 2−242^{-24} and flush to zero. Backprop is linear in the upstream gradient: every VJP is linear in yˉ\bar y, so starting backward from S⋅LS \cdot L instead of LL multiplies every gradient by exactly SS. Choose SS so the small gradients land in fp16’s normal range, and divide the float32 master gradients by SS before the update. With SS a power of two, scaling and unscaling are exact (they only change the exponent).

Dynamic scaling. Too large an SS overflows the large gradients to ∞\infty (and ∞−∞\infty - \infty to NaN). The fix is to detect, not to predict: after backward, if any gradient is non-finite, skip the whole step (one poisoned entry would corrupt every weight it touches through the optimizer), drop the gradients, and multiply SS by the backoff 12\tfrac12. After interval consecutive good steps, multiply SS by the growth factor 2 to probe upward again. The count of good steps restarts after every overflow and every growth. A skipped step costs one batch; a run without scaling silently loses its small gradients. Clipping (M10.4) comes after unscaling, because the clip bound is a norm in true gradient units. bf16 has float32’s range and needs no scaler.

Gradient accumulation. The loss of a batch of NN rows is the mean of per-row losses ℓr\ell_r. Split the rows into micro-batches B1,…,BkB_1, \dots, B_k with nin_i rows each, and let LiL_i be micro-batch ii‘s own mean. Then

L=1N∑rℓr=∑i=1kniN⋅1ni∑r∈Biℓr=∑i=1kniNLi,L = \frac1N \sum_{r} \ell_r = \sum_{i=1}^k \frac{n_i}{N} \cdot \frac{1}{n_i}\sum_{r \in B_i} \ell_r = \sum_{i=1}^k \frac{n_i}{N} L_i,

and since the gradient is linear, ∇L=∑iniN∇Li\nabla L = \sum_i \tfrac{n_i}{N} \nabla L_i. Backpropagating niNLi\tfrac{n_i}{N} L_i for each micro-batch, without zeroing in between (the Tensor adds into .grad, L0.1), accumulates exactly ∇L\nabla L. Only one micro-batch’s activations exist at a time. Weighting by 1k\tfrac1k is the same only when all micro-batches have equal size; the last one of an epoch usually does not. The results agree with one big batch to float32 rounding, not bit for bit, because the sums are added in a different order.

Activation checkpointing. M08.4 counted the trade: keep only some layers’ inputs and recompute the rest during backward. The mechanism is one graph node. checkpoint(fn, x) runs fn under no_grad, so none of its intermediate values are recorded, and returns a node (from_op) whose parents are xx and the parameters fn reads. When backward reaches the node, its VJP reruns fn with grad mode on, from fresh leaf copies of its inputs, backpropagates the upstream gradient through that rerun (which adds the parameter gradients into their .grad directly), and returns the inputs’ gradients. checkpoint_sequential applies M08.4’s checkpoint_schedule: every segment but the last through checkpoint, the last one normally.

Three details make the gradients equal bit for bit. The rerun performs the same floating-point operations on the same values, so it produces the same activations. A dropout layer draws a new mask from its generator on every call, so the node records the generator’s state before the forward pass, sets it back for the rerun, and afterwards restores the state the stream had reached, as if nothing had been rerun. And the segment’s parameters must be parents of the node: the first segment’s input is data, which does not require grad, and from_op records a node only when some parent requires grad, so without the parameters the first layers would silently get no gradient (the classic bug of reentrant checkpointing in PyTorch).

One difference from M08.4’s model: your L0.1 backward() keeps the whole outer graph alive until it returns, where PyTorch frees each saved tensor as soon as backward has used it. So the activation memory held between forward and backward falls as M08.4 predicts, while the peak during backward falls less, because the last segment stays alive while earlier segments are recomputed (test_checkpoint_lowers_peak_memory measures both).

A bf16 matmul. A=[1,1]A = [1, 1], B=[1,2−8]⊤B = [1, 2^{-8}]^\top. Both are bf16 values already. The float32 product is 1+2−8=1.003906251 + 2^{-8} = 1.00390625. Its bf16 neighbours are 1 and 1+2−71 + 2^{-7}, with 1+2−81 + 2^{-8} exactly halfway; ties go to the even mantissa, so autocast returns 1.01.0. An input of 1.011.01 is stored as 1.0078125=1+2−71.0078125 = 1 + 2^{-7}, the nearest bf16 value. For L=∑YL = \sum Y, Yˉ=1\bar Y = 1 and Aˉ=rd(Yˉ) rd(B)⊤=[1,2−8]\bar A = \mathrm{rd}(\bar Y)\,\mathrm{rd}(B)^\top = [1, 2^{-8}], Bˉ=rd(A)⊤rd(Yˉ)=[1,1]⊤\bar B = \mathrm{rd}(A)^\top \mathrm{rd}(\bar Y) = [1, 1]^\top. The same product outside autocast is the float32 1.003906251.00390625.

Accumulation. One weight w=1w = 1, inputs x=[1,2,3,4]x = [1, 2, 3, 4], targets 2x2x, loss the mean of (wx−2x)2(wx - 2x)^2. With w=1w = 1 the residuals are −x-x:

rowsLiL_i∂Li/∂w=mean(2(wx−2x)x)\partial L_i/\partial w = \mathrm{mean}(2(wx - 2x)x)weight
one batch1, 2, 3, 4(1+4+9+16)/4=7.5(1 + 4 + 9 + 16)/4 = 7.52(−1−4−9−16)/4=−152(-1 - 4 - 9 - 16)/4 = -151
micro-batch 111−2-21/41/4
micro-batch 22, 3, 429/329/32(−4−9−16)/3=−58/32(-4 - 9 - 16)/3 = -58/33/43/4

Weighted: 14(−2)+34(−583)=−0.5−14.5=−15\tfrac14(-2) + \tfrac34(-\tfrac{58}{3}) = -0.5 - 14.5 = -15, and the loss 14⋅1+34⋅293=7.5\tfrac14 \cdot 1 + \tfrac34 \cdot \tfrac{29}{3} = 7.5. Weighting each by 12\tfrac12 gives −10.67-10.67; not weighting at all gives −21.3-21.3.

Loss scaling. S=8S = 8, growth 2, backoff 12\tfrac12, interval 2, SGD with η=0.5\eta = 0.5 on one parameter starting at 0. The scaled gradients arrive as below:

stepscaled ggfinite?unscaled g/Sg/Sww afterSS aftergood steps
18yes1−0.5-0.581
2∞\inftyno: skip−0.5-0.540
34yes1−1.0-1.041
48yes2−2.0-2.08 (grown)0
516yes2−3.0-3.081

If the count did not restart at the overflow, step 3 would already be the second good step and grow SS one step early.

Checkpointing six layers. With a budget of 4 units, checkpoint_schedule(6, 4) == [0, 3] (M08.4, section 3). The forward pass runs layers 0 to 2 under no_grad and keeps only their input, then layers 3 to 5 normally. Backward runs layers 5, 4, 3, reaches the checkpoint node, reruns layers 0 to 2 forward, and backpropagates through them. Each of layers 0 to 2 has run forward twice, layers 3 to 5 once, and every gradient equals the plain run’s. These four examples are the first four tests in section 4.

python/tinyllm/train/precision.py
def cast(x: Tensor, dtype: Literal["bf16", "fp16"]) -> Tensor # rounds forward and backward
@contextmanager
def autocast(dtype) -> Iterator[None] # every Tensor matmul: cast(cast(a) @ cast(b))
def autocast_bf16(); def autocast_dtype() -> Optional[str]
class DynamicLossScaler:
def __init__(self, init=2.0**16, growth=2.0, backoff=0.5, interval=2000)
def scale(self, loss: Tensor) -> Tensor # loss * S
def step(self, opt, params, clip=None) -> bool # check, unscale, clip, step, update S
def state_dict(self) -> dict; def load_state_dict(self, sd) -> None
def grad_accumulate(model, micro_batches, loss_fn, scaler=None) -> float # weights n_i / N
def train_step_mixed(model, micro_batches, loss_fn, opt, precision="fp32", scaler=None, clip=None) -> dict
# python/tinyllm/train/recompute.py
def checkpoint(fn, *args, params=(), rngs=()) -> Tensor
def checkpoint_sequential(layers, x, mem_budget_layers, rngs=()) -> Tensor # M08.4's schedule

train_step_mixed is L0.5’s train_step for the capstone: zero the gradients, accumulate the micro-batches under autocast(precision) (none for "fp32"), then either the scaler’s step or a finite-loss check, clipping, and opt.step(). It returns {"loss", "skipped", "scale"} plus "grad_norm" when it clipped. autocast patches Tensor.__matmul__ and __rmatmul__ process-wide while the outermost block is open; blocks nest and the innermost dtype applies.

TestKINDChecksWhy it matters downstream
test_hand_example_bf16_matmulunitsection 3: 1+2−8→11 + 2^{-8} \to 1, 1.01→1.00781251.01 \to 1.0078125, bf16 gradientsyou and the test agree on the emulation
test_hand_example_accumulationunitsection 3: loss 7.5, gradient −15-15 from micro-batches of 1 and 3the capstone’s --accum 4
test_hand_example_loss_scalerunitsection 3’s five-step trace of SS, the count, and wwfp16 runs
test_hand_example_checkpoint_scheduleunitsection 3: forward counts [2,2,2,1,1,1][2, 2, 2, 1, 1, 1], equal gradients--checkpoint-activations
test_cast_rounds_both_directionsunitboth formats, forward and gradient, dtype keptthe op every autocast matmul is built from
test_autocast_matmul_matches_rounded_referencepropertyrandom products, reflected operands, Linearattention and MLP matmuls
test_autocast_grads_are_bf16_and_masters_stay_fp32propertyweight gradients are bf16 values; a 10−310^{-3} update survives in the masterwhy mixed precision converges
test_autocast_restores_matmulboundaryan exception, nesting, an unknown dtypeno bf16 leaking into the next eval
test_fp16_underflow_is_rescued_by_loss_scalingdifferentialtiny gradients vanish in fp16 and match float32 once scaledfp16 hardware
test_scaler_skips_inf_stepspropertyan overflow changes nothing, drops gradients, halves SS until steps succeeda long run survives a spike
test_scaler_unscales_before_clippingdifferentialclip to 1 of a gradient of norm 5 scaled by 1024the clip bound means what it says
test_scaler_state_dict_roundtripunitresume keeps SS and the count; bad arguments raiseTrainRun resumes from checkpoints
test_accumulation_equals_one_big_batchdifferentialmicro-batches of 5, 11, 16 rows equal one of 32 to 10−610^{-6}MS-L11’s loss within 2%
test_train_step_mixed_bf16_tracks_fp32differential60 AdamW steps in bf16 end within 2% of float32MS-L11 in miniature
test_train_step_mixed_rejects_nan_lossboundarya NaN loss raises before any updateL0.5’s guarantee kept
test_checkpoint_grads_bitwise_equaldifferential7 layers, three budgets, float64 and float32, bit for bitcheckpointing never changes training
test_checkpoint_replays_dropout_rngdifferentialsame masks, and the stream ends where the plain run’s doesdropout in the capstone
test_checkpoint_params_get_grads_from_data_inputunitthe first segment’s parameters learn from a data inputthe embedding and first block
test_checkpoint_function_and_no_gradboundarya tensor passed twice, a number argument, no_gradcheckpoint used directly
test_checkpoint_lowers_peak_memorypropertytracemalloc: forward under half, overall peak lowerthe reason to checkpoint
PitfallSymptomCaught by
1. not rounding the matmul outputbf16 runs look more accurate than hardware will betest_hand_example_bf16_matmul (mutant s01)
2. passing the gradient through the cast unroundedfloat32 gradients from a “bf16” backwardtest_cast_rounds_both_directions (mutant s02)
3. rounding only one operandhalf the error budget, unnoticedtest_autocast_matmul_matches_rounded_reference (mutant s03)
4. a cast that changes float64 to float32dtype mixing downstreamtest_cast_rounds_both_directions (mutant s04)
5. no try/finally around the patched matmulan exception leaves every later matmul in bf16test_autocast_restores_matmul (mutant s05)
6. leaving an inner block ends the outer onethe rest of the outer block runs in float32test_autocast_restores_matmul (mutant s06)
7. weighting micro-batches 1/k1/kwrong gradient whenever the sizes differtest_hand_example_accumulation (mutant s07)
8. summing micro-batch lossesgradients kk times too largetest_accumulation_equals_one_big_batch (mutant s08)
9. zeroing gradients between micro-batchesonly the last micro-batch trainstest_hand_example_accumulation (mutant s09)
10. reporting the unweighted mean lossthe logged loss disagrees with the gradient’stest_hand_example_accumulation (mutant s10)
11. never unscalingthe update is SS times too largetest_hand_example_loss_scaler (mutant s11)
12. stepping on non-finite gradientsone overflow makes every weight NaNtest_scaler_skips_inf_steps (mutant s12)
13. keeping the bad gradients after a skipa loop without zero_grad adds onto ∞\inftytest_scaler_skips_inf_steps (mutant s13)
14. not restarting the count after an overflowSS grows right back into overflowtest_hand_example_loss_scaler (mutant s14)
15. clipping before unscalingthe clip acts on S⋅gS \cdot g: tiny updatestest_scaler_unscales_before_clipping (mutant s15)
16. losing the count on resumethe resumed run grows SS latetest_scaler_state_dict_roundtrip (mutant s16)
17. stepping on a NaN loss without a scalerthe run is poisoned silentlytest_train_step_mixed_rejects_nan_loss (mutant s17)
18. checkpointing the last segment tooevery layer runs twice for nothingtest_hand_example_checkpoint_schedule (mutant s18)
19. parameters not parents of the nodethe first segment stops learning on data inputtest_checkpoint_params_get_grads_from_data_input (mutant s19)
20. a tensor passed twice gets its gradient twicedoubled gradients for shared inputstest_checkpoint_function_and_no_grad (mutant s20)
21. not replaying the generatorthe rerun’s dropout mask differs from the forward pass’stest_checkpoint_replays_dropout_rng (mutant s21)
22. not restoring the generator after the rerunlater draws repeat; runs with and without checkpointing divergetest_checkpoint_replays_dropout_rng (mutant s22)
23. running the forward pass with the graph oncorrect gradients, no memory savedtest_checkpoint_lowers_peak_memory (mutant s23)
24. rerunning on the original tensorsthe rerun’s backward runs into the outer graphtest_checkpoint_grads_bitwise_equal (mutant s24)
DirectionModuleHow it uses this
BackL0.1from_op builds cast and the checkpoint node; grad mode decides what the forward pass records
BackL0.3mse is the tests’ loss
BackL0.4Module.parameters() lists a segment’s parents; the tests’ layers
BackM09.1round_to_bf16 and round_to_fp16 are the emulated formats
BackM10.2SGD in the tests
BackM10.3AdamW in the tests and the capstone
BackM10.4clip_grad_norm_ after unscaling
BackM08.4checkpoint_schedule plans the segments
ForwardC1the capstone trainer: bf16, accumulation of 4 micro-batches of 8, checkpointed blocks; MS-L11 compares it with a float32 run

If you skip this module, the capstone’s --bf16 --accum --checkpoint-activations flags have nothing to call.

Your pieceProduction equivalentWhat it addsWhere to look
autocasttorch.autocastper-op cast policies (softmax and norms kept in float32), real bf16 kernelstorch/amp/autocast_mode.py, aten/src/ATen/autocast_mode.cpp
DynamicLossScalertorch.amp.GradScalerper-device inf checks fused into one kernel, optimizer-step skippingtorch/amp/grad_scaler.py
bf16 and fp16NVIDIA Transformer Enginefp8 matmuls with per-tensor delayed scaling, the same idea one format lowertransformer_engine/pytorch/fp8.py
checkpoint_sequentialtorch.utils.checkpointnon-reentrant checkpointing with saved-tensor hooks, RNG state for every devicetorch/utils/checkpoint.py
uniform segmentsMegatron-LM selective recomputationrecompute only attention’s cheap, large softmax and dropout, keep the matmul outputsmegatron/core/tensor_parallel/random.py
grad_accumulateDeepSpeed gradient_accumulation_stepsaccumulation fused with ZeRO’s gradient partitioning (L11.3)deepspeed/runtime/engine.py