🤖 AI Summary
This work addresses the computational redundancy and state loss bottlenecks in Transformer feedforward architectures, where deep-layer information cannot propagate back to shallow layers. To overcome these limitations, we propose the Latent Information Feedback Transformer (LIFT), which pioneers a latent information feedback mechanism during pretraining. Specifically, LIFT leverages the pretrained model distribution to generate dense states and employs teacher forcing to reformulate recurrent learning as a parallelizable prediction task, thereby transcending conventional unidirectional information flow constraints. With only minimal parameter overhead, LIFT significantly outperforms standard Transformers on language modeling and reasoning benchmarks. Notably, smaller LIFT models surpass baseline counterparts trained on eight times more data, while its inference overhead decreases relative to model scale.
📝 Abstract
Transformer language models (LMs) are feed-forward: deep-layer representations are never fed back to shallower layers, and the only pathway for information to flow downward across generation steps is the decoded token. This narrow channel forces models to recompute intermediate results and to discard alternative continuations. In this work, we remove this bottleneck during pretraining, introducing the LIFT (Latent Information Feedback Transformer) architecture and training method which enable LMs to propagate state across generation. We achieve this by turning recurrent-state learning into a teacher-forced prediction problem: each input token is paired with an information-dense state, derived from the next-token distribution of an off-the-shelf pretrained LM. The model, extended with a small number of additional parameters, is then trained to predict both the next token and the next state. As the input states are precomputed, pretraining remains fully parallel across positions. At inference, the model's own predicted states are fed back, with a minor computational overhead that decreases with model size. Experiments with pretrained models ranging from 135M to 1B parameters show that LIFT consistently outperforms standard Transformers and baselines on language modeling, downstream reasoning tasks, and procedural tasks under token-matched budget, while being on par with or ahead of compute-matched Transformers. Moreover, a controlled study on a state-tracking task shows that a tiny LIFT outperforms same-size Transformers trained on 8x more data, even when trained with the states of a Transformer that fails the task. Overall, we show that LMs can learn to exploit deep-to-shallow feedback during pretraining via scalable teacher supervision.