Skip to content
AI.info

Research

Towards Looped Models Done Right, Part II: Rethinking at Fixed Points

Towards Looped Models Done Right, Part II: Rethinking at Fixed Points Overview Research area: Machine learning — efficient training and inference for looped (recurrent-depth) language models. Technica

Towards Looped Models Done Right, Part II: Rethinking at Fixed Points
arXiv
2610.06833
Published
2026-10-05
Authors
Benhao Huang, Chufan Shi, Junlin Chen, Shicheng Wen, Zhengzhong Liu, Eric Xing, Xuezhe Ma

AI summary

Towards Looped Models Done Right, Part II: Rethinking at Fixed Points

Overview

Research area: Machine learning — efficient training and inference for looped (recurrent-depth) language models.

Technical level: Advanced. The paper leans on fixed-point theory, implicit differentiation, Neumann series truncation, and reinforcement-learning post-training, though the experiments are reported in standard perplexity and benchmark terms.

Scope: The paper argues that driving a looped model's recurrent states toward fixed points makes four depth-dependent costs — training backpropagation, KV caching, prefill, and RL rollout scoring — shortcuttable, and it improves the two training ingredients that shape those fixed points: the depth prior and input injection. All results reported here come from the provided content, which is truncated at the RL experiment description (arXiv:2610.06833v1, dated 2026-10-05, by Benhao Huang, Chufan Shi, Junlin Chen, Shicheng Wen, Zhengzhong Liu, Eric Xing, and Xuezhe Ma, affiliated with the Institute of Foundation Models, USC, and CMU; code at https://github.com/ifm-ai/xllm-loop).

What This Paper Is About

Looped language models reuse a shared block of layers many times, so every extra recurrence adds cost in training, decoding, prefill, and reinforcement learning. The authors show that if a model's recurrent states settle near a fixed point, the endpoint of the recurrence can stand in for the whole trajectory, letting each of those costs be cut. The paper then fixes the two mechanisms that determine how good those fixed points are — how training depths are sampled and how the input is injected at each recurrence — and measures the gains from 100M to 1.6B parameters.

Key Contributions

  1. A unified explanation of fixed-point shortcuts. The paper shows that truncated backpropagation through time (TBPTT) alone shapes the fixed points that terminal KV sharing depends on, whereas fixed-depth training with full backpropagation through time (BPTT) does not. At 1.6B, KV sharing drops fixed-depth training's GSM8K accuracy from 50.6 to 21.2.

  2. A learned depth prior replacing Huginn's fixed prior. The prior is learned from prediction feedback via a categorical distribution over depths, with an entropy term that keeps it broad enough for KV sharing and a budget-control term that keeps the expected depth near the target. From 100M to 1.6B, it lowers perplexity and raises the downstream average over Huginn's fixed prior; at 1.6B, with a 3× smaller KV cache, it matches the full-cache downstream average of fixed-depth training.

  3. Orthogonal input injection (OrthoInj). The paper identifies that existing injection schemes leave the state's component along the input unconstrained, so it can amplify or cancel the injected input. OrthoInj projects that component out with an orthogonal projection matrix, keeping the input's strength constant at every recurrence. It improves over the strongest baseline, Parcae-Decay, at every scale.

  4. Two new shortcuts enabled by fixed points. Distilled prefill, trained on a quarter of the teacher's pretraining tokens, prefills up to 1.79× faster on 8K prompts; and rollout-state reuse halves the scoring and backward time of each RL update relative to replaying the loops with full BPTT.

Main Findings

  • Convergence is heterogeneous. Using a tokenwise relative-change measure with threshold τ = 2% and ε = 10⁻⁸, the paper finds that tokens in a sequence in Huginn converge at different depths, not necessarily in causal order, and that convergence depth correlates only weakly with position (Spearman correlation 0.14). Figure 2 reports medians over 128-position intervals.

  • TBPTT both stabilizes training and shapes fixed points. Full BPTT at fixed depth R = 5 fails two fixed-point tests: extrapolating to R = 64 raises perplexity by over 25%, and terminal KV sharing raises it by 21%. TBPTT through the last two recurrences (b = 2) raises perplexity by at most 2% and 3% on those same tests. With sampled depths enabled, training with b ≥ 4 matches full BPTT within 1% perplexity at 400M.

  • Fixed-depth training breaks KV sharing. Under forced terminal KV sharing at R = 5, fixed-depth training has the highest perplexity of the compared priors, and sharing cuts its GSM8K accuracy by over half at the M and L scales (the abstract reports the L-scale drop from 50.6 to 21.2).

  • The learned prior keeps sharing intact. With terminal KV sharing, it lowers validation perplexity relative to Huginn's Fixed PLN-5 by 0.7–1.8% for at most 1.6% more training FLOPs, and raises the downstream average at every scale. The entropy term matters especially at L: without it, GSM8K falls from 47.92 to 36.47. At S scale, freezing the prior at its distribution from the first update lowers validation perplexity by a further 0.8%.

  • The learned prior narrows the looped-versus-unrolled gap. Its downstream-average lead over Untied 4 (same parameters and KV cache) grows from 2.2 points at S to 5.0 at M and 8.9 at L. Its validation perplexity exceeds that of Untied 12 (same logical depth, 3× the non-embedding parameters and KV cache) by 13.6% at S, 8.9% at M, and 6.3% at L. At L it trails Untied 12 by 1.5 points on the downstream average and 0.9 on GSM8K, while leading Untied 4 by 40.7 on GSM8K.

  • OrthoInj beats existing injection schemes at every scale. With terminal KV sharing and PLN-5, OrthoInj achieves the lowest validation PPL (5.67 at S, 3.94 at M, 3.08 at L) and the highest downstream average compared with DEQ-QKV, Huginn-Linear, and Parcae-Decay. It lowers validation perplexity by 0.3–1.4% and ends training with the lowest loss at every scale. Dropping prelude normalization and the projection each help, and their gains add up.

  • Terminal KV sharing is cheap under the right training. At training depth R = 5, sharing keeps 4 of 12 banks and raises perplexity over the full cache by at most 0.85%. With tolerant fixed points, sharing cuts the KV cache 3× at five recurrences, keeping the entropy-regularized learned-prior models within 0.4 points of their full-cache downstream average.

  • Distilled prefill is faster but slightly behind. Trained on a quarter of the teacher's pretraining tokens, the student is up to 1.79× faster on 8K prompts and trails the teacher's downstream average by 1.0–4.5 points; at matched latency it scores 0.5–0.9 points above the teacher prefilling with two recurrences. The distillation loss combines normalized hidden-state regression with a next-token KL term.

  • Rollout-state reuse roughly halves RL update cost. Scoring and backward time per RL update is halved relative to replaying the loops with full BPTT, and the method stays within 1.6 points of that baseline's GSM8K pass@1. The update cost and memory are independent of the number of recurrences R, though rollout generation still takes most of the training time. The RL setup fine-tunes the L-scale learned-prior model with Dr. GRPO at R = 6; the provided text is truncated at this point, so the remaining RL configuration details are not reported in the content supplied.

  • Scale and setup. Models are built from dense Llama Transformer blocks in a 1+2×R+1 layout (one prelude block, a two-block recurrent core, one coda block), giving a logical depth of 2+2R blocks at R recurrences. Scales are S (100M parameters, 21.5B tokens), M (400M, 85.9B), and L (1.6B, 343.6B). Training uses a weighted multilingual mixture including TxT360 with the Jais64k tokenizer, and AdamW with 5% linear warmup then 10% cosine decay to one-tenth of the peak learning rate. Learned priors cap depth at R_max = 64 and start from PLN-5.

Methodology in Plain English

The researchers start from one idea: a looped model applies the same block over and over, and as it does so, its internal state tends to stop changing — it approaches a fixed point. At that point, the route taken to get there matters less than the destination. That means you can backpropagate through only the last few recurrences instead of all of them, cache only the last set of keys and values instead of one per recurrence, send a small distilled network ahead to predict the endpoint instead of running the loops during prefill, and, in reinforcement learning, reuse the states saved during rollout instead of replaying the whole trajectory to compute gradients.

Because these shortcuts only work if the model actually has good fixed points, the paper focuses on two training choices that determine them. The first is the depth prior — the distribution over how many recurrences each training example is run for. The authors replace Huginn's fixed Poisson–log-normal prior with a categorical distribution over depths whose logits are updated using the model's own prediction feedback (converted into an advantage with an EMA baseline, following PopArt-style rescaling), plus an entropy bonus to keep depth coverage broad and a penalty keeping the expected depth near a target budget. The second is input injection — how the input embedding is re-added at each recurrence. Existing schemes let part of the carried-over state point along the input direction, which can either strengthen or cancel the injection; OrthoInj subtracts that component with an orthogonal projection, so each recurrence receives the input at the same strength while the rest of the state evolves freely, and the projected carryover stays contractive.

Evaluation protocols use MixVal (a held-out set from the TxT360 sources) and WikiText-103 for perplexity; LAMBADA, HellaSwag, PIQA, ARC-Easy/Challenge, OpenBookQA, and SciQ for downstream accuracy; and GSM8K, DROP, and MBPP+ for generation. Unless noted, models are evaluated at R = 5 with terminal KV sharing, with perplexity computed by prefilling the first half of each 8,192-token sequence and scoring the second half in one parallel pass.

Why This Matters

Impact on research. The paper reframes the long-running debate over whether looped models should or should not chase fixed points, showing that the two camps differ not just in philosophy but in measurable behavior: Huginn's states approach fixed points and tolerate KV sharing, while Ouro's do not and its accuracy collapses at four loops. It also connects implicit-differentiation theory to practical training recipes by showing that TBPTT, unlike full BPTT, is a single stable estimator that needs no warm-up phase or switch and simultaneously shapes the fixed points that inference shortcuts rely on.

Real-world applications:

  • Serving deep looped models with a KV cache that does not grow with recurrence depth, which directly reduces memory pressure at inference.
  • Faster prompt prefill on long inputs (the distilled student is up to 1.79× faster on 8K prompts), which matters for long-context serving.
  • Cheaper reinforcement-learning post-training for recurrent architectures, since gradients are computed from saved rollout states rather than by replaying the trajectory.
  • A tunable knob (the entropy coefficient λ_H) for trading predictive performance against memory efficiency when a deployment budget is fixed.

Industry relevance. Any deployment that pays per recurrence — inference servers, long-context prefilling, and RL fine-tuning pipelines — sees cost scale with loop count. The paper's claim is that fixed-point shaping decouples those costs from depth, and the reported results at 100M to 1.6B parameters with public benchmarks give a concrete starting point for teams already using looped or recurrent-depth models.

Future Directions

  • Scaling beyond 1.6B. All reported comparisons run from 100M to 1.6B parameters; whether the learned prior's narrowing gap to Untied 12 continues at larger scales is untested in the provided content.
  • The student as a standalone language model or speculative-decoding drafter. The paper explicitly leaves these uses, and possible connections between recurrent state refinement and flow-based language modeling, to future work.
  • Tuning the performance–memory trade-off. The entropy coefficient λ_H is presented as a direct knob balancing predictive performance against KV-cache savings; how to set it automatically for a deployment target is left open.
  • Rollout-state reuse under fuller RL protocols. The provided content is truncated mid-description of the RL setup (Dr. GRPO at R = 6), and the paper notes that rollout generation still dominates training time, so the RL-side savings are partial rather than end-to-end.

Target Audience

This paper suits researchers and engineers working on recurrent-depth or looped language architectures, efficient inference and KV-cache design, or RL post-training for such models. Readers already comfortable with implicit differentiation, fixed-point contractions, and reinforcement-learning objectives will get the most out of it; practitioners focused on serving costs and post-training throughput will find the benchmark and cache results directly applicable, while newcomers will benefit from the paper's plain framing of why fixed points matter even if the mathematical derivations require background.

Authors’ abstract

Every recurrence of a looped language model adds cost in training, decoding, prefill, and reinforcement learning (RL). The closer recurrent states get to fixed points, the less the path to them matters. This enables truncated backpropagation in training; terminal key-value (KV) sharing for decoding with almost no loss in accuracy; a distilled student that prefills up to 1.79x faster; and RL updates that compute gradients from saved rollout states, 2x faster than backpropagating through the replayed trajectory. We therefore improve the two components of training that shape these fixed points: the depth prior and input injection. Fixed-depth training breaks KV sharing, and Huginn's broad depth prior supports sharing but dilutes supervision at the target depth more than sharing requires; we learn the prior from prediction feedback, with an entropy term that keeps it broad. Existing injection schemes let the state's component along the input amplify or cancel the injection; we remove this component with orthogonal injection. From 100M to 1.6B parameters, the learned prior and orthogonal injection lower perplexity at every scale relative to Huginn's prior and existing injection schemes, respectively. At 1.6B, the learned prior with a 3x smaller KV cache matches the downstream average of fixed-depth training with the full cache.

Read the original paper