Backward-State Policy Is Part of the Learning Algorithm
Organizations: Beijing Tongming Lake Information Technology Application Innovation Center (TLAIC), China
Abstract
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
Appendix figures & tables9 assets
Supplementary material from the paper’s appendix.
Appendix
| Policy | Softmax’s backward reads | Gradient of reads | Row of | |
|---|---|---|---|---|
| U/R (reference) | ||||
| R/R | ||||
| U/U | ||||
| N/N | ||||
| R/U |
| Statement | Registered rule | Measured | Outcome |
|---|---|---|---|
| First configuration: width 256, | |||
| 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 | and under 10 %; baseline | , 0.19 %; baseline | held |
| S3. It is kept when recomputation includes the loss | , the baseline’s sign, over 30 % | , same sign, 106.12 % | held |
| S4. Without restoration, the dropout mask changes | agreement between and ; 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 |
| Statement | Registered rule | Measured | Outcome |
|---|---|---|---|
| First configuration: | |||
| LS1. Under the likelihood, reuse changes nothing on average | and under 5 % of the constructed loss’s effect | , 0.17 % | held |
| LS2. Under the likelihood, the mean error is | ; | relative error 0.004385; above 5 | held |
| LS3. Under the constructed loss, reuse shifts the error by | ; | relative error -0.002878; above 5 | held |
| LS4. Under the constructed loss, the independent copy’s mean error is zero | and under 5 % of LS2’s error | , 0.91 % | missed: |
| LS5a. Dividing the reconstruction by removes the likelihood’s error | and under 5 % of it | , 0.10 % | held |
| Signs of the two samples | (shared sample’s deviation) | (fixed) |
|---|---|---|
| Predicted mean | ||
| Measured mean | 0.0625 | 0 |
| 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 | |||
| GPUs micro-batch accumulation |
| 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 |
| Policy | Gain use reads ( ) | Weight use reads ( ) | 6 bits | 5 bits |
| U/R (reference) | ✓ | ✓ | ||
| U/U | ✓ | ✓ | ||
| U/N | ✓ | |||
| R/R | ✓ | ✓ | ||
| R/N | ✓ | |||
| N/R | ✓ |
| 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 |
| What the backward reads | Label | Late max logit | Late training | Validation |
|---|---|---|---|---|
| (a) TE’s FP8 backward, through the wrapper; reference: the emulated backward with from | ||||
| TE’s output, passed through | fails | 26368.0 | +0.0876 | +0.0454 |
| New rounding of , rounding seed 1 | useful | 166.0 | +0.0009 | |
| New rounding of , rounding seed 2 | useful | 218.0 | +0.0010 | |
| Nearest rounding of | fails | 11520.0 | +0.0282 | +0.0217 |
| Reference | 197.0 | |||