Research
Quantifying the Stability of Multi-Step Reasoning via Error Amplification
Quantifying the Stability of Multi-Step Reasoning via Error Amplification Overview Research area: Machine learning theory and large language model training, specifically the stability and error propag

- arXiv
- 2610.06404
- Published
- 2026-10-05
- Authors
- Dongyue Li, Ziniu Zhang, Minxuan Duan, Hongyang R. Zhang
AI summary
Quantifying the Stability of Multi-Step Reasoning via Error AmplificationOverview
Research area: Machine learning theory and large language model training, specifically the stability and error propagation of multi-step / chain-of-thought reasoning in autoregressive transformers.
Technical level: Advanced. The paper combines Jacobian-based error analysis, convergence proofs for transformer training dynamics on linear and quadratic regression tasks, and empirical fine-tuning experiments on algorithmic reasoning benchmarks.
Scope: The paper derives an "error amplification factor" from products of input-space Jacobian spectral norms, proves when transformers converge to a stable regime, and proposes a training method combining chain-of-thought length compression with quantization-aware training to control accumulated inference error.
What This Paper Is About
When a language model solves a problem step by step, each generated step becomes part of the input for the next one. A small mistake early on can therefore be fed back into the model and grow as generation continues, so the final-answer error can be much larger than the error at any individual step. The authors ask what actually determines the stability of such multi-step reasoning, and whether that quantity can be measured, bounded, predicted, and directly controlled during training.
Key Contributions
-
A measurable error-amplification bound. Using a first-order Taylor expansion through the input space, the authors prove (Lemma 3.1) that the gap between the inference loss and the loss conditioned on correct intermediate steps is governed by products of the spectral norms of input-space Jacobians, summarized as the error amplification factor ∑_{i=1}^{T-1} ρ_i^(T). This factor can grow exponentially with the number of reasoning steps T.
-
A convergence theory showing the factor can decay. For one-layer multi-head transformers trained on in-context weight-prediction tasks for linear and quadratic functions, the authors prove the model converges to a solution where the per-step Jacobian spectral norms are strictly less than one, so the amplification factor decays with T and the inference loss is O(1/poly(d)) with T = Θ(log d).
-
A novel three-head transformer construction. For the quadratic-function setting, Lemma 3.3 shows there exist weight matrices and elementwise nonlinearities such that a three-head transformer exactly implements a gradient descent step, extending prior one-head analyses of linear functions.
-
A training method with two knobs. The authors propose (i) chain-of-thought length compression, which uniformly subsamples T′ = ⌊λ·T⌋ intermediate steps, and (ii) quantization-aware training using straight-through estimators, hypothesized to regularize the trace of the loss Hessian and therefore the input Jacobian norms.
Main Findings
-
The bound tracks real inference losses. On Bellman-Ford, breadth-first search, Dijkstra, and Prim's algorithm with Qwen-1.5B, the inference loss and the bound both scale exponentially with T (measured for T from 1 to 6), and the bound qualitatively tracks the loss growth. The authors state the bound is worst-case and not necessarily numerically tight: the bound-to-gap ratio is around 8 for Qwen-1.5B and around 10^5 to 10^6 for Llama-1B and Gemma-2B.
-
Theoretically stable training exists. In the linear-function setting, the trained transformer reaches J_{i-1}^{(i)} ≈ I − (η/n)XX^⊤, so ρ_i^(T) ≈ ‖I − (η/n)XX^⊤‖^(T−i) and the accumulated errors form a converging geometric series. A similar form holds for quadratic functions with D(θ_i) = diag(3(X^⊤θ_i)^⊙2 − (X^⊤θ*)^⊙2).
-
Near-zero inference loss is empirically observed in the controlled setting. Figure 3 reports that inference loss converges close to zero for both linear functions (one-head linear transformer, d = 10, n = 20, η = 0.4) and quadratic functions (three-head nonlinear transformer, n = 200, η = 10^{-3}), with Gaussian noise σ = 0.002, and that longer steps yield lower inference loss.
-
Length compression reduces amplification dramatically. Intermediate step counts achieve the lowest inference loss on Dijkstra and Prim. Subsampling steps reduces the error amplification factor by 92× and 181× compared to training with all steps, though it increases the per-step errors ε_t, so the best λ balances the two effects.
-
Quantization further reduces amplification. With 1-bit quantization, ∑_{t=1}^{T-1} ρ_t^(T) is reduced by 3.6× and 5.8× compared to full precision on Dijkstra and Prim respectively, even though quantization increases per-step prediction error.
-
Improved accuracy across seven evaluations. Across five CLRS-Text graph algorithmic tasks and two LEGO symbolic state-tracking tasks, the method averages a 3.5% improvement over baselines. Best reported accuracies for Algorithm 1 include Dijkstra 76.2 ± 1.7, Bellman-Ford 80.3 ± 0.1, Prim 66.2 ± 0.3, BFS 88.3 ± 1.3, DFS 74.4 ± 1.5, Cyclic 94.1 ± 1.3, and Symmetry 62.2 ± 0.8, all averaged over three random seeds.
-
Better length generalization. On inputs 10% longer than training sequences, the method beats baselines by an average of 8.2%, with accuracies such as Dijkstra 44.3 ± 0.8, Bellman-Ford 50.1 ± 0.1, Prim 42.5 ± 0.6, BFS 59.6 ± 0.6, DFS 35.4 ± 0.8, Cyclic 87.2 ± 0.2, and Symmetry 59.1 ± 0.7. At 20% longer inputs, absolute accuracy falls below 25% for every method, so the authors report the 10% setting as their main result.
-
Comparison against latent reasoning. Combining sequence compression with quantization-aware training surpasses the latent-reasoning baselines (implicit CoT and Coconut) by up to 6.5%.
-
Direct measurement of regularization on Qwen-1.5B. Across the three tasks in Figure 1, the method reduces ∑_{i=1}^{T-1} ρ_i^(T) by 43% relative to baselines, with a 72% average reduction in inference loss.
Methodology in Plain English
The authors start by writing down what happens when a model's own predictions replace the correct intermediate steps. Using a first-order Taylor expansion through the input space, they express the difference between the inference-time loss and the training-time loss as a sum over generation steps, where each step contributes its per-step error multiplied by a factor ρ that sums the products of Jacobian spectral norms along all propagation paths from that step forward. This ρ-based sum is the error amplification factor.
They then pick two problems simple enough to analyze fully: predicting the weight vector of a linear function, and of a quadratic function, from in-context examples where the ground-truth intermediate steps are gradient descent iterates. They track the training dynamics of a one-layer multi-head transformer and show it converges to a solution whose per-step Jacobian is approximately I minus a scaled data covariance term, which makes the amplification factor a decaying geometric series. They also construct an explicit three-head transformer that carries out a gradient descent step for the quadratic case.
For the practical algorithm, they take the theoretical message literally: shorten the reasoning chain to reduce the number of multiplication factors, and regularize the Jacobian norms themselves. Length compression uniformly subsamples a fraction λ of the intermediate steps during training. Quantization-aware training simulates low-bit weight rounding in the forward pass and uses a straight-through estimator for gradients, motivated by the idea that this penalizes the Hessian trace, which decomposes into the input-space Jacobians.
Empirically, they measure Jacobian spectral norms with power iteration using Jacobian-vector products on the continuous input-embedding space, taking the maximum norm across the dataset, and they fine-tune Qwen-1.5B, Llama-1B, and Gemma-2B with LoRA (rank 16 to 64) and AdamW (learning rate 1×10⁻⁵ to 4×10⁻⁵), sweeping λ over 0.1, 0.2, 0.5, and 1, and bit-widths over 1, 2, 3, and 4 bits.
Why This Matters
For research, the paper turns vague intuitions about "error compounding" in chain-of-thought into a concrete, computable quantity tied to input-space Jacobian norms, and it connects a training-time regularizer to that quantity with both a proof in a tractable setting and measurements on real fine-tuned models. It also extends the line of work proving what transformers converge to on in-context regression tasks from linear to quadratic functions.
Real-world applications:
- API and agent pipelines that run multi-step tool calls or multi-hop reasoning, where one wrong intermediate output can corrupt an entire trajectory.
- Algorithmic and code-execution assistants, such as models that simulate graph algorithms, data structures, or program state, where a single mis-tracked variable invalidates everything downstream.
- Robotics and control planning, where plans are executed as long sequences of dependent decisions and late-stage failures are costly.
- Long-horizon document and data processing, where each extracted field feeds into the next stage of a pipeline.
For industry, the practical appeal is that the intervention is a fine-tuning recipe — subsampling training chains and quantizing weights — rather than a redesign of the model architecture. Quantization-aware training matters doubly because it is already used to shrink models for deployment, and this work suggests it also reduces error propagation.
Future Directions
-
Better length generalization. The paper reports that at 20% longer inputs every method, including theirs, falls below 25% accuracy, and explicitly leaves improving this to future work. The bound suggests two separate mechanisms to attack: covariate shift in per-step errors, and extra geometric multiplication when the step-to-step Jacobian norm exceeds one.
-
Tightening the bound. Remark 3.2 states the bound is worst-case and the bound-to-gap ratio ranges from about 8 (Qwen-1.5B) to 10^5–10^6 (Llama-1B, Gemma-2B), so a sharper, more numerically faithful version is an open problem.
-
Choosing λ and bit-width principledly. The best λ must trade off a reduced amplification factor against larger per-step errors ε_t; the paper says only that the optimum balances these, without a closed-form rule.
-
Extending the theory beyond linear and quadratic functions. The convergence proofs cover one-layer multi-head transformers on weight-prediction tasks; whether the same decay of the amplification factor holds for deeper models and more complex tasks is not established.
The paper states that code for replicating the empirical findings is available at https://github.com/VirtuosoResearch/Multi-step-reasoning-experiments.
Target Audience
Machine learning theory researchers working on transformer training dynamics and in-context learning; LLM practitioners interested in chain-of-thought reliability, latent reasoning, and length generalization; and engineers deploying quantized models for long-horizon reasoning pipelines, who want a training-time lever backed by both proof and measurement. Readers need comfort with Jacobians, spectral norms, and convergence analysis, though the "What This Paper Is About" framing is accessible to anyone familiar with chain-of-thought prompting.
Authors’ abstract
We consider the stability of multi-step reasoning processes, which have extensive applications in language models, including chain-of-thought and algorithmic reasoning. While longer sequences of reasoning can improve a model's generation capability at test time, the errors due to intermediate reasoning steps can accumulate in autoregressive generation, and thus grow substantially at the end. In this paper, we ask: What are the key factors determining the stability of multi-step reasoning? First, we show an inference error bound governed by the product of spectral norms of the Jacobians taken through the input space across generation steps. This product can be viewed as an error amplification factor, which could scale exponentially with the number of reasoning steps, serving as a quantitative measure of reasoning stability. Second, we analyze this measure in transformer models trained to predict simple tasks like linear and quadratic functions. We theoretically prove that the transformer model converges to a solution where the stability measure decays, thus yielding nearly zero inference loss over (arbitrarily) long steps. Finally, the stability analysis leads to several algorithmic implications for controlling the stability, through (i) chain-of-thought length compression that reduces the sensitivity of each step, and (ii) quantization-aware training that regularizes the input Jacobian norms. We validate the proposed algorithms by fine-tuning language models on graph-algorithmic reasoning tasks and symbolic state-tracking tasks. Across seven evaluations, our algorithms improve over baseline comparisons by 3.5% on average, and by 8.2% for longer-length inputs. Ablation analysis validates that the stability measure is drastically reduced by 3-8$\times$, confirming the regularization effect on the spectral norms of the (input space) Jacobians.