Looped transformers improve decoding efficiency without extra training

Decoding Looped Transformers Better for (Almost) Free

Machine Learning

Summary

Looped transformers use the same part of a model multiple times to save parameters. Usually, only the last pass helps pick the next word, but earlier passes contain useful clues that are wasted. The authors introduce LoopCD, a method that compares early and late passes during decoding to make better choices without extra training. This improves performance while allowing fewer passes, which saves computing work.

What this means in practice

  • For machine learning engineers: Cut inference compute costs in transformer-based language models by reducing recurrent loops while maintaining or improving output quality.
  • For software developers: Use more efficient decoding to speed up code generation models producing correct programs with fewer model passes.

Authors

Weihao Liu, Huangjie Zheng, Tianrong Chen, Rohit Dilip, Richard He Bai, Yizhu Jiao, Yuyang Wang, Ruixiang Zhang

Abstract

Looped Transformers achieve parameter efficiency by repeatedly executing a shared block across recurrent loops. Each loop yields an intermediate representation decodable for the same next token, yet standard decoding discards earlier states. Because earlier loops embody less computation, recurrence inherently supplies aligned weak-and-strong prediction pairs without auxiliary models or external training. We introduce LoopCD, a training-free contrastive decoding framework that guides token selection by contrasting the final prediction with an earlier recurrent pass, operating either in logit space with one extra output pass (LoopCD-Logits) or in hidden-state space with zero output overhead (LoopCD-Hidden). Across four looped Transformer families, LoopCD delivers substantial, consistent gains at full recurrent depth: LoopCD-Logits raises Ouro-2.6B-Thinking's AIME 2024 pass@1 from 61.88% to 73.33%, while LoopCD-Hidden lifts Huginn's HumanEval pass@1 from 22.56% to 31.71%. Crucially, these performance gains enable halving the number of recurrent loops while still matching or exceeding full-depth unguided baselines, reducing forward FLOPs by 22.5% to 48.2%. By transforming intermediate recurrent states into effective guidance signals, LoopCD achieves superior decoding quality while substantially reducing inference compute.