REST: Retrieval-Based Speculative Decoding
Zhenyu He, Zexuan Zhong, Tianle Cai, Jason D. Lee, Di He
Introduction
Transformer-based Large Language Models (LLMs) have emerged as a foundation model in natural language processing (Vaswani et al., 2017; Devlin et al., 2019; Brown et al., 2020; Zhang et al., 2022; Scao et al., 2022; Chowdhery et al., 2022; Zeng et al., 2022; Touvron et al., 2023). While they achieve impressive performance across various tasks, the inference cost is huge in practical scenarios. During inference, the model autoregressively uses the preceding context to generate the next token. Each iteration requires reloading the billion-parameter LLM from the High-Bandwidth Memory (HBM) to the on-chip cache of modern accelerators like GPUs, merely for computing the next token, making the whole generation inefficient and time-consuming.
A recent direction in accelerating the LLM generation is to reduce the number of forward processes with LLMs while guaranteeing the quality of the output sequence simultaneously. Speculative decoding (Leviathan et al., 2023; Chen et al., 2023; Miao et al., 2023; Spector and Re, 2023) is one of the typical approaches in this direction. Intuitively, speculative decoding methods leverage a small LM to generate tokens with less computational cost. During inference, the method first uses the small LM to create a draft token sequence and then uses the LLM for verification. If the predictions from both models are consistent, we accept the draft and return it to the user. Here, the actual token generation is carried out using the small LM, and the large LM is only used to validate the draft, which can be performed in parallel and requires reloading the memory only once. Consequently, the entire framework of speculative decoding reduces the overall inference cost.
However, obtaining a high-quality draft model remains challenging: It must balance small size and strong predictive power while matching the vocabulary of the base model; also, it should integrate well into a distributed system for serving. Therefore, people often need to train a draft model specifically for their model and use cases (Chen et al., 2023; Miao et al., 2023; Cai et al., 2023). In this study, rather than relying on an additional small LM, we investigate using a data corpus directly to construct the draft token sequence in speculative decoding. We develop a retrieve-based approach, called Retrieval-Based Speculative Decoding (REST) (Figure 1). Compared to previous approaches, our retrieval-based system replaces the parametric draft model with a non-parametric retrieval datastore, which can easily port to any LLM and accelerate its inference.
To use REST, the first step is constructing the datastore. In this paper, we leverage either the pretraining data corpus or the instruction-tuning data corpus to build our datastore, which serves as the source for the draft token sequence. During each inference step, we first use previous tokens (pre-generated tokens or prompt tokens) as queries to identify exact matches in the datastore. The subsequent tokens from these exact matches are candidate tokens. A Trie is constructed using these candidates. The nodes with the highest frequencies are selected as the draft tokens. This sequence then undergoes verification by the LLM through a single forward pass, aided by a meticulously designed attention mask known as tree attention Cai et al. (2023); Miao et al. (2023); Spector and Re (2023). Finally, we directly sample tokens from the conditional probability. As many subsequences during generation likely appear in the datastore, REST can frequently produce multiple tokens per step.
We conduct extensive experiments to test the efficiency and effectiveness of REST in different scenarios. For the code domain, we use a portion of Python pretraining code (2.7M samples) from The Stack Kocetkov et al. (2022) as the datastore and accelerate CodeLlama Rozière et al. (2023) 7B and 13B respectively. The results show on HumanEval Chen et al. (2021) REST achieves to speedup. For the general domain, we construct a datastore using UltraChat Ding et al. (2023), containing around 774K conversations. The results show on MT-Bench Zheng et al. (2023) REST accelerates 7B and 13B Vicuna Chiang et al. (2023) by to respectively.
Related Work
Improving the efficiency of LLM inference has been an emergent research direction in recent years. Broadly, previous attempts can be divided into two categories: lossless acceleration and lossy acceleration. Lossy acceleration approaches aim to learn efficient models that can execute faster and act similarly to a target LLM. These methods include pruning (Wang et al., 2021; Hubara et al., 2021; Frantar and Alistarh, 2023), quantization (Yao et al., 2022; Park et al., 2022; Dettmers et al., 2022; Frantar et al., 2022; Xiao et al., 2023; Liu et al., 2023) and knowledge distillation (Sanh et al., 2019). Lossless acceleration strategies focus on directly accelerating the target LLM from different perspectives, such as memory and IO optimization (Dao et al., 2022; Dao, 2023; Kwon et al., 2023; Sheng et al., 2023), and ways to reduce the function calls of LLM during decoding, e.g., speculative decoding (Stern et al., 2018; Leviathan et al., 2023; Chen et al., 2023; Miao et al., 2023; Spector and Re, 2023; Cai et al., 2023). This work falls within the second branch. Speculative decoding (Leviathan et al., 2023; Chen et al., 2023; Miao et al., 2023; Spector and Re, 2023) leverages a smaller model to generate a draft and use LLM to verify the draft tokens with a single forward pass. In this framework, blockwise parallel decoding (Stern et al., 2018) and Medusa (Cai et al., 2023) train multiple heads based on the LLM for draft token generation.
Our method diverges from these approaches by retrieving draft tokens from a datastore, presenting a novel avenue for efficiency improvement in large language model generation. While there is a similar study, LLMA (Yang et al., 2023), that employs retrieval to accelerate generation, our work distinguishes itself in two primary ways: (1) The LLMA approach is tailored towards scenarios where referred contexts (as in Retrieval-Augmented Generation and Cache-Assisted Generation) are provided during generation. It retrieves draft tokens from these referred contexts. In contrast, our method retrieves draft tokens from a comprehensive datastore, thereby not being confined to a small context. (2) In the LLMA framework, the retrieved instance is typically limited to one or a handful. Our method, however, is designed to handle a much larger number of retrieved instances. This difference in approach allows us to leverage a wider information base during the generation process.
Retrieval-Based Speculative Decoding
In this section, we first provide notations and a background overview of speculative decoding and then introduce our proposed REST framework.
We use to denote a token where is the vocabulary. At each time step , given the preceding context , the autoregressive decoding method generates the token at position according to:
where is the conditional probability distribution calculated by the LLM with parameter . In this process, a forward run of the LLM is required at each step of generation. This is significantly time-consuming due to the memory bandwidth and cannot fully exploit the computational power of modern GPU hardware (Shazeer, 2019).
Speculative decoding aims to reduce the computational cost during inference by reducing the count of executions with . In addition to the LLM , speculative decoding leverages another language model of a much smaller size with parameter . At step , the method operates by iteratively executing the following steps.
Although the tokens are still generated one by one, the computational cost of this process is reduced as it uses instead of .
Draft verification
Draft acceptance
2 Our Approach: REST
While in the classic speculative decoding, a smaller LM is used as the draft model, finding a high-quality draft model is usually challenging for several reasons: (1) For efficiency, the draft model needs to be lightweight enough to not introduce much overhead. (2) For quality, it needs to predict the LLM output accurately. (3) For system integration, it needs the same vocabulary set as the LLM, and an architecture that distributes easily with a similar configuration to the LLM Chen et al. (2023). These challenges require carefully selecting or even training custom draft models for each new LLM.
In this paper, we solve the challenges differently. We develop a training-free approach to speculative decoding that can easily integrate with any new model to accelerate inference. Instead of relying on a parametric draft model, our method Retrieval-Based Speculative Decoding (REST) proposes using retrieval for draft construction. An overview of REST is shown in Figure 1. In the following, we first describe constructing a datastore and operations on it, then demonstrate using it for draft construction and verification. Together, REST provides an efficient, high-quality, and easy-to-integrate solution for accelerating the inference of LLMs.
REST operates based on a pre-built datastore , where represents a context and represents the corresponding continuation of the context . Given a text/code corpus, we construct the datastore using the prefix context and the corresponding continuation at each position.
Retrieving from the datastore
At inference, given a context , our objective is to construct the draft tokens which are likely the continuations of the context. Different from vanilla speculative decoding that uses a small LM to construct the draft, we leverage the built datastore and directly retrieve draft tokens from the datastore. We first use the context to retrieve context-continuation pairs from the datastore and construct a set of continuation candidates :
where implements a retrieval process in the datastore that returns a set of context-continuation pairs by using as the query. It is straightforward to use recent dense retrieval models Khandelwal et al. (2020); Karpukhin et al. (2020) to find contexts that are similar to . However, using dense retrievers adds additional overhead during inference. We instead use a fast exact-match method to retrieve continuation candidates.
Our retrieval process is shown in Algorithm 1. We aim to find contexts in that match the longest suffix of . We employ a greedy strategy and start from a pre-defined match length upper limit . For each suffix length , we obtain the context ’s suffix with tokens (line 5), and obtain all the contexts that match as a suffix (line 6). If at least one context in matches the current (i.e., ), we return the corresponding context-continuation pairs as the retrieval result; otherwise we decrease the matching length by one and try to match a shorter suffix (line 7). We use a suffix array Manber and Myers (1993) to implement efficient exact match in datastore for a given . The retrieval process leads to negligible overhead () in our experiments (see details in Section 5).
Draft construction from retrieved results
The retrieved result includes possible continuations of the context . For each , any prefix of can serve as draft tokens of in the speculative decoding and be further verified by the LLM. Note that the retrieved set of continuation candidates can be large. It is not feasible to use all candidates as draft tokens and feed them into the LLM for verification. Here we present how we select high-quality draft tokens from the retrieved set . A naive strategy is to sample a subset of sequences in as the draft tokens. However, this is suboptimal as the shared prefixes of continuations in may be considered and verified multiple times.
We select draft tokens from the retrieved result using a Trie. In the Trie, the unique path from a node to the root node corresponds to a prefix of . For each node, we assign a weight reflecting the number (frequency) of the corresponding prefix that appears in the retrieved candidates. As shown in Algorithm 2, we first construct a Trie using all sequences in , and the node weight is updated when a candidate is inserted into the Trie (lines 2-7). The Trie data structure allows us to prioritize tokens using the weights and select high-frequency prefixes (lines 8-15). In the practical implementation, we choose a subtree that contains the top nodes with the highest weights, which equals to selecting the top high-frequency prefixes as the draft sequences.
Draft verification of REST
In REST, multiple draft sequences may be retrieved from the datastore. While one might initially approach the drafts independently and feed them into the LLM as distinct sequences in a batch, practical observations reveal that many drafts share common prefixes. This leads to redundant computation of Transformer layers on these shared prefixes across different sequences, resulting in a waste of computational power. To optimize the efficiency, we construct a pseudo sequence from the subtree using breadth-first search. By definition, it can be immediately obtained that each draft constitutes a sub-sequence of this pseudo sequence, and any shared prefix appears only once. To correctly execute LLM on this pseudo sequence, we implement a carefully designed attention mask in each attention layer, ensuring that the computation of each token precisely reflects its dependencies in the original draft sequence. This attention strategy is also known as tree attention (Cai et al., 2023; Miao et al., 2023; Spector and Re, 2023).
Draft acceptance of REST
We adopt a similar acceptance strategy compared to the original speculative decoding. By feeding the drafts into LLM, we obtain the conditional distribution at each position given by and then check the correctness of the draft token. All correct tokens from the start will be accepted, and the draft tokens after the first mistake will be rejected.
Comparison with existing approaches
Although REST follows a schema similar to that of speculative decoding, it offers significant advantages over existing approaches. Current speculative decoding methods rely on a high-quality small model to generate draft tokens Leviathan et al. (2023); Chen et al. (2023). Such methods must strike a balance between a small size and strong predictive power, while also matching the vocabulary of the base model. Moreover, they require additional GPU memory and introduce complexity during inference. In contrast, REST directly retrieves draft tokens from a datastore, which can be easily integrated with language models of any size, vocabulary, or architecture. Different from Stern et al. (2018) and Cai et al. (2023) which train specialized modules to create a draft model, REST eliminates the need for any additional training steps and can serve as a plug-and-play solution of efficient decoding across different models. Furthermore, the effectiveness of REST is affected by the quality of retrieval results. This opens up the opportunities to further enhance REST by using a better/larger datastore or an advanced retrieval model. We also note that in addition to using REST directly, it is possible to combine REST with the vanilla speculative decoding. This combination can enhance the generation speed of the small LM. We leave this for future work.
Experiments
We implement two sampling mechanisms: greedy sampling and nucleus sampling (Holtzman et al., 2019) for the LLM. Greedy sampling selects the token with the highest probability at each step. Nucleus sampling, also known as top- sampling, generates tokens by sampling from the most probable tokens in the model’s predicted distribution until their cumulative probability reaches the threshold . It is worth noting that under our approach, we only accept draft tokens if they match the tokens sampled from the LLM. As a result, the sequences produced using REST are identical to those generated by standard autoregressive generation.
Datasets and models
We conduct experiments on two datasets: HumanEval (Chen et al., 2021) and MT-Bench (Zheng et al., 2023). HumanEval is a dataset that includes 164 human-written Python programming problems. The goal for the models is to generate code solutions using provided docstrings as prompts. On the other hand, MT-Bench contains 80 multi-turn questions designed to emulate real-world multi-turn dialogues. We compare the generation speed of standard autoregressive generation with REST, focusing on both the HumanEval and MT-Bench datasets. For HumanEval, we perform 1-shot evaluation for greedy sampling and 10-shot evaluation for nucleus sampling and employ the CodeLlama (Rozière et al., 2023). While for MT-Bench, we perform 1-shot evaluation for both greedy sampling and nucleus sampling and utilize Vicuna (Chiang et al., 2023). We test both the 7B and 13B configurations of CodeLlama and Vicuna, with a maximum generation limit of 512 tokens and 1024 tokens, respectively. All experiments are conducted on a single NVIDIA A6000 GPU and 96 CPU cores. All results are averaged across three different runs.
Hyperparameters
When performing exact match in the datastore, the starting context suffix length, , is set to 16, and is progressively reduced by one until we find matching contexts in the datastore. The length of each retrieved continuation candidate denoted as , is truncated to 10. Empirical results from Medusa (Cai et al., 2023) suggest 64 draft tokens to be an optimal computation configuration. Hence, we limit the maximum number of selected draft tokens in the constructed Trie to 64, designated as .
Metrics
The first metric we use is Mean Token Time, which is the average generation time of one token for the LLM. Another metric, Mean Generated Length, is calculated as the ratio of the length of the generated tokens to the number of forward steps taken by the original LLM. Formally, if denotes the length of the generated tokens and represents the number of forward steps, the Mean Generated Length, , is given by:
Note that the Mean Generated Length () acts as the upper limit of the speedup that REST can achieve, ignoring the overhead for retrieving and constructing draft tokens.
Datastores
For CodeLlama, we construct a datastore using a portion of the Python pretraining code from The Stack (Kocetkov et al., 2022). This dataset comprises approximately 2.7M Python code samples and results in a datastore with a size of 27GB. On the other hand, for Vicuna, we construct a datastore using data derived from UltraChat Ding et al. (2023). This dataset consists of around 774K conversations from ChatGPT, yielding a datastore with a size of 12GB.
2 Main Results
Table 1 compares the generation speed of REST and the speed of the standard autoregressive decoding approach.
Regarding generation speed, REST demonstrates a significant speed enhancement, achieving to increase for CodeLlama in the HumanEval benchmark. The MT-Bench benchmark also reveals a speedup for Vicuna when using our method, with a factor ranging from to . These empirical results lend weight to the effectiveness of our method for speeding up the generation process of LLMs. Note that the speedup of nucleus sampling is not as good as that of greedy sampling. We speculate that this drop in performance is caused by the randomness introduced by nucleus sampling.
Another intriguing observation that emerges from these results is the domain-dependent nature of the speed improvements. This characteristic has also been noted in other methods like speculative decoding Chen et al. (2023) and Medusa Cai et al. (2023). Specifically, the speedup achieved with REST is significantly greater in the HumanEval benchmark than in the MT-Bench benchmark, suggesting that the effectiveness of REST may vary depending on the specific domain.
Additionally, it is important to note that the average time (divided by the total number of tokens) required for retrieval (which includes the time taken to construct the Trie) is less than 1 ms. This time is very small and can, for all practical purposes, be considered negligible. This negligible retrieval time further underscores the efficiency of REST.
Ablation Study
To gain a deeper understanding of our method, we conduct a series of ablation studies and analyses focused on each individual component.
Increasing the size of the datastore is an effective strategy for enhancing the accuracy of retrieved draft tokens in the Trie, which in turn can significantly boost generation speed. In Table 2, we show that as the datastore size increases, both the Mean Generated Length and Mean Token Time correspondingly improve. However, it’s important to note that the speedup growth is not as pronounced as that of the Mean Generated Length. This discrepancy could be attributed to the overhead of getting draft tokens. We assume that in industry applications, there will be ample disk storage to build a large datastore and ample CPU cores for fast retrieval. We also visualize the trend of scaling the retrieval datastore size in Figure 2. From this, we can infer that there is still potential to achieve even faster speeds with a larger datastore.
Effect of the maximum number of draft tokens
Increasing the volume of draft tokens can potentially lead to a higher Mean Generated Length by the LLM. However, this also escalates the computational burden on GPUs during verification. As shown in Figure 3, an initial speed increase is observed as the maximum number of draft tokens increases. However, beyond the threshold of 48 draft tokens, the speed stabilizes to an average of approximately 11.75 ms per token. When the token count exceeds 200, it leads to a slowdown. Therefore, while it is possible to achieve similar speeds with a large maximum number of draft tokens, it’s more efficient to limit the number to a smaller one to avoid unnecessary strain on GPUs.
Effect of draft token selecting strategies
We compare selecting draft tokens in the Trie with randomly sampling retrieved continuation candidates as draft tokens. For an equitable comparison, we employ a random sampling technique to sample at most eight sequences from all the retrieved candidates. Furthermore, each sequence is truncated to a maximum length of 8. This results in a maximum number of 64 draft tokens, corresponding to the maximum number of selected draft tokens from the Trie. The data presented in Table 3 indicates that selecting draft tokens from the Trie, as opposed to employing a random sampling approach, enhances the performance.
Visualization of matched suffix length
The distribution of matched suffix length of the context is illustrated in Figure 4. From this graphic, it is apparent that almost all cases contain a matched suffix length. Notably, shorter suffix lengths ranging from 2 to 9 comprise the majority of the matched cases, accounting for a substantial 85% of the total. In contrast, longer suffix lengths, which range from 10 to 16, constitute a minority, making up only 15% of the cases.
Effect of the choice of the maximum suffix length
We vary the value of to test the generation speed of REST. The outcomes of this study are depicted in Figure 5. An interesting observation is that when the value of is set to less than 6, there is a substantial increase in the generation time. Conversely, when exceeds 6, the generation speed remains consistently high and appears to be largely unaffected by further changes to the value. Hence, in practice, there is no substantial need to expend excessive efforts in selecting the precise optimal value of .
Conclusion
In this work, we propose REST: retrieval-based speculative decoding. Instead of requiring a small LM, REST employs a datastore for retrieving and employing draft tokens. We construct a Trie to select the most probable draft tokens. REST is not only straightforward to implement but also easily integrates into the generation processes of any existing language models without necessitating additional training.
Future Directions
We consider four future directions that can further boost the performance of REST:
In this work, we construct datastores from pretraining datasets or instruction-tuning datasets. However, for improved alignment with the original model, it might be advantageous to consider constructing datastores from content generated by the model itself.
In this work, we directly implement REST to enhance the generation speed of the LLM, while it’s possible to combine REST with speculative decoding to enhance the generation speed of the small LM.
For situations where resources are limited, it’s worthwhile to explore methods of minimizing the datastore size without compromising performance.
Integration of in-context abilities. For instance, the challenge of retrieving personalized variable names in code generation—a task that inherently requires understanding context—raises an interesting question: How can we empower retrieval methodologies to effectively deal with such complexities?