Break the Sequential Dependency of LLM Inference Using Lookahead Decoding

Yichao Fu, Peter Bailis, Ion Stoica, Hao Zhang

Introduction

Large language models (LLMs) are transforming the AI industry. As they are increasingly integrated into diverse applications such as search (Team et al., 2023) and chatbots (Ouyang et al., 2022), generating long sequences at low-latency using LLMs is becoming one significant requirement. However, current LLMs generate text based on (Touvron et al., 2023a, b; Jiang et al., 2023; OpenAI, 2023) autoregressive decoding, which falls short in efficiency, primarily for two reasons. First, autoregressive decoding generates only one token at a time. Hence, the overall generation time is proportional to the number of decoding steps. Second, each decoding step largely underutilizes the parallel processing capabilities of modern accelerators (e.g., GPUs). Given the pressing need for low latency in various applications, improving autoregressive decoding remains a central challenge.

Several approaches have been proposed – one such approach is speculative decoding (Chen et al., 2023; Leviathan et al., 2023) and its variants (He et al., 2023; Stern et al., 2018; Cai et al., 2024; Li et al., 2023; Liu et al., 2023; Miao et al., 2023). These methods all follow a guess-and-verify approach: they use a draft model to speculate several subsequent tokens and then use the original (base) LLM to verify these tokens in parallel. Since the draft model requires much fewer resources and the cost of verifying multiple tokens in parallel is similar to the cost of generating a single token, these methods can achieve considerable speedups. However, their speedups are bounded by the token acceptance rate (§4.1), i.e., the fraction of tokens generated by the draft model that passes the verification test of the base model. This is because every token that fails verification needs to be regenerated by the base model. In the worst case, if most proposed tokens fail verification, these methods may slow down the decoding process. Therefore, achieving a high acceptance rate is essential for these methods. Unfortunately, training a draft model to achieve a high acceptance rate is non-trivial, and the trained draft model does not generalize across base models and datasets.

To address these problems, this paper develops Lookahead Decoding. We build upon a key observation: autoregressive decoding can be equivalently formulated as solving a non-linear system via the fixed point Jacobi iteration method (§2), which we term as Jacobi decoding (Santilli et al., 2023). Each Jacobi decoding step can generate multiple tokens in parallel at different positions. Although these tokens may appear at incorrect positions, we can leverage this parallel generation approach to have the LLM generate several disjoint n-grams in parallel in a single step. These n-grams could potentially be integrated into future parts of the generated sequence, pending verification by the base model to maintain the output distribution.

Lookahead Decoding takes advantage of the particular characteristics of autoregressive decoding, which is bounded by the memory bandwidth–as each generated token depends on all tokens before it–rather than compute, by using the available cycles to generate and verify nn-grams (subsequent tokens) at virtually no additional cost. In a nutshell, Lookahead Decoding consists of a lookahead branch that generates nn-grams and a verification branch that verifies nn-grams, both executing in a single step. To improve efficiency, we use an nn-gram pool to cache the historical nn-grams generated so far. This way, Lookahead Decoding can significantly reduce the latency of LLM inference just by exploiting the compute resources that autoregressive decoding would leave unused. More importantly, Lookahead Decoding scales with the compute – we show that it can linearly reduce the number of decoding steps relative to the log(FLOPs) allocated per step.

We have implemented the algorithm in both Python and CUDA, compatible with memory-efficient attention algorithms (e.g., FlashAttention (Dao, 2023)), and supports various sampling methods without changing the output distribution. We also scale it to multiple GPUs, resulting in Lookahead Parallelism. We evaluate Lookahead Decoding on the popular LLaMA-2 (Touvron et al., 2023b) models. It achieves 1.8x speedup on the challenging multi-turn chat dataset MT-Bench (Zheng et al., 2023) and up to 4x speedup in code completion tasks with Lookahead Parallelism on 8 GPUs. Lookahead Decoding showed significant potential in lowering the latency for latency-sensitive tasks. Our contributions are summarized as follows.

We design Lookahead Decoding, a new lossless, parallel decoding algorithm to accelerate LLM inference without needing any auxiliary component.

We reveal Lookahead Decoding’s scaling behavior: it linearly reduces the number of decoding steps according to per-step log⁡(\log(FLOPs)). This enables trade-offs between the number of decoding steps and per-step FLOPs, making it future-proof.

We show it benefits from the latest memory-efficient attentions and is easily parallelizable by developing its distributed CUDA implementations.

We evaluated Lookahead Decoding and demonstrate its effectiveness under different settings.

Background

In this section, we formulate both autoregressive and Jacobi decoding from the lens of solving nonlinear systems.

Causal Attention in Decoder Models. Most contemporary LLMs are composed of two core components: token-wise modules (including MLP and normalization (Ba et al., 2016; Zhang & Sennrich, 2019)) and attention (Vaswani et al., 2023) modules. Tokens interact with each other in the attention modules, while in other token-wise modules, they are processed without exchanging information with each other.

The attention layer encompasses three input elements: query Q\mathbf{Q}, key K\mathbf{K}, and value V\mathbf{V}, with the ii-th token in each denoted as Qi\mathbf{Q}_{i}, Ki\mathbf{K}_{i}, and Vi\mathbf{V}_{i}, respectively. The attention layer executes the following operation: O=softmax(QKT)V\mathbf{O}=\textrm{softmax}\left(\mathbf{Q}\mathbf{K}^{T}\right)\mathbf{V}. A lower triangular mask applied to QKT\mathbf{Q}\mathbf{K}^{T} in causal attentions (specific to decoder models) ensures that Oi\mathbf{O}_{i} is calculated only from Qi\mathbf{Q}_{i} and Kj\mathbf{K}_{j}, Vj\mathbf{V}_{j} where j≤ij\leq i. Because all other layers in the LLM perform token-wise operations, for any given model input x\mathbf{x} and output o\mathbf{o}, oi\mathbf{o}_{i} (ii-th token in o\mathbf{o}) is exclusively influenced by xj\mathbf{x}_{j} (jj-th token in x\mathbf{x}) where j≤ij\leq i.

We define x0\mathbf{x}^{0} as the prompt tokens given by the user. The LLM needs to generate an output sequence (of length mm) from x0\mathbf{x}^{0}. Denote yiy_{i} as the token generated at step ii. The autoregressive decoding process of mm tokens can be seen as solving the following mm problems one by one (assume greedy sampling):

Guess-And-Verify Paradigm. The Guess-And-Verify decoding paradigm speculates multiple potential future tokens and subsequently confirms the correctness of these speculations within a single decoding step. Take speculative decoding with greedy sampling as an example: at step tt, with the prompt x0\mathbf{x}^{0} and tokens y1:t−1\mathbf{y}_{1:t-1} generated so far, we can use a draft model to autoregressively generate a draft sequence yt:t+n−1\mathbf{y}_{t:t+n-1} of length nn. Because yt:t+n−1\mathbf{y}_{t:t+n-1} is known a priori, we then use the LLM to solve Eqs 2 in parallel, obtaining yt:t+n′\mathbf{y}^{\prime}_{t:t+n}. Then, we verify if yt+iy_{t+i} is equal to yt+i′y_{t+i}^{\prime} for each ii from i=0i=0 to i=n−1i=n-1. If there is a match, we accept this token and proceed; otherwise, we stop checking and drop subsequent tokens. Finally, we update y\mathbf{y} with all accepted tokens.

As stated in §1, these approaches depend on a good draft model, which is hard to obtain and cannot generalize.

Jacobi Decoding. By notating f(yi,y1:i−1,x0)=yi−arg⁡ ⁣max⁡PM(yi∣y1:i−1,x0)f(y_{i},\mathbf{y}_{1:i-1},\mathbf{x}^{0})=y_{i}-\arg\!\max P_{M}(y_{i}|\mathbf{y}_{1:i-1},\mathbf{x}^{0}), we can transform Eqs 1 into the following non-linear system of equations (Song et al., 2021; Santilli et al., 2023):

We can solve this non-linear system using Jacobi iteration by iteratively updating all yiy_{i} from a random initial guess y0\mathbf{y}^{0}, along the trajectory y1,...,yt,...\mathbf{y}^{1},...,\mathbf{y}^{t},..., until converging to the fixed point solution ym\mathbf{y}^{m}. We detail this algorithm, termed as Jacobi decoding, in Appendix Algorithm 1. This process guarantees to return the solution of all mm variables yiy_{i} in at most mm iterations, as the very first token of each Jacobi update matches autoregressive decoding. Sometimes, more than one token might be correctly generated in a single iteration, potentially reducing the number of decoding steps. It is worth noting that, as yt\mathbf{y}^{t} is generated based on the past value yt−1\mathbf{y}^{t-1} on the trajectory, any two adjacent tokens from yt−1\mathbf{y}^{t-1} and yt\mathbf{y}^{t} can form a meaningful 2-gram.

Limitations of Jacobi Decoding. Empirically, we observe Jacobi decoding can hardly reduce decoding steps, even if it can generate multiple tokens per step. This is because the generated tokens are often put in the wrong positions of the sequence, and correctly placed tokens are frequently replaced by subsequent Jacobi iterations. These prevent it from achieving wall-clock speedup.

Lookahead Decoding

Lookahead Decoding leverages Jacobi decoding’s ability to generate many tokens in one step but addresses its limitation. Fig. 1 illustrates its workflow. The key design in Lookahead Decoding is to keep track of the trajectory of Jacobi decoding and generate nn-gram from this trajectory. This is achieved by maintaining a fixed-sized 2D window, with the two dimensions corresponding to the sequence and the time axis, respectively, to generate multiple disjoint nn-grams from the Jacobi iteration trajectory in parallel. We call this process the lookahead branch. In addition, Lookahead Decoding introduces an nn-gram pool to cache these nn-grams generated along the trajectory. Promising nn-gram candidates are verified later by a designed verification branch to preserve the LLM’s output distribution; if passing verification, those disjoint n-grams are integrated into the sequence. The detailed algorithm is shown in Algorithm 2 in Appendix.

Lookahead Decoding uses a fixed-sized 2D window for efficient nn-gram generation. In contrast to the original Jacobi decoding, which only uses the history tokens from the last step (or equivalently, it generates 2-grams), Lookahead Decoding generates many nn-grams, with n≥2n\geq 2, in parallel by using the n−1n-1 past steps’ history tokens, effectively leveraging more information from the trajectory. The fixed-sized 2D window in the lookahead branch is characterized by two parameters: (1) WW defines the lookahead size into future token positions to conduct parallel decoding; (2) NN defines the lookback steps into the past Jacobi trajectory to retrieve nn-grams. See Algorithm 2 for a detailed process.

An example of the lookahead branch with W=5W=5 and N=4N=4 is in Fig. 2 (b), in which we look back N−1=3N-1=3 steps and look ahead 55 tokens for each step. The blue token with the digit 0 is the current step’s (tt) input, and the orange, green, and red tokens were generated in previous lookahead branches at steps t−3t-3, t−2t-2, and t−1t-1, respectively. The digit on each token shows its relative position to the current input (i.e., the blue one labeled as 0). In the present stage, we perform a modified Jacobi iteration to generate new tokens for all 5 positions, following the trajectory formed by the preceding 3 steps. Once generated, we collect and cache them in the nn-gram pool (n=4n=4) – for instance, a 4-gram consists of the orange token at position 1, the green token at position 2, the red token at position 3, and a newly generated token.

The most outdated tokens in both dimensions (time and sequence) will be removed, and newly generated tokens will be appended to the lookahead branch to maintain a fixed window size for each step. For example, we will remove all orange and green tokens with position 1 in Fig. 2. We then form a new lookahead branch with green tokens with indices 2, 3, 4, 5, all red tokens, and all newly generated tokens for the next step.

2 Verification Branch

Lookahead Decoding preserves the output distribution via its verification branch. We first discuss how to verify in greedy sampling. Recall in speculative decoding: the verification is performed by sending the draft tokens to the LLM to get an output for each draft token, then progressively checking if the last token’s corresponding output, generated by the target LLM, exactly matches the draft token itself (§2). The verification branch in Lookahead Decoding resembles this process, despite verifying many draft nn-gram candidates in parallel. In particular, We first look up from the nn-gram pool for “promising” nn-grams – by checking if a nn-gram starts with a token that exactly matches the last token of the current ongoing sequence. We then use the LLM to verify all these nn-grams in parallel, following a similar fashion as in speculative decoding. See Algorithm 3 in the Appendix for the detailed procedures.

We next discuss how to support more advanced sampling. Previous research (Miao et al., 2023) has developed efficient tree-based verification for speculative decoding with sampling support, where multiple draft sequences derived from a token tree can be verified in parallel. However, it does not apply to Lookahead Decoding as our verification works on disjoint nn-grams instead of trees. We improve it by progressively verifying along the nn-gram length and removing nn-grams with mismatched prefixes. Besides, speculative decoding style verification requires the probability distribution where the draft token is sampled to update the probability distribution when the draft token is rejected. Because we store all nn-grams in a pool instead of discarding them each step, we would need huge memory to store the probability distributions (each of vocabulary size) for the entire nn-gram pool. The key to overcome this is to leverage the mechanism that the verification is indifferent to how draft tokens were sampled – different sampling methods (e.g., greedy sampling) only influence the acceptance rate but keep the output distribution. We can force greedy sampling at the nn-gram generation (lookahead branch), in which the probability distribution degenerates into a one-hot vector. Hence we only need to store which token is selected. We elaborate the approach in Algorithm 4, prove its correctness in Appendix B, and verify its quality and speedups in §5.3.

It is expected to have an increasingly large nn-gram cache hence a growing verification branch as decoding progresses. We set a cap of GG to limit the maximum number of promising candidates run in parallel in the verification branch to manage the verification cost. Empirically we suggest to set GG proportional to WW to balance generation and verification. In practice, we simply set G=WG=W.

3 Decode, Predict, and Verify in The Same Step

At execution, the lookahead and verification branches can be integrated into one decoding step to leverage parallel processing. This requires a designated attention mask, as shown in Fig. 2 (b). This attention mask is straightforwardly derived following the principle that each token is only visible to the tokens with a larger position index than itself (§2). For example, only the green token at position 5 and all orange tokens are visible to the red token 6. The tokens in the lookahead branch are not visible to the tokens in the verification branch, and vice versa.

Integration with FlashAttention. FlashAttention (Dao et al., 2022; Dao, 2023) can vastly accelerate the training and inference of LLMs by saving memory I/O on the slow memory hierarchy. It forces a causal mask (e.g., Fig. 2 (a)) to avoid all token interactions outside a lower triangular scope, which is not suitable for Lookahead Decoding as we take a more subtle attention mask (e.g., Fig. 2 (b)) for different WW, NN, and GG. To solve this, we hardcode Lookahead Decoding’s attention pattern with adjustable WW, NN, and GG in FlashAttention. Applying FlashAttention to Lookahead Decoding brings about 20% end-to-end speedup compared to a straightforward implementation on top of native PyTorch in our experiments (§5.2).

4 Lookahead Parallelism

Lookahead Decoding is easy to parallelize on multiple GPUs for both lookahead and verification branches. Parallelizing the lookahead branch is achieved by noting that the lookahead computation is composed of several disjoint branches. For example, the branch with green 1 and red 2 tokens does not have interaction with the branch with the tokens green 3 and red 4 in Fig. 2 (b). We can put these disjoint branches onto different GPUs without introducing communication during the inference computation. Parallelizing the verification branch is done by assigning multiple nn-gram candidates to different devices. Because the verification of each candidate, by design, is independent of others, this will not cause communication.

Fig. 3 shows an example of parallelizing the lookahead branch and verification branch in Fig. 2 (b) to four GPUs. This workload allocation will have the orange token 0,1,2,3 and the input token 0 be redundantly placed and computed. However, it can essentially save communication volume during the whole forward pass. We only need to synchronize the generated tokens on each device after the forward pass. We can further scale the WW, NN, and GG with multiple GPUs’ increased FLOPs to obtain a lower latency according to Lookahead Decoding’s scalability (§4).

We name this new parallelism as lookahead parallelism (LP). Unlike previous parallelism methods (including pipeline and tensor parallelisms) that shard the model parameters or states across different GPUs, LP maintains an entire copy of the model for each GPU (thus needing more memory) and allows distributing tokens to different GPUs. Hence, LP is advantageous in inference as it introduces near-zero communication per step while existing model parallelism methods (Narayanan et al., 2021; Shoeybi et al., 2019) involve a large communication overhead on the critical path of each decoding step.

Scaling Law of Lookahead Decoding

Since Lookahead Decoding introduces flexible parameters WW and NN associated with the cost of each parallel decoding step. This section investigates the scaling law between compute FLOPs and the theoretical speedup, and compares it to speculative decoding.

Speculative decoding uses the draft model to speculate one token sequence at each step. We represent the probability of each token in the sequence passing the verification of the LLM by β\beta (acceptance rate) and notate its expectation E(β)=αE(\beta)=\alpha. If we use the draft model to guess γ\gamma tokens per step, the expectation of the number of accepted tokens is denoted as (Leviathan et al., 2023):

Instead of speculating one sequence every time, we would speculate bb sequences. We assume that bb sequences, each of γ\gamma tokens, are sampled as each token will have the same acceptance rate of β\beta. Under this setting, the expectation of the number of accepted tokens is denoted as follows:

See derivations in Appendix C for Eq. 4 and Eq. 5. Note that when b=1b=1, Eq. 5 falls back to Eq. 4.

2 Estimating Speedup for Lookahead Decoding

We define the S=step compression ratio\mathcal{S}=\textit{step compression ratio} as the number of autoregressive steps divided by the number of Lookahead Decoding steps to generate the same length of the sequence. As the number of generated tokens equals the autoregressive steps, it can be denoted as:

Lookahead Decoding speculates bb sequences every time as in Eq. 5. In each step, we will search nn-grams in the pool starting with the current input token and have at most GG speculations of length N−1N-1. As we set G=WG=W (§3.2), we have G=W=bG=W=b and N−1=γN-1=\gamma using the notations in Eq. 5. In practice, we cannot expect each step to have equally good speculations (i.e., acceptance rate with E(β)=αE(\beta)=\alpha). We assume that, on average, for every ff step, we have one good speculation with E(#tokens)E(\#tokens) tokens accepted, and for the other f−1f-1 steps, we fall back to autoregressive decoding due to bad speculations. We use this ff to bridge S\mathcal{S} and E(#tokens)E(\#tokens) per step as follows:

We can plot the curve indicated by our formulation with one specific setting as in Fig. 4 (b). We find that the trend of our empirical experiments (LLaMA-2-Chat-7B on MT-Bench with G=WG=W as in Fig. 4 (a)) align well with the formulation to some extent. From this formulation, we conclude that we can linearly reduce the number of decoding steps according to per-step log⁡(b)\log(b) given a large enough γ\gamma. In contrast to speculative decoding, Lookahead Decoding will not meet an upper bound indicated in Eq. 4 by simultaneously increasing γ\gamma and bb. This reveals the scaling law of Lookahead Decoding to linearly reduce decoding steps according to per-step log⁡(\log(FLOPs)) given a large enough NN, since per-step FLOPs is roughly proportional to the number of input tokens (i.e., (W+G)∗(N−1)(W+G)*(N-1)). The scaling law also suggests Lookahead Decoding’s strong scaling to multiple GPUs, in which we can obtain an even greater per-token latency reduction by using more FLOPs, which is advantageous for latency-sensitive tasks.

Evaluation Results

Model and testbed. We used various versions of the LLaMA-2 (Touvron et al., 2023b) and CodeLlama (Roziere et al., 2023) models, including the 7B, 13B, 34B, and 70B sizes, on two GPU setups S1 and S2. S1 is equipped with NVIDIA A100 GPUs with 80GB of memory. On S1, the 7B, 13B, and 34B models are deployed on a single A100, while the 70B model utilizes 2 A100s with pipeline parallelism supported by Accelerate (Gugger et al., 2022). S2 is a DGX machine with 8 NVIDIA A100 GPUs with 40GB memory and NVLink. All models serve with FP16 precision and batch of 1 if not specified (Cai et al., 2024; He et al., 2023).

Datasets. We benchmarked Lookahead Decoding’s performance across a broad spectrum of datasets and tasks. MT-Bench (Zheng et al., 2023) is a diverse set of multi-turn questions with many unique tokens. GSM8K (Cobbe et al., 2021) contains a set of math questions, in which we use the first 1k questions. HumanEval (Chen et al., 2021) covers both code completion and infilling tasks. We also test on MBPP (Austin et al., 2021) dataset for instruction-based code generation, and on ClassEval (Du et al., 2023) for class-level code completion. To control generation length in code generation tasks, we set the maximum sequence length to 512 and 2,048 on HumanEval and ClassEval, respectively, aligned with prior setups (Ben Allal et al., 2022; Du et al., 2023). Tab. 1 lists detailed settings. In addition, we validate the effectiveness of sampling (§3.2) on XSum (Narayan et al., 2018) and CNN/Daily Mail (See et al., 2017) datasets.

Baseline Settings. Our primary baseline is HuggingFace’s implementation of greedy search (Wolf et al., 2020). Additionally, we employ FlashAttention (Dao et al., 2022; Dao, 2023) as a stronger baseline to assess the performance of FlashAttention empowered Lookahead Decoding. In distributed settings, we evaluate LP against TP (supported by deepspeed (Aminabadi et al., 2022)) and PP (supported by accelerate (Gugger et al., 2022)). We measure the throughput of single batch inference against these baseline settings (Cai et al., 2024; He et al., 2023).

Fig. 5 shows the end-to-end performance of Lookahead Decoding when compared with HuggingFace’s implementation of greedy search on S1. The used tasks and models are shown in Tab. 1. Across various datasets, Lookahead Decoding demonstrates a 1.5x-2.3x speedup. Generally, our method exhibits better performance in code completion tasks (e.g., 2.3x), given the higher occurrence of repetitive tokens during code completions, making predictions easier. Besides, smaller models also exhibit a higher speedup when compared to larger models. This is because Lookahead Decoding trades per-step FLOPs with a step compression ratio (§4). A larger model requires more FLOPs and quickly hits the GPU FLOPs cap compared to a smaller model. So, it shows a lower ability to compress decoding steps given the same GPU setting.

2 Performance with LP and FlashAttention

We evaluated the performance of Lookahead Decoding with LP and FlashAttention augmentation on S2 with greedy search. The used tasks and models are shown in Tab. 1. The results for the 7B and 13B models are in Fig. 6 and Fig. 7, respectively. FlashAttention speeds up the PyTorch implementation of Lookahead Decoding by 20%. Notably, FlashAttention-integrated Lookahead Decoding shows 1.8x speedups for the 7B model on MT-Bench compared with autoregressive decoding with FlashAttention (i.e., 1.9x vs 1.07x in Fig. 6). We did a strong scaling of the workloads to multiple GPUs for distributed settings (i.e., increasing GPUs but not increasing workloads). The multiple GPU settings of both TP (w/ DeepSpeed) and PP (w/ Accelerate) bring slowdowns (i.e., 0.75x-0.82x). The results echos DeepSpeed’s documentation (dee, 2023). However, with Lookahead Decoding, we can further utilize the FLOPs of multiple GPUs to reduce the inference latency (e.g., 4x on ClassEval).

3 Generation Quality of Lookahead Decoding

We assess the generation quality of Lookahead Decoding on LLaMA-2-7B-Chat model with the prompts in Appendix D on summarization datasets (Chen et al., 2023; Leviathan et al., 2023) in Tab. 2. Whether the sampling is activated, Lookahead Decoding can reserve the output distribution quality, which is evaluated in rouge-1, rouge-2, and rouge-L (Lin, 2004), while achieving 1.46x-1.60x speedups compared with autoregressive decoding. Using sampling gives smaller speedups as the acceptance ratio is lower according to the sampling verification algorithm 4, which aligns with the results in the previous research (Chen et al., 2023; Leviathan et al., 2023). We further verify that using greedy sampling and advanced integrations will not change the generation quality in Appendix E.

4 Ablation Study

In this section, we study the importance of the lookahead and verification branch in achieving a high speedup. We experiment on LLaMA-2-7B-Chat and MT-Bench on S1 with various settings. The results are shown in Tab. 3.

We ablate the importance of lookahead branch by comparing the performance of using a lookahead branch to the recent methods of using prompts as reference (Yang et al., 2023; Saxena, 2023). This comparison assumes that Lookahead Decoding does not use the prompt to build the n-gram pool. We use the implementation in transformers v4.37 of prompt lookup as a baseline (②, with prompt_lookup_num_tokens=10). We also use prompt to build n-gram pool to augment Lookahead Decoding (③④⑥⑨). The results show that although using a minimal lookahead branch (W=1W=1) with various N,GN,G settings (③④⑤⑥) can obtain a decent speedup on MT-Bench, it is still not as good as using balanced branches (⑧). We can find that prompt lookup can surpass prompt as reference implementation in Lookahead Decoding. This is because our method checks if nn-gram starts with one token that exactly matches the last generated token while prompt lookup in transformers v4.37 checks several starting tokens for a better speculation.

We ablate the importance of verification branch by reporting the speedup of using a tiny verification branch and a large lookahead branch (⑦, G=1G=1) . It shows lower performance due to lower potential in accepting speculations compared with a balanced branches (⑧).

Besides, our evaluation shows that using prompt as reference can further boost Lookahead Decoding (⑧ and ⑨). We have integrated them in our implementation.

5 Discussion and Limitation

The main limitation of Lookahead Decoding is that it requires extra computations. Our experimental results show that on A100, the configuration in Tab. 4 works near optimally in most cases for single batch serving. Because the per-step FLOPs are roughly proportional to the number of per-step input tokens, which is (W+G)∗(N−1)(W+G)*(N-1). If we ignore the attention cost’s increase with sequence length, the 7B, 13B, and 34B models require 120x, 80x, and 56x extra FLOPs per step, respectively. Since the LLM decoding is memory bandwidth-bound rather than compute-bound, these extra FLOPs only turn into a limited wall-clock slowdown for each step.

Given this, Lookahead Decoding needs large surplus FLOPs to obtain high speedups. Running in compute-bound environments (e.g., serving with a large batch size) may cause slowdowns. Another example is shown in Fig. 8, where lower speedup is observed when the GPU’s cap FLOPs is smaller (e.g., on RTX 3090 GPUs).

Based on §4, we need to exponentially increase the per-step FLOPs to obtain a linear reduction in decoding steps. Hence, the setting in Tab. 4 faces a diminishing return. However, when FLOPs are not rich, we see that a gentle speedup (e.g., 30%30\% on RTX 3090 and >50%>50\% on A100) on MT-Bench easily achievable, as in Fig. 8, which is a free lunch that requires no extra model, training, or changing the output distribution.

Related Work

Speculative decoding (Chen et al., 2023; Leviathan et al., 2023) pioneer in speedup autoregressive decoding with a draft model. Different methods for obtaining speculations are researched. Specinfer (Miao et al., 2023) uses many draft models obtained from distillation, quantization, and pruning to conduct speculations together. Medusa (Cai et al., 2024), OSD (Liu et al., 2023), and EAGLE (Li et al., 2023) use training to obtain a draft model. REST (He et al., 2023) uses the finetuning dataset itself as a datastore to lookup speculations, while other works (Yang et al., 2023; Saxena, 2023) uses prompt as a reference for speculations. Different from these methods, Lookahead Decoding uses LLM’s parallel generation ability for speculations. Sampling methods are also researched. Specinfer maintains output distribution by a tree-based sampling algorithm. Medusa uses a typical acceptance scheme to accelerate when the temperature is large but does not persist on an exact output distribution. Lookahead Decoding follows Specinfer to maintain output distribution but with multiple disjoint nn-grams.

Conclusion

In this paper, we present Lookahead Decoding to parallelize the autoregressive decoding of LLMs without changing the output distribution. It shows notable speedup without a draft model and can linearly decrease the decoding steps with exponential investment in per-step FLOPs.

References

Appendix A Algorithms

Appendix B Proof: Output distribution preserved disjoint n-gram verification

The sampling verification in Lookahead Decoding is adapted from the algorithm in Specinfer but with all speculations generated by the greedy sample. It does not change the output distribution from a fundamental point that how the draft model generates speculations is unimportant.

Theorem A For a given LLM, prompt and previously generated tokens x=(x1,x2,...,xi)\mathbf{x}=(x_{1},x_{2},...,x_{i}), and GG speculations s=(s1,s2,...,sG)\mathbf{s}=(s_{1},s_{2},...,s_{G}) of next token xi+1x_{i+1}. Each speculation token is sampled by a greedy sample (i.e., probability of 1). We use P(v∣x)P(v|\mathbf{x}) to represent the probability of xi+1=vx_{i+1}=v sampled by the LLM and use Q(v∣x)Q(v|\mathbf{x}) to represent the probability of xi+1=vx_{i+1}=v sampled by our proposed algorithm 4. We use P(v)P(v) and Q(v)Q(v) for short. We need to prove P(v)=Q(v)P(v)=Q(v) for any GG, and any vv and sjs_{j} from the full vocabulary VV.

The proof of this part corresponds to line 14 to line 44 in algorithm 4. Given speculations s\mathbf{s}, we use aj(v)a_{j}(v) to represent the probability that the token vv is accepted by the jj-th speculation (line 18-line 30), and rj(sj)r_{j}(s_{j}) is the probability that the token sjs_{j} is rejected by the jj-th speculation (line 30-line 35), where sjs_{j} is the jj-th speculation’s token. Moreover, aG+1′(v)a_{G+1}^{\prime}(v) is the probability of being accepted by the sampling at line 41. For simplicity, we use aja_{j} to represent aj(v)a_{j}(v), aj′a_{j}^{\prime} to represent aj′(v)a_{j}^{\prime}(v), and use rjr_{j} to represent rj(sj)r_{j}(s_{j}). We use P1\mathcal{P}_{1} to represent the probability distribution obtained in line 13 and Pj\mathcal{P}_{j} to present the updated probability before the jj-th speculation. We have P1(v)=P(v)\mathcal{P}_{1}(v)=P(v) as P1(v)\mathcal{P}_{1}(v) is never updated. We define QG(v)Q_{G}(v) is the probability of xi+1=vx_{i+1}=v sampled by algorithm 4 when we have GG speculations. Then we should have:

QG(v)=a1+r1a2+r1r2a3+...+aG∏k=1G−1rk+aG+1′∏k=1GrkQ_{G}(v)=a_{1}+r_{1}a_{2}+r_{1}r_{2}a_{3}+...+a_{G}\prod\limits_{k=1}^{G-1}r_{k}+a_{G+1}^{\prime}\prod\limits_{k=1}^{G}r_{k}

We use induction to prove QG(v)=P(v)Q_{G}(v)=P(v) for any G≥1G\geq 1, any v∈Vv\in V, and any sj∈Vs_{j}\in V with 1≤j≤G1\leq j\leq G:

When G=1G=1, we have QG(v)=a1+r1a2′Q_{G}(v)=a_{1}+r_{1}a_{2}^{\prime}. The initial guess s1s_{1} can be either the same at vv or be different from vv.

When s1=vs_{1}=v, a1a_{1} equals P1(v)\mathcal{P}_{1}(v) at line 17, which is the same as P(v)P(v) as it is never updated. Upon this, we have r1=1−a1=1−P1(v)=1−P(v)r_{1}=1-a_{1}=1-\mathcal{P}_{1}(v)=1-P(v). And, a2′a_{2}^{\prime} is the updated probability P2(v)\mathcal{P}_{2}(v) at line 42. Since P2(v)\mathcal{P}_{2}(v) is set to zero once rejected at line 32, a2′=0a_{2}^{\prime}=0. In this case, QG(v)=P(v)+(1−P(v))∗0=P(v)Q_{G}(v)=P(v)+(1-P(v))*0=P(v).

When s1≠vs_{1}\neq v, a1a_{1} should be even if sis_{i} is accepted. Moreover, we have r1=1−P1(s1)r_{1}=1-\mathcal{P}_{1}(s_{1}). Then P2(v)\mathcal{P}_{2}(v) is updated to P1(v)1−P1(s1)\frac{\mathcal{P}_{1}(v)}{1-\mathcal{P}_{1}(s_{1})} at lines 32 and 33. Then a2’=P2(v)a_{2}’=\mathcal{P}_{2}(v). In this case, QG(v)=0+r1∗P(v)r1=P(v)Q_{G}(v)=0+r_{1}*\frac{P(v)}{r_{1}}=P(v).

When G=gG=g holds, which means Qg(v)=a1+r1a2+...+ag∏k=1g−1rk+ag+1′∏k=1grk=P(v)Q_{g}(v)=a_{1}+r_{1}a_{2}+...+a_{g}\prod\limits_{k=1}^{g-1}r_{k}+a_{g+1}^{\prime}\prod\limits_{k=1}^{g}r_{k}=P(v) for any sj,v∈Vs_{j},v\in V, 1≤j≤g1\leq j\leq g.

We prove Qg+1(v)=Qg(v)−ag+1′∏k=1grk+ag+1∏k=1grk+ag+2′∏k=1g+1rk=P(v)Q_{g+1}(v)=Q_{g}(v)-a_{g+1}^{\prime}\prod\limits_{k=1}^{g}r_{k}+a_{g+1}\prod\limits_{k=1}^{g}r_{k}+a_{g+2}^{\prime}\prod\limits_{k=1}^{g+1}r_{k}=P(v) for the same sj,v∈Vs_{j},v\in V, 1≤j≤g1\leq j\leq g, and any sg+1∈Vs_{g+1}\in V.

When sg≠vs_{g}\neq v, we have ag=0a_{g}=0. If Pg(v)=0\mathcal{P}_{g}(v)=0, we have all ag+1′=0a_{g+1}^{\prime}=0, ag+1=0a_{g+1}=0 and ag+2′=0a_{g+2}^{\prime}=0. It ensures that Qg+1(v)=Qg(v)−0+0+0=Qg(v)=P(v)Q_{g+1}(v)=Q_{g}(v)-0+0+0=Q_{g}(v)=P(v).

If Pg(v)≠0\mathcal{P}_{g}(v)\neq 0, Pg+1(v)=Pg(v)rg=P1(v)∏k=1grk\mathcal{P}_{g+1}(v)=\frac{\mathcal{P}_{g}(v)}{r_{g}}=\frac{\mathcal{P}_{1}(v)}{\prod\limits_{k=1}^{g}r_{k}} since sg≠vs_{g}\neq v by observation.

Then we have Qg+1(v)=Qg(v)−Pg+1(v)∏k=1grk+ag+1(v)∏k=1grk+ag+2′∏k=1g+1rk=Qg(v)−P1(v)∏k=1grk∏k=1grk+ag+1∏k=1grk+ag+2′∏k=1g+1rk=Qg(v)−P(v)+ag+1(v)∏k=1grk+ag+2′∏k=1g+1rkQ_{g+1}(v)=Q_{g}(v)-\mathcal{P}_{g+1}(v)\prod\limits_{k=1}^{g}r_{k}+a_{g+1}(v)\prod\limits_{k=1}^{g}r_{k}+a_{g+2}^{\prime}\prod\limits_{k=1}^{g+1}r_{k}=Q_{g}(v)-\frac{\mathcal{P}_{1}(v)}{\prod\limits_{k=1}^{g}r_{k}}\prod\limits_{k=1}^{g}r_{k}+a_{g+1}\prod\limits_{k=1}^{g}r_{k}+a_{g+2}^{\prime}\prod\limits_{k=1}^{g+1}r_{k}=Q_{g}(v)-P(v)+a_{g+1}(v)\prod\limits_{k=1}^{g}r_{k}+a_{g+2}^{\prime}\prod\limits_{k=1}^{g+1}r_{k}. Here we have another two cases:

➀ If sg+1=vs_{g+1}=v, ag+1∏k=1grk=P1(vi=v)∏k=1grk∏k=1grk=P(v)a_{g+1}\prod\limits_{k=1}^{g}r_{k}=\frac{\mathcal{P}_{1}(v_{i}=v)}{\prod\limits_{k=1}^{g}r_{k}}\prod\limits_{k=1}^{g}r_{k}=P(v) and ag+2′=Pg+2(v)=0a_{g+2}^{\prime}=\mathcal{P}_{g+2}(v)=0. We have Qg+1(v)=Qg(v)−P(v)+P(v)+0=Qg(v)=P(v)Q_{g+1}(v)=Q_{g}(v)-P(v)+P(v)+0=Q_{g}(v)=P(v).

➁ If sg+1≠vs_{g+1}\neq v, ag+1=0a_{g+1}=0. ag+2′∏k=1g+1rk=Pg+2(v)∏k=1g+1rk=P1(v)∏k=1g+1rk∏k=1g+1rk=P(v)a_{g+2}^{\prime}\prod\limits_{k=1}^{g+1}r_{k}=\mathcal{P}_{g+2}(v)\prod\limits_{k=1}^{g+1}r_{k}=\frac{\mathcal{P}_{1}(v)}{\prod\limits_{k=1}^{g+1}r_{k}}\prod\limits_{k=1}^{g+1}r_{k}=P(v). So Qg+1(v)=Qg(v)−P(v)+0+P(v)=Qg(v)=P(v)Q_{g+1}(v)=Q_{g}(v)-P(v)+0+P(v)=Q_{g}(v)=P(v)

When sg=vs_{g}=v, Pg+1(v)\mathcal{P}_{g+1}(v) is set to zero at line 32 after this step. In this case, we have ag+1′=Pg+1(v)=0a_{g+1}^{\prime}=\mathcal{P}_{g+1}(v)=0, ag+1=Pg+1(v)=0a_{g+1}=\mathcal{P}_{g+1}(v)=0 and ag+2′=0a_{g+2}^{\prime}=0. It makes that Qg+1(v)=Qg(v)−0+0+0=Qg(v)=P(v)Q_{g+1}(v)=Q_{g}(v)-0+0+0=Q_{g}(v)=P(v)

This part of the proof guarantees that from line 14 to line 44 in Algorithm 4, any new token appended to o\mathbf{o} can follow the original distribution of the LLM. Line 21 to line 28 guarantees that sequences in V\mathbf{V} share the same prefix of length i−1i-1 in every iteration. This further guarantees that P\mathcal{P} from D[j]i\mathbf{D}[j]_{i} is the same for all jj, follows the wanted distribution. Thus, the correctness of the whole sampling algorithm is proved.

Appendix C Derivation of Expectation of The Number of Accepted Tokens

We first start with single-candidate speculation. We need to obtain the probability of accepting ii tokens as P(#accepted tokens=i)P(\#accepted\ tokens=i) for all possible ii. Since the speculation’s length is γ\gamma, the probability of accepting ii tokens with i≥γ+2i\geq\gamma+2 is 0. P(#accepted tokens=1)P(\#accepted\ tokens=1) is the probability of the first token being rejected, which is 1−α1-\alpha. The probability P(#accepted tokens=i)=P(#accepted tokens=i−1)∗αP(\#accepted\ tokens=i)=P(\#accepted\ tokens=i-1)*\alpha, for all i≤γi\leq\gamma. The probability P(#accepted tokens=γ+1)P(\#accepted\ tokens=\gamma+1) is accepting all tokens, which is αγ\alpha^{\gamma}. Thus we have the following, which is Eq. 4:

We then investigate the case of speculations with a batch size of bb. We need to obtain the probability of accepting ii tokens as P(#accepted tokens=i)P(\#accepted\ tokens=i). Since all speculations’ length is γ\gamma, the probability of accepting aa tokens with a≥γ+2a\geq\gamma+2 is 0. We use pip_{i} to denote (1−αi)b(1-\alpha^{i})^{b}, which is the probability that at most ii tokens are accepted in all bb speculations. For all i≤γi\leq\gamma, we should have P(#accepted tokens=i)=pi−pi−1P(\#accepted\ tokens=i)=p_{i}-p_{i-1}. And, the probability P(#accepted tokens=γ+1)P(\#accepted\ tokens=\gamma+1) should be (1−pγ)(1-p_{\gamma}). Thus we have the following, which is Eq. 5:

Appendix D Prompt for LLaMA-2-Chat on Summarization Tasks

We use the following as the prompt for summarization task, modified from (Ruan et al., 2023).

Appendix E Verification of Generation Quality for Greedy Sampling and Advanced Supports

Generation Quality with Greedy Search is not changed. Theoretically, Lookahead Decoding does not change the output generation of greedy search due to the verification mechanism. However, Lookahead Decoding’s output does not perfectly align with the huggingface’s implementation of greedy search in practice. We attribute this discrepancy to numerical accuracy issues. To substantiate this claim, we compared the output results as follows. We use the LLaMA-2-7b-Chat model’s single precision (FP32) inference with huggingface’s greedy search on 160 turns on the MT-Bench dataset as a baseline. With single precision inference, the outputs of Lookahead Decoding (on 1GPU, 4GPUs, and 8GPUs) are the same as the output of the baseline. With half-precision (FP16) inference, huggingface’s greedy search has 35 out of 160 (w/o FlashAttention) and 42 out of 160 (w/ FlashAttention) answers not perfectly aligned with the baseline output. In contrast, Lookahead Decoding and its integration with FlashAttention and multi-GPU inference has 35-44 results different from the baseline output under different settings. We claim this result can show that Lookahead Decoding can retain the output distribution using a greedy search within the numerical error range (not worse than huggingface’s half-precision inference). Besides, Tab. 2 also strengthens the statements for greedy search.

Generation Quality with LP and FlashAttention Augmentation is not changed. We verify that FlashAttention and LP Support will not change the compression ratio (S\mathcal{S}) of vanilla Lookahead Decoding. We compared each 18 generations of Lookahead Decoding w/ FlashAttention and w/o FlashAttention (7B and 13B model on MT-Bench, HumanEval, and ClassEval); the average S\mathcal{S} w/ FlashAttention is 3.267 while w/o FlashAttention is 3.259, with less than 0.3% differences. We also compared 6 generations of Lookahead Decoding on a single GPU and 12 generations with LP (7B model on MT-Bench, HumanEval, and ClassEval, both with N=5N=5, W=15W=15, and G=15G=15). The average S\mathcal{S} on a single GPU is 2.558, while on multiple GPUs, it is 2.557, with less than 0.1% differences. We claim that our advanced support does not change S\mathcal{S}.