SPEED: Speculative Pipelined Execution for Efficient Decoding

Coleman Hooper, Sehoon Kim, Hiva Mohammadzadeh, Hasan Genc, Kurt Keutzer, Amir Gholami, Sophia Shao

Introduction

The Transformer neural network architecture has recently revolutionized NLP, providing massive accuracy gains across a range of tasks . In particular, there has been growing interest in applying Transformers decoders for generative tasks . Unlike Transformer encoders which can process an entire input sequence in parallel, Transformer decoders must be applied autoregressively at inference time as each input token depends on the output classification for the previous token. This means that they exhibit low arithmetic intensity and are typically memory bandwidth-bound . For small batch sizes (as is typical for edge deployment scenarios ), it is extremely difficult to achieve any parallelism. In order to accelerate memory bandwidth-bound decoder inference, we must reduce the number of memory operations required.

In this work, we aim to reduce the latency of memory bandwidth-bound decoder inference by employing speculative execution in order to process tokens at different positions in the sequence in parallel. When employing speculative execution, the forward passes for future tokens are started using speculative output values from earlier tokens. By starting future tokens, we can process them in parallel with finishing the forward passes for earlier tokens. If a prediction is later found to be wrong, we must invalidate all future inferences that were started based on the speculative output value from the incorrect prediction. By still following all iterations through to completion, we can ensure that full model accuracy is maintained.

On its own, speculative execution would not lead to performance benefits within a single network. As shown in (b) in Figure 1, different tokens in the sequence would need to be processed by different layers in the network at the same time, meaning that the number of memory operations required for performing inference would not be reduced (even assuming perfect prediction). Additionally, to support inference on low-resource edge devices, it is crucial to reduce the model’s memory footprint. Parameter sharing is a common method for model compression in Transformer networks . However, although it reduces the size of the network, parameter sharing doesn’t typically provide significant speedup as the standard computation must still be performed for all layers in the network. Even if inference is memory-bound, parameter sharing only reduces the number of memory operations required if the entire model fits in local cache memory, which is restrictive and hardware-dependent.

However, in a network which employs parameter sharing, speculative execution allows us to amortize the memory operations required for the weight matrices across different tokens in the sequence. By employing speculative execution in networks with parameter sharing, we can pipeline inference, thereby reducing memory operations. Each pipeline stage corresponds to passing several tokens at different positions through the same set of linear layers (since the parameters for these linear layers are shared across decoder blocks). Our speculative execution approach therefore allows us to accelerate decoder inference with networks that employ parameter sharing as a model compression method. We believe that our speculative execution approach can make parameter sharing an advantageous model compression strategy for both shrinking the static model size and for accelerating inference.

Method

The parameter sharing scheme in this work corresponds to the “CYCLE" configurations from , meaning that if a group of two decoder layers is shared three times, a forward pass consists of alternating between going through layer 1 and layer 2 three times. Our cyclic parameter sharing scheme is outlined graphically in part (a) of Figure 2. During fine-tuning, we incorporate a weighted loss function inspired by the work of . The purpose is to adapt the output classifier so that it can make early predictions during inference using the output logits from earlier decoder layers (i.e., after different repetitions of decoder layer groups). More formally, the shared loss function is given by Lw=∑i=1GwiLiL_{w}=\sum_{i=1}^{G}w_{i}L_{i}, where GG is the number of decoder layer groups, and LiL_{i} and wiw_{i} correspond to the loss and the applied weighting for group ii, respectively. The default weighting scheme we use was the linear weighting described in , which is given as wi=i/(∑ii)w_{i}=i/(\sum_{i}i). Note that this weighting scheme intentionally weights the loss for later layers higher to ensure the final output accuracy is not degraded. The training scheme using a shared classifier is also illustrated in part (a) of Figure 2.

2 Speculative Pipelined Execution

Diagrams (b) and (c) in Figure 2 outline how the forward pass is performed in our speculative approach. Our decoding algorithm is outlined in detail in Algorithm 1 (Appendix B). In essence, SPEED speculatively predicts future tokens based on early predictions and then concatenates them with the current token for their parallel processing. The crucial feature of SPEED is its invalidation logic since speculative predictions can be sometimes incorrect. To achieve this, our framework keeps track of previous classifications for each token at the previous stage (i.e., before passing through a decoder layer group) and performs the invalidation logic whenever subsequent classifications change after the current stage (i.e., after passing through a decoder layer group). In such a case, any future iterations that have been speculatively initiated using the previous classifications must be flushed out and restarted.

Another crucial implementation detail is that the internal logic in the attention module and the internal Key/Value (KV) cache management logic both need to be modified to facilitate pipelining. The KV cache corresponds to intermediate activations associated with earlier tokens in the sequence, which are required for calculating later tokens. The KV cache management logic has to ensure that when future tokens are invalidated, all previous KV cache updates corresponding to these tokens are also invalidated. These modifications play a key role in ensuring that the final output classification for each token remains unaffected by speculation.

Results

We implement SPEED within the T5X repository, which is built on top of the JAX framework. Our implementation for pipelined inference used a custom decoding function in the T5X framework. Our initial profiling runs also indicated that the existing greedy decoding function in JAX had greater runtime overhead than our custom decoding algorithm, likely due to additional optional arguments that were unused in our experiments. In order to benchmark the networks without parameter sharing, we therefore implemented a stripped-down greedy decoding function to serve as a fair baseline since it has minimal added control logic.

We use the baseline T5-Base decoder-only model architecture , which has 12 decoder layers, a hidden dimension of 768, 12 attention heads each with dimension 64, and an FFN dimension of 2048 (the default configuration in T5X ). We keep all model architecture parameters constant across all experiments aside from the number of decoder layers. We use the 12-layer network as a baseline for comparison since it has the same number of total layers as the configurations with parameter sharing. We pretrain each network from scratch on C4 , and we finetune networks on translation and summarization tasks . For configurations using parameter sharing, parameter sharing is incorporated throughout pretraining and finetuning. Additional training details are provided in Appendix C.

2 Main Results

Figure 3 shows the accuracy versus efficiency tradeoff comparisons. When employing parameter sharing with SPEED, we observe significant speedups relative to the baseline 12-layer decoder network, achieving close to the same runtime as the shorter decoder network without parameter sharing across all benchmarks other than WMT-ENDE. Additionally, our parameter-sharing configurations attain significantly higher accuracy than the shallow decoder baselines. This demonstrates how the SPEED approach allows for improving accuracy for a fixed model size without a significant runtime penalty. We further experiment with deepening the decoder with parameter sharing by sharing parameters more times such that the total number of layers is increased. We find that deepening the decoder generally improves accuracy with minimal runtime penalty; as such, we believe that this is a promising approach to further boost accuracy for a fixed model size without much latency overhead. Appendix D.1 provides analysis for the accuracy of predictions made at early layers with SPEED. A detailed analysis of performance implications of the SPEED approach (and analysis of the lesser speedups we observe for WMT-ENDE) is provided in Appendix D.

Conclusion

We present a novel decoding strategy that allows for pipelined execution in Transformer decoders with parameter sharing. We describe the modifications required to the model architecture to leverage pipelined execution to reduce memory traffic (namely, cyclic parameter sharing in the decoder module). We observe consistent accuracy gains across all tasks for an equivalent model size, with only a small latency penalty. These results demonstrate the accuracy and performance benefits of our pipelined inference approach, showing how SPEED allows for deeper decoder configurations with parameter sharing to improve accuracy for a fixed parameter budget and minimal latency penalty.

References

Appendix A Related Work

Prior works have explored parameter sharing in encoder-only , encoder-decoder , and decoder-only Transformers as a method for reducing the size of the network by sharing parameters across all layers in the encoder and/or all layers in the decoder. explored only sharing parameters amongst a subset of layers and found that cyclic parameter sharing schemes outperformed sharing across all layers. There have also been several works on speculative decoding which aim to produce a set of “draft” tokens autoregressively using a smaller network and then correct them (in parallel) using a larger network . Our work instead aims to support speculative execution within a single network in order to accelerate inference with parameter sharing networks.

There are also prior works that aim to accelerate decoder inference through early exiting, where inference is terminated early when the model is confident that it can already predict the next token . Our work also leverages similar intuition, namely that while certain predictions truly benefit from the models’ full capacity, other continuations are more trivial and can be solved with reduced compute . However, our proposed approach for accelerating decoder inference has advantages over typical early exiting approaches. Although early exit can be applied to an existing network and doesn’t require pretraining, our method is guaranteed to always achieve the same accuracy as the baseline network with parameter sharing since it fixes any mistakes. Our method also reduces the model size through parameter sharing (in addition to the speedup from pipelined execution).

Appendix B Algorithm

Algorithm 1 provides the detailed outer-loop decoding algorithm for pipelining decoder inference. The “iteration_indices" variable is responsible for both tracking which pipeline stages have a valid token and also what iteration in the sequence these valid tokens are at. For example, if the model has 6 groups of shared decoder layers and three valid tokens (corresponding to iterations 3, 2, and 1 in the sequence) which are entering layers 1, 3, and 5 in the network, “iteration_indices" will be equal to (3, -1, 2, -1, 1, -1).

Appendix C Training Details

We used the default SentencePiece tokenizer with a vocabulary size of 32K, and we used tied input and output embeddings . We pretrained on C4 (Colossal Clean Crawled Corpus) for 524,288 steps using a batch size of 128 . C4 is a large dataset of filtered English text scraped from the web . We used a base learning rate of 1 with a square root decaying learning rate scheduler and with 10K warmup steps. We focused on two particular sequence-to-sequence tasks during finetuning: translation and summarization. For translation, we used the WMT English to German dataset as well as the English to German Paracrawl-Paragraph translation dataset . The Paracrawl-Paragraph dataset was used to also evaluate on a translation dataset with longer source and target context lengths (since it consists of full paragraph translations). For summarization, we used the CNN/DailyMail dataset, which consists of news articles written by journalists at CNN and the Daily Mail , as well as the English to English split of the Wikilingua multilingual summarization dataset . We finetuned for 262,144 steps using a batch size of 128 and dropout of 0.1, using input/target sequence lengths of 512/512 across all tasks. When finetuning, we used a constant learning rate of 0.001 with 1K warmup steps. We evaluated checkpoints every 5,000 steps during finetuning on the validation set, and then reported results on the test set using the checkpoint with the best accuracy on the validation set. Both training and inference arithmetic were performed in BF16 precision. We used TPU v2-8 machines on Google Cloud Platform for training experiments, and we launched these experiments using Skypilot .

Appendix D Performance Analysis

In order to assess the accuracy of the predictions made from our network at earlier layers, we profiled the proportion of predictions that were flipped between pairs of layers during inference. Figure 4 shows the prediction consistencies for 4x3 and 2x6 network configurations across all tasks. Upon examining these numbers and plots, we found that the model is able to make the majority of predictions accurately at early layers. Across all three configurations, the proportion of predictions that would need to be corrected after the first layer was between 13-17% for the 2x6 configuration and between 6-14% for the 4x3 configuration, showing that the majority of predictions were correct at early layers. Additionally, we found that a very small percentage of predictions flipped at later layers. This shows that the model tends to converge to the final answer and does not experience much oscillation between different predictions.

D.2 WMT-ENDE Performance

The primary reason that we observed greater latency penalties with our approach for WMT-ENDE compared with the other translation and summarization tasks was due to its shorter output generation lengths. The benefits from our pipelined decoding approach come from being able to process multiple tokens in parallel, and in the first few and last few iterations with our method, there will be fewer tokens in the pipeline. This means that the first iterations and final iterations in pipelined decoding aren’t completely overlapped. This is only a limiting factor for tasks with shorter generated sequence lengths (where the average number of tokens generated is close to the number of shared groups of layers in the network). The generation lengths for WMT-ENDE are typically shorter than summarization tasks and paragraph-level translation, which leads to increased latency penalties.

D.3 CNN/DM Performance

With CNN/DM, we actually observed reduced latency when inferring the 4x3/4x4 configurations relative to the 4-layer network without parameter sharing. This is unexpected, since even assuming perfect prediction for the networks with parameter sharing, the latency would not be less than the baseline 4-layer network (assuming the same output generation length). However, it is possible for the parameter sharing configurations to exhibit lower latency due to differences in the average generation lengths for the networks with parameter sharing relative to the network without parameter sharing (as if the average generation length is shorter for the networks with parameter sharing, they could have lower average latency).

D.4 General Discussion

There are several factors which impact the runtime when employing SPEED.

One factor is the generation length, as the first iterations and final iterations in pipelined decoding aren’t completely overlapped (as discussed in Appendix D.2). This limits the runtime gains from SPEED for tasks with short output generation lengths.

Because the embedding matrix is large, it can actually end up consuming a large portion of the memory bandwidth (and hence the runtime) for smaller models. This is a crucial reason why the latency gains aren’t linear as you go from a 12-layer network down to a 2-layer network even without considering parameter sharing or speculative execution.

One additional performance implication is that if the pipeline is too deep (i.e. layers are shared too many times), this can lead to greater misprediction penalties. Improving prediction consistency is therefore crucial for improving runtime with deeper decoder configurations.