LoopAhead: Parallelling Recurrent Depth In Looped Transformers
Abstract
Looped Transformers increase computation per token by repeatedly applying a shared backbone, offering a way to improve language modeling and downstream task performance at a fixed backbone parameter count. However, each loop depends on the preceding loop, so increasing recurrent depth also increases decoding latency. Existing approaches include Parallel Loop Transformers, which reduce latency by overlapping loop computation across token positions, but yield rapidly diminishing marginal accuracy gains at larger depths. To address this challenge, we introduce LoopAhead, which processes hidden states from multiple loop depths in parallel for the same token prefix. A lightweight predictor estimates inputs to later loops so they can start before earlier loops finish. A corrector then refines the outputs of these later loops using newly available results from earlier loops. All components are trained jointly using only the final next-token prediction loss. In experiments with 1.4B models, deeper LoopAhead configurations achieve progressively higher average zero-shot accuracy across eight tasks, rising from 53.41% to 56.38%. Vanilla looped models and PLT show limited gains over the same depth range. Compared with a nine-loop vanilla model, LoopAhead uses five sequential stages and additional parallel computation to reduce decoding latency by 46%.
est. 32% chance this paper gets accepted at ICLR 2027.
What do you think this paper will get?
All positions stay anonymous.