Better & Faster Large Language Models via Multi-token Prediction

Fabian Gloeckle, Badr Youbi Idrissi, Baptiste Rozière, David Lopez-Paz, Gabriel Synnaeve

Introduction

Humanity has condensed its most ingenious undertakings, surprising findings and beautiful productions into text. Large Language Models (LLMs) trained on all of these corpora are able to extract impressive amounts of world knowledge, as well as basic reasoning capabilities by implementing a simple—yet powerful—unsupervised learning task: next-token prediction. Despite the recent wave of impressive achievements (OpenAI, 2023), next-token prediction remains an inefficient way of acquiring language, world knowledge and reasoning capabilities. More precisely, teacher forcing with next-token prediction latches on local patterns and overlooks “hard” decisions. Consequently, it remains a fact that state-of-the-art next-token predictors call for orders of magnitude more data than human children to arrive at the same level of fluency (Frank, 2023).

In this study, we argue that training LLMs to predict multiple tokens at once will drive these models toward better sample efficiency. As anticipated in Figure 1, multi-token prediction instructs the LLM to predict the nn future tokens from each position in the training corpora, all at once and in parallel (Qi et al., 2020).

While multi-token prediction has been studied in previous literature (Qi et al., 2020), the present work offers the following contributions:

We propose a simple multi-token prediction architecture with no train time or memory overhead (Section 2).

We provide experimental evidence that this training paradigm is beneficial at scale, with models up to 13B parameters solving around 15% more code problems on average (Section 3).

Multi-token prediction enables self-speculative decoding, making models up to 3 times faster at inference time across a wide range of batch-sizes (Section 3.2).

While cost-free and simple, multi-token prediction is an effective modification to train stronger and faster transformer models. We hope that our work spurs interest in novel auxiliary losses for LLMs well beyond next-token prediction, as to improve the performance, coherence, and reasoning abilities of these fascinating models.

Method

Standard language modeling learns about a large text corpus x1,…xTx_{1},\ldots x_{T} by implementing a next-token prediction task. Formally, the learning objective is to minimize the cross-entropy loss

where PθP_{\theta} is our large language model under training, as to maximize the probability of xt+1x_{t+1} as the next future token, given the history of past tokens xt:1=xt,…,x1x_{t:1}=x_{t},\ldots,x_{1}.

In this work, we generalize the above by implementing a multi-token prediction task, where at each position of the training corpus, the model is instructed to predict nn future tokens at once. This translates into the cross-entropy loss

To make matters tractable, we assume that our large language model PθP_{\theta} employs a shared trunk to produce a latent representation zt:1z_{t:1} of the observed context xt:1x_{t:1}, then fed into nn independent heads to predict in parallel each of the nn future tokens (see Figure 1). This leads to the following factorization of the multi-token prediction cross-entropy loss:

In practice, our architecture consists of a shared transformer trunk fsf_{s} producing the hidden representation zt:1z_{t:1} from the observed context xt:1x_{t:1}, nn independent output heads implemented in terms of transformer layers fhif_{h_{i}}, and a shared unembedding matrix fuf_{u}. Therefore, to predict nn future tokens, we compute:

for i=1,…ni=1,\ldots n, where, in particular, Pθ(xt+1∣xt:1)P_{\theta}(x_{t+1}\mid x_{t:1}) is our next-token prediction head. See Appendix B for other variations of multi-token prediction architectures.

One big challenge in training multi-token predictors is reducing their GPU memory utilization. To see why this is the case, recall that in current LLMs the vocabulary size VV is much larger than the dimension dd of the latent representation—therefore, logit vectors become the GPU memory usage bottleneck. Naive implementations of multi-token predictors that materialize all logits and their gradients, both of shape (n,V)(n,V), severely limit the allowable batch-size and average GPU memory utilization. Because of these reasons, in our architecture we propose to carefully adapt the sequence of forward and backward operations, as illustrated in Figure 2. In particular, after the forward pass through the shared trunk fsf_{s}, we sequentially compute the forward and backward pass of each independent output head fif_{i}, accumulating gradients at the trunk. While this creates logits (and their gradients) for the output head fif_{i}, these are freed before continuing to the next output head fi+1f_{i+1}, requiring the long-term storage only of the dd-dimensional trunk gradient ∂Ln/∂fs\partial L_{n}/\partial f_{s}. In sum, we have reduced the peak GPU memory utilization from O(nV+d)O(nV+d) to O(V+d)O(V+d), at no expense in runtime (Table S5).

Inference

During inference time, the most basic use of the proposed architecture is vanilla next-token autoregressive prediction using the next-token prediction head Pθ(xt+1∣xt:1)P_{\theta}(x_{t+1}\mid x_{t:1}), while discarding all others. However, the additional output heads can be leveraged to speed up decoding from the next-token prediction head with self-speculative decoding methods such as blockwise parallel decoding (Stern et al., 2018)—a variant of speculative decoding (Leviathan et al., 2023) without the need for an additional draft model—and speculative decoding with Medusa-like tree attention (Cai et al., 2024).

Experiments on real data

We demonstrate the efficacy of multi-token prediction losses by seven large-scale experiments. Section 3.1 shows how multi-token prediction is increasingly useful when growing the model size. Section 3.2 shows how the additional prediction heads can speed up inference by a factor of 3×3\times using speculative decoding. Section 3.3 demonstrates how multi-token prediction promotes learning longer-term patterns, a fact most apparent in the extreme case of byte-level tokenization. Section 3.4 shows that 44-token predictor leads to strong gains with a tokenizer of size 3232k. Section 3.5 illustrates that the benefits of multi-token prediction remain for training runs with multiple epochs. Section 3.6 showcases the rich representations promoted by pretraining with multi-token prediction losses by finetuning on the CodeContests dataset (Li et al., 2022). Section 3.7 shows that the benefits of multi-token prediction carry to natural language models, improving generative evaluations such as summarization, while not regressing significantly on standard benchmarks based on multiple choice questions and negative log-likelihoods.

To allow fair comparisons between next-token predictors and nn-token predictors, the experiments that follow always compare models with an equal amount of parameters. That is, when we add n−1n-1 layers in future prediction heads, we remove n−1n-1 layers from the shared model trunk. Please refer to Table S14 for the model architectures and to Table S13 for an overview of the hyperparameters we use in our experiments.

To study this phenomenon, we train models of six sizes in the range 300M to 13B parameters from scratch on at least 91B tokens of code. The evaluation results in Figure 3 for MBPP (Austin et al., 2021) and HumanEval (Chen et al., 2021) show that it is possible, with the exact same computational budget, to squeeze much more performance out of large language models given a fixed dataset using multi-token prediction.

We believe this usefulness only at scale to be a likely reason why multi-token prediction has so far been largely overlooked as a promising training loss for large language model training.

2 Faster inference

We implement greedy self-speculative decoding (Stern et al., 2018) with heterogeneous batch sizes using xFormers (Lefaudeux et al., 2022) and measure decoding speeds of our best 4-token prediction model with 7B parameters on completing prompts taken from a test dataset of code and natural language (Table S2) not seen during training. We observe a speedup of 3.0×\mathbf{3.0\times} on code with an average of 2.5 accepted tokens out of 3 suggestions on code, and of 2.7×2.7\times on text. On an 8-byte prediction model, the inference speedup is 6.4×6.4\times (Table S3). Pretraining with multi-token prediction allows the additional heads to be much more accurate than a simple finetuning of a next-token prediction model, thus allowing our models to unlock self-speculative decoding’s full potential.

3 Learning global patterns with multi-byte prediction

To show that the next-token prediction task latches to local patterns, we went to the extreme case of byte-level tokenization by training a 7B parameter byte-level transformer on 314B bytes, which is equivalent to around 116B tokens. The 8-byte prediction model achieves astounding improvements compared to next-byte prediction, solving 67% more problems on MBPP pass@1 and 20% more problems on HumanEval pass@1.

Multi-byte prediction is therefore a very promising avenue to unlock efficient training of byte-level models. Self-speculative decoding can achieve speedups of 6 times for the 8-byte prediction model, which would allow to fully compensate the cost of longer byte-level sequences at inference time and even be faster than a next-token prediction model by nearly two times. The 8-byte prediction model is a strong byte-based model, approaching the performance of token-based models despite having been trained on 1.7×1.7\times less data.

4 Searching for the optimal n𝑛n

To better understand the effect of the number of predicted tokens, we did comprehensive ablations on models of scale 7B trained on 200B tokens of code. We try n=1,2,4,6n=1,2,4,6 and 88 in this setting. Results in table 1 show that training with 4-future tokens outperforms all the other models consistently throughout HumanEval and MBPP for pass at 1, 10 and 100 metrics: +3.8%, +2.1% and +3.2% for MBPP and +1.2%, +3.7% and +4.1% for HumanEval. Interestingly, for APPS/Intro, n=6n=6 takes the lead with +0.7%, +3.0% and +5.3%. It is very likely that the optimal window size depends on input data distribution. As for the byte level models the optimal window size is more consistent (8 bytes) across these benchmarks.

5 Training for multiple epochs

Multi-token training still maintains an edge on next-token prediction when trained on multiple epochs of the same data. The improvements diminish but we still have a +2.4% increase on pass@1 on MBPP and +3.2% increase on pass@100 on HumanEval, while having similar performance for the rest. As for APPS/Intro, a window size of 4 was already not optimal with 200B tokens of training.

6 Finetuning multi-token predictors

Pretrained models with multi-token prediction loss also outperform next-token models for use in finetunings. We evaluate this by finetuning 7B parameter models from Section 3.3 on the CodeContests dataset (Li et al., 2022). We compare the 4-token prediction model with the next-token prediction baseline, and include a setting where the 4-token prediction model is stripped off its additional prediction heads and finetuned using the classical next-token prediction target. According to the results in Figure 4, both ways of finetuning the 4-token prediction model outperform the next-token prediction model on pass@k across kk. This means the models are both better at understanding and solving the task and at generating diverse answers. Note that CodeContests is the most challenging coding benchmark we evaluate in this study. Next-token prediction finetuning on top of 4-token prediction pretraining appears to be the best method overall, in line with the classical paradigm of pretraining with auxiliary tasks followed by task-specific finetuning. Please refer to Appendix F for details.

7 Multi-token prediction on natural language

To evaluate multi-token prediction training on natural language, we train models of size 7B parameters on 200B tokens of natural language with a 4-token, 2-token and next-token prediction loss, respectively. In Figure 5, we evaluate the resulting checkpoints on 6 standard NLP benchmarks. On these benchmarks, the 2-future token prediction model performs on par with the next-token prediction baseline throughout training. The 4-future token prediction model suffers a performance degradation. Detailed numbers are reported in Appendix G.

However, we do not believe that multiple-choice and likelihood-based benchmarks are suited to effectively discern generative capabilities of language models. In order to avoid the need for human annotations of generation quality or language model judges—which comes with its own pitfalls, as pointed out by Koo et al. (2023)—we conduct evaluations on summarization and natural language mathematics benchmarks and compare pretrained models with training sets sizes of 200B and 500B tokens and with next-token and multi-token prediction losses, respectively.

For summarization, we use eight benchmarks where ROUGE metrics (Lin, 2004) with respect to a ground-truth summary allow automatic evaluation of generated texts. We finetune each pretrained model on each benchmark’s training dataset for three epochs and select the checkpoint with the highest ROUGE-L F1F_{1} score on the validation dataset. Figure 6 shows that multi-token prediction models with both n=2n=2 and n=4n=4 improve over the next-token baseline in ROUGE-L F1F_{1} scores for both training dataset sizes, with the performance gap shrinking with larger dataset size. All metrics can be found in Appendix H.

For natural language mathematics, we evaluate the pretrained models in 8-shot mode on the GSM8K benchmark (Cobbe et al., 2021) and measure accuracy of the final answer produced after a chain-of-thought elicited by the fewshot examples. We evaluate pass@k metrics to quantify diversity and correctness of answers like in code evaluations and use sampling temperatures between 0.2 and 1.4. The results are depicted in Figure S13 in Appendix I. For 200B training tokens, the n=2n=2 model clearly outperforms the next-token prediction baseline, while the pattern reverses after 500B tokens and n=4n=4 is worse throughout.

Ablations on synthetic data

What drives the improvements in downstream performance of multi-token prediction models on all of the tasks we have considered? By conducting toy experiments on controlled training datasets and evaluation tasks, we demonstrate that multi-token prediction leads to qualitative changes in model capabilities and generalization behaviors. In particular, Section 4.1 shows that for small model sizes, induction capability—as discussed by Olsson et al. (2022)—either only forms when using multi-token prediction as training loss, or it is vastly improved by it. Moreover, Section 4.2 shows that multi-token prediction improves generalization on an arithmetic task, even more so than tripling model size.

Induction describes a simple pattern of reasoning that completes partial patterns by their most recent continuation (Olsson et al., 2022). In other words, if a sentence contains “AB” and later mentions “A”, induction is the prediction that the continuation is “B”. We design a setup to measure induction capability in a controlled way. Training small models of sizes 1M to 1B nonembedding parameters on a dataset of children stories, we measure induction capability by means of an adapted test set: in 100 stories from the original test split, we replace the character names by randomly generated names that consist of two tokens with the tokenizer we employ. Predicting the first of these two tokens is linked to the semantics of the preceding text, while predicting the second token of each name’s occurrence after it has been mentioned at least once can be seen as a pure induction task. In our experiments, we train for up to 90 epochs and perform early stopping with respect to the test metric (i.e. we allow an epoch oracle). Figure 7 reports induction capability as measured by accuracy on the names’ second tokens in relation to model size for two runs with different seeds.

We find that 2-token prediction loss leads to a vastly improved formation of induction capability for models of size 30M nonembedding parameters and below, with their advantage disappearing for sizes of 100M nonembedding parameters and above.Note that a perfect score is not reachable in this benchmark as some of the tokens in the names in the evaluation dataset never appear in the training data, and in our architecture, embedding and unembedding parameters are not linked. We interpret this finding as follows: multi-token prediction losses help models to learn transferring information across sequence positions, which lends itself to the formation of induction heads and other in-context learning mechanisms. However, once induction capability has been formed, these learned features transform induction into a task that can be solved locally at the current token and learned with next-token prediction alone. From this point on, multi-token prediction actually hurts on this restricted benchmark—but we surmise that there are higher forms of in-context reasoning to which it further contributes, as evidenced by the results in Section 3.1. In Figure S14, we provide evidence for this explanation: replacing the children stories dataset by a higher-quality 9:1 mix of a books dataset with the children stories, we enforce the formation of induction capability early in training by means of the dataset alone. By consequence, except for the two smallest model sizes, the advantage of multi-token prediction on the task disappears: feature learning of induction features has converted the task into a pure next-token prediction task.

2 Algorithmic reasoning

Multi-token prediction improves algorithmic reasoning capabilities as measured by this task across task difficulties (Figure 8). In particular, it leads to impressive gains in out-of-distribution generalization, despite the low absolute numbers. Increasing the model size from 30M to 100M parameters, on the other hand, does not improve evaluation accuracy as much as replacing next-token prediction by multi-token prediction does (Figure S16). In Appendix K, we furthermore show that multi-token prediction models retain their advantage over next-token prediction models on this task when trained and evaluated with pause tokens (Goyal et al., 2023).

Why does it work? Some speculation

Why does multi-token prediction afford superior performance on coding evaluation benchmarks, and on small algorithmic reasoning tasks? Our intuition, developed in this section, is that multi-token prediction mitigates the distributional discrepancy between training-time teacher forcing and inference-time autoregressive generation. We support this view with an illustrative argument on the implicit weights multi-token prediction assigns to tokens depending on their relevance for the continuation of the text, as well as with an information-theoretic decomposition of multi-token prediction loss.

Not all token decisions are equally important for generating useful texts from language models (Bachmann and Nagarajan, 2024; Lin et al., 2024). While some tokens allow stylistic variations that do not constrain the remainder of the text, others represent choice points that are linked with higher-level semantic properties of the text and may decide whether an answer is perceived as useful or derailing.

Multi-token prediction implicitly assigns weights to training tokens depending on how closely they are correlated with their successors. As an illustrative example, consider the sequence depicted in Figure 9 where one transition is a hard-to-predict choice point while the other transitions are considered “inconsequential”. Inconsequential transitions following a choice point are likewise hard to predict in advance. By marking and counting loss terms, we find that nn-token prediction associates a weight of n(n+1)2\frac{n(n+1)}{2} to choice points via their correlates, and a smaller weight of nn to inconsequential points. Please refer to Appendix L.3 for more details. Generally, we believe that the quality of text generations depends on picking the right decisions at choice points, and that nn-token prediction losses promote those.

2 Information-theoretic argument

Language models are typically trained by teacher-forcing, where the model receives the ground truth for each future token during training. However, during test time generation is unguided and autoregressive, whereby errors accumulate. Teacher-forcing, we argue, encourages models to focus on predicting well in the very short term, at the potential expense of ignoring longer-term dependencies in the overall structure of the generated sequence.

To illustrate the impact of multi-token prediction, consider the following information-theoretic argument. Here, XX denotes the next future token, and YY the second-next future token. The production of both of these tokens is conditioned on some observed, input context CC, that we omit from our equations for simplicity. When placed before token XX, vanilla next-token prediction concerns the quantity H(X)H(X), while multi-token prediction with n=2n=2 aims at H(X)+H(Y)H(X)+H(Y). We decompose these two quantities as:

By discarding the term H(Y∣X)H(Y\mid X)—which appears again when predicting at the following position—we observe that 2-token prediction increases the importance of I(X;Y)I(X;Y) by a factor of 22. So, multi-token predictors are more accurate at predicting tokens XX that are of relevance for the remainder of the text to come. In Appendix L.2, we give a relative version of the above equations that shows the increased weight of relative mutual information in a loss decomposition of 2-token prediction loss.

Related work

Dong et al. (2019) and Tay et al. (2022) train on a mixture of denoising tasks with different attention masks (full, causal and prefix attention) to bridge the performance gap with next token pretraining on generative tasks. Tay et al. (2022) uses the span corruption objective, which replaces spans of tokens with special tokens for the encoder and the decoder then predicts the contents of those spans. Unlike UniLM, this allows full causal training with teacher forcing. Similarly, Yang et al. (2019) train on permuted sequences, while conserving the original positional embeddings, effectively training the model to predict various parts of the sequence given a mix of past and future information. This permuted language modeling is the closest task to ours since it allows predicting beyond the next token. However all of these language modeling tasks train on a small percentage of the input text: on average only 15% of the tokens are backwarded through. For Dong et al. (2019), where the masking is done in BERT style, it is hard to mask more than 15% since it destroys too much information. For Tay et al. (2022), it is technically possible to have a larger proportion but in practice, the settings used have between 15% and 25% of masked tokens. (Yang et al., 2019) also makes it possible to train on the whole sequence since it is only permuted, and no information is lost. Yet, in practice, since the completely random permutation is very hard to reconstruct, only 15% are predicted for training stability reasons.

Multi-token prediction in language modelling

Qi et al. (2020) argue that multi-token prediction encourages planning, improves representations and prevents the overfitting on local patterns that can result from teacher-forced training. However, their technical approach replicates the residual stream nn-fold while ours allows for compute-matched comparisons and makes the residual representations participate more directly in the auxiliary loss terms. Stern et al. (2018) and Cai et al. (2024) propose model finetunings with multi-token prediction for faster inference but do not study the effects of such a loss during pretraining. Pal et al. (2023) use probing methods to show that next-token prediction models are able to predict additional consecutive tokens to a certain extent, but less so than our models which are specifically trained for this task. Jianyu Zhang (2024) observe improvements in language modelling tasks with multi-label binary classification over the occurrence of vocabulary words in the future as an auxiliary learning task.

Self-speculative decoding

Stern et al. (2018) are, to the best of our knowledge, the first to suggest a speculative decoding scheme for faster inference. Our architecture replaces their linear prediction heads by transformer layers, but is otherwise similar. By reorganizing the order of the forward/backward, we can use all loss terms instead of stochastically picking one head for loss computation. Cai et al. (2024) present a more elaborate self-speculative decoding scheme that uses the top-kk predictions of each head instead of the best one only. It can be used with the multi-token prediction models we train.

Multi-target prediction

Multi-task learning is the paradigm of training neural networks jointly on several tasks to improve performance on the tasks of interest (Caruana, 1997). Learning with such auxiliary tasks allows models to exploit dependencies between target variables and can even be preferable in the case of independent targets (Waegeman et al., 2019). While more specifically tailored architectures for multi-target prediction are conceivable (Spyromitros-Xioufis et al., 2016; Read et al., 2021), modern deep learning approaches usually rely on large shared model trunks with separate prediction heads for the respective tasks (Caruana, 1997; Silver et al., 2016; Lample et al., 2022) like we do. Multi-target prediction has been shown to be a successful strategy in various domains, e.g. for learning time series prediction with more distant time steps in the future as auxiliary targets (Vapnik and Vashist, 2009) or for learning from videos with several future frames (Mathieu et al., 2016; Srivastava et al., 2016) or representations of future frames (Vondrick et al., 2016) as auxiliary targets.

Conclusion

We have proposed multi-token prediction as an improvement over next-token prediction in training language models for generative or reasoning tasks. Our experiments (up to 7B parameters and 1T tokens) show that this is increasingly useful for larger models and in particular show strong improvements for code tasks. We posit that our method reduces distribution mismatch between teacher-forced training and autoregressive generation. When used with speculative decoding, exact inference gets 3 times faster.

In future work we would like to better understand how to automatically choose nn in multi-token prediction losses. One possibility to do so is to use loss scales and loss balancing (Défossez et al., 2022). Also, optimal vocabulary sizes for multi-token prediction are likely different from those for next-token prediction, and tuning them could lead to better results, as well as improved trade-offs between compressed sequence length and compute-per-byte expenses. Finally, we would like to develop improved auxiliary prediction losses that operate in embedding spaces (LeCun, 2022).

Impact statement

The goal of this paper is to make language models more compute and data efficient. While this may in principle reduce the ecological impact of training LLMs, we shall be careful about rebound effects. All societal advantages, as well as risks, of LLMs should be considered while using this work.

Environmental impact

In aggregate, training all models reported in the paper required around 500K GPU hours of computation on hardware of type A100-80GB and H100. Estimated total emissions were around 50 tCO2eq, 100% of which were offset by Meta’s sustainability program.

Acknowledgements

We thank Jianyu Zhang, Léon Bottou, Emmanuel Dupoux, Pierre-Emmanuel Mazaré, Yann LeCun, Quentin Garrido, Megi Dervishi, Mathurin Videau and Timothée Darcet and other FAIR PhD students and CodeGen team members for helpful discussions. We thank Jonas Gehring for his technical expertise and the original Llama team and xFormers team for enabling this kind of research.

References

Appendix A Additional results on self-speculative decoding

Appendix B Alternative architectures

The architecture described in Section 2 is not the only sensible option, but proved technically viable and well-performing in our experiments. We describe and compare alternative architectures in this section.

Replicating the unembedding matrix nn times is a simple method for implementing multi-token prediction architectures. However, it requires matrices with shapes (d,nV)(d,nV) in the notation of Section 2, which is prohibitive for large-scale trainings.

Linear heads

Apart from using a single transformer layer for the heads HiH_{i}, other architectures are conceivable. We experimented with a single linear layer without any nonlinearity as heads, amounting to linear probing of the model’s residual representation zz. Architectures with more than one layer per head are also possible, but we did not pursue this direction further.

Causal and anticausal variant

Instead of making the prediction heads Pi(xt+i ∣ zt:1)P_{i}(x_{t+i}\,|\,z_{t:1}) architecturally independent of each other, we can also allow them to rely on other heads’ (pre-unembedding) outputs. In a causal variant, later prediction heads are applied on top of the previous ones, i.e. the ii-th prediction head PiP_{i} is given by

In another anticausal variant, the network starts by predicting the most distant tokens before gradually refining up to the following token:

These architectures likewise allow a sequential forward/backward order as the parallel architecture from Section 2. This is described in Figure S11.

Appendix C Training speeds

Appendix D Finetuning

Appendix E Additional results on model scaling behavior

Appendix F Details on CodeContests finetuning

Appendix G Additional results on natural language benchmarks

We evaluate the models from Section 3.7 on standard natural language processing benchmarks: ARC Challenge [Yadav et al., 2019], COPA [Roemmele et al., 2011], Hellaswag [Zellers et al., 2019], Natural Questions [Kwiatkowski et al., 2019], PIQA [Bisk et al., 2019], SIQA [Sap et al., 2019] and TriviaQA [Joshi et al., 2017].

Appendix H Additional results on abstractive text summarization

In this section, we report comprehensive evaluation results on summarization tasks for the 7B parameter models trained on 200B and 500B tokens of natural language from Section 3.7.

Appendix I Additional results on mathematical reasoning in natural language

Appendix J Additional results on induction learning

Appendix K Additional results on algorithmic reasoning

We investigate the following computation-sharing hypothesis for explaining the efficacy of multi-token prediction as training loss.

The prediction difficulty of different tokens in natural text varies greatly. Some tokens may be the continuations of partial words that are uniquely determined from their preceding context without any effort, while others may require to predict theorem names in difficult mathematical proofs or the correct answer to an exam question. Language models with residual connections have been shown to refine their output token distribution with each successive layer, and can be trained with early exit strategies that spend variable amounts of computational resources per token position. Multi-token prediction losses explicitly encourage information-sharing between adjacent token positions and can thus be viewed as a method to learn allocating computational resources in language models more efficiently to the tokens that benefit most of it.

To check the truth of this hypothesis, we augment the polynomial arithmetic task from Section 4.2 with a varying number of pause tokens [Goyal et al., 2023] inserted between the question and a token that denotes the beginning of the answer. Pause tokens introduce additional computational resources that can be expended for computations that are expected to be useful later on in the sequence, in other words: to start thinking about the answer. According to the computation-sharing hypothesis, multi-token prediction models learn information-sharing and thus computation-sharing between token positions more easily, and may be better at making use of these additional computational resources than next-token prediction models are. In Figure S15, we show the evaluation results on the polynomial arithmetic task with a fixed number of pause tokens inserted both at training and evaluation time. Multi-token prediction models likewise outperform next-token prediction models on these task variants across task difficulties and model sizes. However, we do not see strong evidence of a widening or shrinking of this gap i.e. we cannot conclude from these experiments on the veracity of the computation-sharing hypothesis.

In Table S11, we report results from another experiment in the same spirit: by adding spaces and newlines to HumanEval and MBPP prompts, we add “pause tokens” in a somewhat natural way. According to these results, multi-token prediction models have a slight advantage at using this additionally provided compute, but the effect is marginal.

Appendix L Additional intuitions on multi-token prediction

In Section 5.2, we argued that multi-token prediction reduces the distribution mismatch between teacher-forced training and autoregressive evaluation of language models. Scheduled sampling [Bengio et al., 2015] is a curriculum learning method that likewise aims to bridge this gap in sequence prediction tasks by gradually replacing more and more input tokens with model-generated ones.

While effective in areas such as time series forecasting, scheduled sampling is, in our opinion, inapplicable to language modelling due to the discrete nature of text. Replacing ground truth input sequences by interleavings of ground truth and model-generated tokens frequently results in ungrammatical, factually wrong or otherwise incoherent text, which should be avoided at all cost. Moreover, unlike multi-token prediction, the technique originally developed for recurrent neural networks cannot easily be adapted for parallel training setups like the ones of transformer models.

L.2 Information-theoretic argument

We give details on the information-theoretic terms appearing in the decomposition in Section 5.2 and derive a relative version that similarly allows to decompose multi-token prediction losses. As in Section 5.2, denote by XX the next token and by YY the second-next one, and omit conditioning on the preceding context CC for ease of notation. In Section 5.2, we decomposed H(X)+H(Y)H(X)+H(Y)—the quantity of interest for 2-token prediction models—as follows:

Let us explain each of the terms. The entropy terms denote the uncertainty contained in the ground-truth random variables XX and YY. In particular, they do not refer to model predictions. The term H(Y∣X)H(Y\mid X) is a classical next-token entropy for the prefix (C,X)(C,X). The conditional entropy H(X∣Y)H(X\mid Y) is a more theoretical entity not modelled by causal models. It describes the uncertainty about XX given the prefix CC and suffix YY, and therefore captures the local variations of XX that do not affect the continuation of the text YY. The mutual information I(X;Y)I(X;Y) on the other hand describes the information about YY contained in XX (and vice versa) and therefore captures the variations of XX which constrain the continuation of the text.

However, the argument given in Section 5.2 relies on the assumption that multi-token prediction losses obey a similar decomposition as the sum of the ground-truth entropies themselves. Let us make this rigorous. Denote by p(x,y)p(x,y) the joint distribution of XX and YY, by p(x)p(x) (short for pX(x)p_{X}(x)) the marginal distribution of XX and by p(y)p(y) the one of YY. Denote the densities of the model’s predictions by q(x,y)q(x,y), q(x)q(x) and q(y)q(y), respectively, conditional distributions by p(x∣y)p(x\mid y) and Kullback-Leibler divergence from qq to pp by D(p  ∥  q)D(p\;\|\;q) and cross-entropy from qq to pp by H(p,q)H(p,q).

The conditional cross-entropy H(pX∣Y,qX∣Y)H(p_{X\mid Y},q_{X\mid Y}) of XX conditioned on YY from qq to pp is defined as the expectation under yy of the cross-entropy between the distributions pXp_{X} and qXq_{X} conditioned on yy, in formulas:

The relative mutual information Ip∥q(X;Y)I_{p\|q}(X;Y) of XX and YY from qq relative to pp is defined by

We have Ip∥q(X;Y)=H(pX,qX)+H(pY,qY)−H(p,q)I_{p\|q}(X;Y)=H(p_{X},q_{X})+H(p_{Y},q_{Y})-H(p,q), Ip∥p(X;Y)=Ip(X;Y)I_{p\|p}(X;Y)=I_{p}(X;Y) reduces to standard mutual information under the distribution pp and Ip∥q(X;Y)I_{p\|q}(X;Y) is symmetric in XX and YY but can be negative.

We have the following relative version of the decomposition H(X)=H(X∣Y)+I(X;Y)H(X)=H(X\mid Y)+I(X;Y).

H(pX,qX)=H(pX∣Y,qX∣Y)+Ip∥q(X;Y).H(p_{X},q_{X})=H(p_{X\mid Y},q_{X\mid Y})+I_{p\|q}(X;Y).

Symmetrizing, we get the desired relative version of H(X)+H(Y)=H(X∣Y)+2I(X;Y)+H(Y∣X)H(X)+H(Y)=H(X\mid Y)+2I(X;Y)+H(Y\mid X):

Setting pp to be the empirical distribution of the training data, the left-hand side describes the cross-entropy loss used to train 2-token prediction models. The right-hand side gives the decomposition into a local cross-entropy term, a mutual information term with weight two and a shifted next-token cross-entropy term. We interpret this as follows: by adding the term H(pY,qY)H(p_{Y},q_{Y}) to the loss, 2-token prediction incentivizes models to precompute features which will become useful for predicting YY in the next step and increases the weight of the relative mutual information term in the loss. What does relative mutual information actually mean? By interpreting Kullback-Leibler divergence D(p  ∥  q)D(p\;\|\;q) as the average number of bits needed in addition to send data from pp with a code optimized for qq instead of pp, we see that minimizing

means minimizing the average number of additional bits needed to send data from pp with a code optimized for qq that treats XX and YY as independent compared to one that does not. If this number is small, qq managed to exploit the mutual information of XX and YY under pp.

L.3 Lookahead reinforces choice points

Training with multi-head prediction increases the importance of choice points in the loss in comparison to inconsequential decisions. To make this argument, we present a simplified model of language modelling. Consider a sequential decision task and a model MM that is trained in a teacher-forced way on optimal trajectories. We distinguish choice points –transitions that lead to different outcomes – and inconsequential decisions which do not (Figure S17 (a) and (b)).

More formally, assume that the language model is deployed in a reinforcement learning setting like in reinforcement learning from human feedback [Ouyang et al., 2022] (states are prompts followed by the partial sequence of tokens xt:1x_{t:1} generated so far, actions are single tokens xt+1x_{t+1} to generate, rewards are external R(xt:1)R(x_{t:1})). The quantity

is the value of the state xt:1x_{t:1} following the policy π\pi, while

quantifies the importance of the decision xt+1x_{t+1} on the value thereafter. Choice points can formally be viewed as steps tt for which σπ(xt:1)\sigma_{\pi}(x_{t:1}) is large, while inconsequential points are steps where it is low. Note that for completion models, there is no explicit reward, and our argument is merely meant to illustrate what we mean by choice points.

Derailing denotes a situation where autoregressive generation of trajectories from MM at inference time results in bad outcomes after MM made a mistake on a choice point. Even if subsequently, MM acts optimally given this choice, the final outcome can be significantly worse than the outcome of the optimal trajectory.

As argued in Section 5.1, we believe that this model captures important features of training and inference with language models: choice points are semantically important turning points in the generated texts, such as the final answer to a question or a specific line of code, while inconsequential decisions can be a choice among synonyms or of variable names in code.

L.4 Factorization orders

Causal language modelling factorizes probabilities over text sequences xt⋯x1x_{t}\cdots x_{1} classically as

While moving forward in time is certainly the most natural choice of factorization order, there exist cases where it is suboptimal. In inflectional languages, for instance, agreement between related sentence parts is a frequent pattern with one word directing the grammatical forms of others. Consider the German sentence

Wie konnten auch Worte meiner durstenden Seele genügen?roughly: How could words be enough for my thirsty soul?

Friedrich Hölderlin, Fragment von Hyperion (1793)

where "genügen" requires a dative case object and then "Seele" requires the possessive pronoun "mein" to be in female singular dative form "meiner" and the participle "durstend" to be in female singular dative form in weak declination "durstenden" because it follows "meiner". In other words, the factorization order

Wie konnten auch Worte →\rightarrow genügen →\rightarrow Seele →\rightarrow meiner →\rightarrow durstenden?

is arguably an easier one for constructing the above sentence. Humans as well as language models therefore have to perform this factorization (which deviates from the causal order in which predictions take place!) within their latent activations, and a 44-token prediction loss makes this easier as it explicitly encourages models to have all information about the successive 4 tokens in its latent representations.

Appendix M Training hyperparameters