cs.LGSep 30, 2026

Backward-State Policy Is Part of the Learning Algorithm

Authors: Shuxiao Xie, Shuyang Xie, Dezhi Ran, Wei Yang, Tao Xie

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

Explore similar work

Jun 22, 2026cs.LG

FORGE: Fused On-Register Gradient Elimination for Memory-Efficient LLM Training

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.
Sep 30, 2026cs.LG

Low-Discrepancy Dither for Quantized Recurrent State Caches

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.
Jul 9, 2026cs.LG

The Silent Freeze: Predicting When Low-Precision Training Stops Learning

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 124124-million-parameter GPT-2 transformer whose weights are constrained to the 88-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×128\times in precision, small networks, and a CNN on MNIST: a computable axis of low-precision training, not diffuse noise.