Improving Efficient Recurrent Depth Via Next-Latent Alignment
Abstract
Looped Transformers share parameters across depth, allowing for parameter-efficient language modeling at the cost of latency and KV-cache size scaling with loops. Parallel Loop Transformers (PLTs) condition each loop on the previous loop's output, at the previous token position, removing the latency penalty incurred by looped Transformers by allowing all loops to decode simultaneously. However, PLTs face degraded absolute and per-loop performance relative to looped Transformers. This paper introduces a simple intervention that improves the performance of PLTs, while preserves decode-time parallelism and adding only a small parameter and compute overhead. By instantiating a small MLP at the loop boundary and training it jointly with the model to predict same-position hidden states at previous loops, we restore monotonic improvement with loop count, recover most of a looped Transformer's per-loop gain, and close a substantial fraction of the remaining gap to a parameter- and FLOP-matched looped Transformer. Our results suggest that parallel recurrence works best when the state transitions available at decoding time are trained alongside the representations that depend on them.