Low-precision training rounds tensors that the backward pass reads again, often for several gradients; each use can read the forward's rounded value, the original, or a new random rounding. This backward-state policy looks like a memory and precision detail, settled by copy accuracy and final loss. We argue that it is part of the learning algorithm, and that neither check shows whether it is right. Copy accuracy does not decide the outcome: in three pairs of 390M runs with an emulated FP8 backward, training fails when attention's backward reuses the forward's rounded output and succeeds with a new rounding from the same distribution. Even the most accurate copy, the original itself, can be wrong by our reference: the gradient of the forward pass as it actually ran, with gradients passed through rounding unchanged. For example, a normalization output stored in low precision feeds two gradients: the gain's gradient needs the original, but the next layer's weight gradient needs the rounded value that layer multiplied. Final loss, the other check, does not rule out the error of reading the original for both: it persists in models trained with such a store, while planned loss comparisons stay within a margin fixed in advance. We therefore derive from this reference which value each use must read, or which substitute gives the same gradient on average with the forward held fixed, and check these per-use requirements on single operators, without training. In three tests using PyTorch and Transformer Engine, the requirements predicted beforehand whether reuse changes what the backward computes on average relative to an independent copy, and every prediction held. Backward-state policy is thus part of the learning algorithm: it should be specified and checked use by use, not settled by copy accuracy and final loss.
Figures & tables
Figure 1: Same rounding distribution, opposite outcomes. Training loss of 390M runs with an emulated FP8 attention backward, forked from one checkpoint; all consume a stochastic rounding of the attention output O , and within each panel they differ only in how the correction term D is computed. (a) The first of three pairs: reusing the forward’s rounding fails (final loss 7.314); a new rounding from the same distribution trains (3.256), like a control whose D reads no copy of O . (b) The runs of (a), repeated on newer versions of PyTorch and Transformer Engine.
Figure 2: Changing only the attention backward removes the failure. Training loss of 390M models trained from scratch over the 20k-step schedule, one panel per seed, shown every 250th step. With TE’s FP8 attention the loss jumps and stays high on every seed, ending 5.097, 5.120 and 5.225 nats above BF16 (the same model trained with FP8 disabled) in validation loss. Keeping its FP8 forward and replacing only the attention backward with a BF16 one tracks BF16, ending 0.0029, 0.0059 and 0.0052 nats above it. Seed 2’s BF16 run used PyTorch 2.13 and TE 2.13 instead of 2.11 and 2.10 (Appendix E ).
Figure 3: Each use has its own reference value. (a) Policies, written gain use/weight use; the marks rate each use’s gradient ( ≈ : on average over the new rounding, the forward fixed; N/N hands both uses one new rounding). U/R is the reference; ★ : matches it at both uses, at least on average. (b) Where each policy’s average update is zero in the two-parameter model of Section 3.2 : ★ marks the minimum of the expected loss (shading); vermillion dots: the gain use reads R.
Figure 4: Every named interval but one lies inside the margin fixed in advance. Held-out loss of the 390M runs trained with a 6-bit (a) or 5-bit (b) store, evaluated without it, first policy minus second: model-based 95% intervals (bars) and the two seeds (markers). Only the R/R − U/R interval at 5 bits crosses the margin.
Figure 5: Comparable final loss, different trained weights. (a) Root-mean-square gain of the normalization before each MLP at the end of the 6-bit runs, divided by U/R’s at the same layer and seed; each band spans its group’s runs on both seeds (U/R’s at 1). (b) The mean gain over layers during training for U/R, R/R and BF16 without the store (seed 0).
Figure 6: Registered predictions about reuse held. Lines: predictions written before each test; dots: measurements. (a) The use’s mean error, as a percentage of its size without checkpointing (random state not restored). (b) Mean change per gradient entry, divided by its probability: shared minus independent copy under each loss, and either copy’s error under the likelihood; the inset enlarges ε=0.1 . (c) Query gradient, independent minus shared sample, divided by the softmax scale; mean of four sample pairs.
Appendix figures & tables9 assets
Supplementary material from the paper’s appendix.
Appendix
Policy
Softmax’s backward reads
Gradient of V reads
E[Δs∣F]
Row j of E[ΔV∣F]
U/R (reference)
J(p)
pq
0∗
0∗
R/R
J(pq)
pq
δ
0∗
U/U
J(p)
p
0∗
−σj2vj⊤
N/N
J(pf)
pf
−ΣrGp
−σj2vj⊤
R/U
J(pq)
p
δ
−σj2vj⊤
Appendix
Table 1: Average errors of the reference and four policies at the two uses of an attention row’s rounded probabilities. Policy codes name the softmax’s backward first; averages are over the forward’s rounding and any new rounding, the forward’s inputs fixed; δ is given in equation 9 . Zeros marked * hold in every draw.
Statement
Registered rule
Measured
Outcome
First configuration: width 256, n=512
S1. Restoring the random state makes checkpointing change nothing
record bit-identical to the baseline’s, wherever the checkpointed part ends; masks agree exactly
bit-identical; masks agree
held
S2. Without restoration, the effect is removed when recomputation stops at the rounding
z<2 and under 10 %; baseline z>5
z=0.06 , 0.19 %; baseline z=30.59
held
S3. It is kept when recomputation includes the loss
z>5 , the baseline’s sign, over 30 %
z=33.30 , same sign, 106.12 %
held
S4. Without restoration, the dropout mask changes
agreement between 0.30 and 0.85 ; exactly 1 with restoration
0.4998; all entries with restoration
held
S5. PyTorch’s own check stays silent
no error or warning in any condition
none
held
Appendix
Table 2: Checkpointing: every registered statement and its outcome. “Removed” and “kept” refer to the use’s mean error when the random state is not restored. Percentages and ratios are relative to the baseline without checkpointing, in magnitude.
Statement
Registered rule
Measured
Outcome
First configuration: ε=0.1
LS1. Under the likelihood, reuse changes nothing on average
z<2 and under 5 % of the constructed loss’s effect
z=0.68 , 0.17 %
held
LS2. Under the likelihood, the mean error is coshε−1
∣relative error∣<0.02 ; z>5
relative error 0.004385; z above 5
held
LS3. Under the constructed loss, reuse shifts the error by −εsinhε
∣relative error∣<0.05 ; z>5
relative error -0.002878; z above 5
held
LS4. Under the constructed loss, the independent copy’s mean error is zero
z<2 and under 5 % of LS2’s error
z=2.033 , 0.91 %
missed: z>2
LS5a. Dividing the reconstruction by coshε removes the likelihood’s error
z<2 and under 5 % of it
z=0.29 , 0.10 %
held
Appendix
Table 3: Log-softmax: every registered statement and its outcome. “Likelihood” is the negative log-likelihood with a fixed target, “constructed” the loss whose incoming gradient equals the rounding error. Errors are per entry, divided by the entry’s probability; relative errors are against the registered constant.
Signs (σ,τ) of the two samples
h=σa (shared sample’s deviation)
h=a (fixed)
(+,+)
0
0
(+,−)
2a2
2a2
(−,+)
2a2
−2a2
(−,−)
0
0
Predicted mean
a2=0.0625
0
Measured mean
0.0625
0
Appendix
Table 4: FP8 attention: the eight predicted cells and the measurement. T is the query gradient from the independent sample minus that from the shared one, divided by the softmax scale.
390M, 20k schedule
390M, store runs
1.5B
Used in
Sections 2.2 – 2.3
Section 3
Section 2.2
Layers; width; MLP width
16; 1024; 4096
16; 1024; 4096
16; 2048; 8192
Query heads; key-value heads
8; 4
8; 4
16; 8
Input and output embeddings
tied
tied
untied
Sequences × tokens per step
128×4096
128×4096
256×4096
GPUs × micro-batch × accumulation
8×4×4
8×4×4
8×4×8
Appendix
Table 5: Training recipes. “Decay” is the length of the final decay; “none” holds the peak rate to the end. The forks of Sections 2.1 and 2.3 use the first column’s model and batch on their own schedules (Appendices D.2 and E ).
Build
PyTorch
TE
cuDNN
Runs
Earlier
2.11.0+cu128
2.10.0+769ed778
9.16
the forks of Figure 1 a; every FP8 run of Figure 2 and its seed-0 and seed-1 BF16 runs; the runs of Section 2.3 ; the 1.5B runs
Newer (cu129)
2.13.0+cu129
2.13.0+28777046
not named
the seed-2 BF16 run of Figure 2 ; a seed-0 BF16 repeat; the 6-bit store runs
Newer (cu130)
2.13.0+cu130
2.13.0+28777046
9.20
Figure 1 b; the 20k runs on newer software (Section 2.2 ) and their own BF16 runs; the 5-bit store runs; the continuations of Section 3.3 ; the TE test of Section 4 , on one H20
Appendix
Table 6: Software builds. PyTorch and TE build strings, with cuDNN versions where recorded.
Policy
Gain use reads ( zg )
Weight use reads ( uW )
6 bits
5 bits
U/R (reference)
z
uq
✓
✓
U/U
z
u
✓
✓
U/N
z
uf
✓
R/R
uq/γ
uq
✓
✓
R/N
uq/γ
uf
✓
N/R
uf/γ
uq
✓
Appendix
Table 7: Store policies. The value each use reads, and the store widths at which each policy was trained. The code names the gain use first. BF16 is the model trained without the store.
Run
Seed
Outcome
Late training
Validation
390M, 20k schedule, earlier software; reference: BF16 run of the same seed (seed 2’s: cu129 build)
TE’s FP8 attention
0, 1, 2
fails at 4500, 5500, 5000
BF16 backward
0
useful
0.0037
0.0029
1
useful
0.0071
0.0059
2
useful
0.0061
0.0052
Replacement of saved output
0
useful
0.006281395
0.0054
Appendix
Table 8: Further runs of Section 2 . For runs that train usefully, the gaps in nats to the block’s reference: late training loss (the mean over the last 200 steps) and final validation loss. For failing runs, the step, counted from the start of training, at which the excursion that fails the run opens.
What the backward reads
Label
Late max logit
Late training
Validation
(a) TE’s FP8 backward, through the wrapper; reference: the emulated backward with D from P
TE’s output, passed through
fails
26368.0
+0.0876
+0.0454
New rounding of O′ , rounding seed 1
useful
166.0
+0.0009
−0.0010
New rounding of O′ , rounding seed 2
useful
218.0
+0.0010
−0.0010
Nearest rounding of O′
fails
11520.0
+0.0282
+0.0217
Reference
197.0
Appendix
Table 9: Runs forked at step 1999, grouped by construction. Labels by each block’s rule; a failing run’s training loss had an excursion. Late max logit: the largest attention logit over all layers in steps 2400 to 2999. Gaps in nats to the block’s reference: mean training loss over steps 2800 to 2999, and validation loss at step 2999.
Reverse-mode differentiation computes every weight gradient, writes it to memory, and only then lets the optimizer read it back. This two-phase schedule sets the memory ceiling of modern training: at the seam between the phases, every layer's gradient is live at once. We argue that this materialized gradient is an artifact of how differentiation is staged, not a quantity that learning requires -- and we eliminate it. FORGE folds the optimizer step into the backward pass and applies it one tile at a time, entirely in registers, so each gradient tile is consumed the instant it is produced and never becomes a tensor. The fusion changes only when the update happens, not what it computes: in full precision the fused step is provably exact -- the identical optimizer update, for every element-wise rule -- and that exactness survives tensor- and sequence-parallel sharding; in the bf16 and 8-bit regimes used in practice it is faithful rather than bit-identical, its deviation bounded and, for the weight store, rendered unbiased by stochastic rounding. Because each gradient tile is born and consumed in the same registers, it is never converted down to bf16 to be stored and read back; FORGE thus preserves the full-precision fidelity that both bf16 and 8-bit optimizers lose to that conversion. Nor is the method tied to one architecture or one optimizer: linear layers are ubiquitous, and FORGE reclaims the gradient memory of any of them under any element-wise rule. Empirically FORGE more than halves the memory of an optimizer step and, at the small batch sizes typical of fine-tuning and continued pretraining, runs about 1.5x faster; integrated into tensor-parallel Megatron-LM it fits 8B training at four times the micro-batch a standard optimizer allows on the same GPUs.
Mamba-style and hybrid language models compress their past into a fixed-size recurrent state that is rewritten at every generated token. Storing this state in low precision saves memory bandwidth, but every rounding error is fed back into the next update and can accumulate over long generations. Production systems round the state stochastically; we ask which rounding rule such caches should use. We find that a deterministic golden-ratio Weyl dither, which needs no random numbers, consistently brings the quantized model closer to the full-precision one than stochastic rounding, across pure and hybrid models, storage formats, and long decoding horizons, at no extra cost. Round-to-nearest behaves differently: because it discards small updates, its error keeps growing, so it can look best in short evaluations yet falls far behind over long generations. A discrepancy analysis explains this ordering, and we document implementation pitfalls that silently remove the benefit.
Training in reduced floating-point precision can silently halt learning: when a gradient-descent weight update falls below half the unit in the last place (ULP) of the weight, it rounds away and that coordinate freezes while its gradient is still nonzero. The freeze is deterministic, governed by a per-coordinate half-ULP condition, and predictable from a high-precision trajectory and the target mantissa length alone, without low-precision data. In a small GPT trained under the standard AdamW-plus-cosine recipe with bf16-equivalent stored weights, training proceeds normally and then permanently freezes just past mid-run, within four steps of the a-priori prediction. In a 124-million-parameter GPT-2 transformer whose weights are constrained to the 8-bit floating-point grid after every optimizer step, with no master weights, the dense weights freeze at initialization in both fp8 formats -- predicted \emph{a priori} from an fp32 reference -- and validation loss plateaus while full precision keeps improving. Stochastic rounding removes the persistent freeze, and the same reference predicts that too. The condition transfers across frozen-feature regression, a mantissa-truncation emulator spanning 128× in precision, small networks, and a CNN on MNIST: a computable axis of low-precision training, not diffuse noise.
Zekai Shang
University of Illinois at Urbana-Champaign, Champaign, IL, USA