Memory-efficient Transformers via Top-$k$ Attention
Ankit Gupta, Guy Dar, Shaya Goodman, David Ciprut, Jonathan Berant
Introduction
The Transformer architecture has been successful in a wide range of natural language processing tasks, including machine translation , language modeling , question-answering , and many more. Transformers pre-trained on large amounts of text with a language modeling (LM) objective, have become the standard in NLP, exhibiting surprising amounts of linguistic and world knowledge .
Compared to prior methods, top- attention has multiple attractive properties:
Top- attention has the same memory footprint as Performer , a state-of-the-art attention variant with linear time and memory complexity, on very long inputs (orange curve, Fig. 1, top-right), while being as fast as vanilla attention, and even faster than linear variants on inputs of length up to (Figure 1, bottom-left). This allows us, e.g., to train a typical 12-layer Transformer decoder over -long inputs on a GiB GPU (Figure 3a).
Top- attention is a highly accurate approximation to vanilla attention and is a plug-and-play replacement at both multi-head attention and feed-forward layers of a Transformer. This is unlike past attention variants that require an expensive corrective pre-training stage to adjust model weights to the new variant, which can be prohibitive for large models. We show top- attention can replace vanilla attention in a zero-shot inference setup and at fine-tuning time without any corrective pre-training.
We extensively evaluate top- attention on a wide range of tasks and demonstrate its mentioned advantages. Training from scratch, we show top- attention performs as well as vanilla self-attention on Long Range Arena, a benchmark dedicated to evaluating the ability of transformers to handle long sequences, and in a language modeling task (WikiText-103). Second, we show top- attention can be used as a drop-in replacement for vanilla attention at inference time without any additional training at the feed-forward layer of the UnifiedQA model on 12 different question answering (QA) datasets, reducing the number of keys used per query by more than 99%. Last, we show top- attention obtains similar performance to vanilla attention on a wide range of QA tasks when fine-tuning T5 , without the need for any corrective pre-training.
Overall, our results demonstrate that top- attention is a simple and effective method for dramatically reducing the memory footprint of Transformers without loss of accuracy that can allow resource-constrained researchers enjoy the benefits of large pre-trained Transformer-based models. Our code is available at https://github.com/ag1988/top_k_attention.
Efficient Transformer through Top-kk Attention
In this section, we briefly review the Transformer architecture, its sparse approximations, and show how to cast the feed-forward layer into the query-key-value framework (§2.1). We then describe top- attention and our memory-efficient implementation for it (§2.2).
A Transformer is a stack of layers each consisting of multi-head attention and feed-forward sub-layers. Its contextualizing component is the multi-head attention defined as follows.
Sparse approximations
The attention function (Eq. 1) requires the computation of containing entries and can be expensive for long sequences ( and are typically the sequence length). To alleviate this issue, sparse attention variants relax this requirement and compute only a few entries of , masking out the rest. For a binary mask ,
The sparsity of can be leveraged via customized implementations of matrix product and, thus Eq. 2 can be significantly cheaper to compute compared to Eq. 1.
Feed-forward as attention
In the feed-forward layer, a -hidden layer fully-connected network is applied identically to every input token. As observed in past work , a feed-forward layer can be cast into the query-key-value framework as:
2 Top-kk Attention
In this work we propose top- attention, where for each query, we mask out all but its largest dot products with the keys, that is, in each row of we only keep its largest elements and mask out the rest:
A simple way to implement attention and reduce its peak memory consumption is to chunk queries: instead of processing all the queries at once, we partition the queries into chunks and process them sequentially, one chunk at a time. For a chunk size , the rows of are grouped into contiguous chunks of size and the attention function (Eq. 1, 3, 4) is computed using as inputs where denotes the subset of corresponding to chunk .
During inference, once a query chunk is fully processed, the intermediate activations produced during its processing can be discarded and, hence, the peak memory required to process all queries is bounded by the memory required to process a single chunk. Therefore, modulo the storage required for and the outputs themselves, the peak memory usage reduces from to which is linear with respect to for a fixed chunk size .
Chunk size provides a simple way to trade-off between the maximum memory usage and the slowdown due to the sequential processing of chunks. Fig. 2 shows memory and time for different chunk sizes for a single BERT-base self-attention layer over a sequence of length . We observe that chunk sizes yield a good trade-off between time and memory.
Input checkpointing
Taking inspiration from gradient checkpointing , we observe that if the inputs are available during the backward pass, we can re-compute and then use the produced intermediate activations to compute from . Once is computed, we can again discard the intermediate activations and gradients produced during this step and move on to the next chunk. This ensures that the peak memory usage during the backward pass through the attention layer is bounded by the memory required to backpropagate through a single chunk.
To summarize, a customized backward pass allows us to utilize query chunking, both during forward and backward passes, and only requires us to cache the inputs to the attention function. For a stack of attention layers and fixed , this reduces the peak memory usage from to .
As described above, the combination of query chunking and input checkpointing provides a simple method for reducing the memory-footprint of vanilla attention, independent of top- attention. Indeed, our benchmarking experiments in §3 demonstrate this. However, a drawback of this approach is that, during the backward pass, an implicit second forward pass is performed to re-compute the intermediate activations as described above. This can potentially increase the compute (FLOPs) required for a combined forward and backward pass by . We now describe how to further improve both compute and memory by combining query chunking and input checkpointing with top- attention.
Improving efficiency through top-kk attention
Benchmarking
In this section, we benchmark top- attention in terms of time and memory, and compare it to vanilla attention, query-chunking without the top- operation, and to Performer , as a representative of state-of-the-art linear attention variants. We separately benchmark (a) a single self-attention layer over long sequences, (b) a single feed-forward layer with a large feed-forward dimension, and (c) a -layer Transformer decoder with same architecture as BERT-base .
For all models, we benchmark by running a forward and backward pass over random inputs. Each measurement is an average over 3 runs on an Nvidia A100 GPU and is discarded if memory usage exceeds GiB. We use causal masking for self-attention layers to highlight the simplicity of our approach that can seamlessly handle arbitrary attention masks, unlike other methods , where implementing causal masking requires customized CUDA implementations. For Performer, we use 256 random features, and the CUDA implementation from .
Multi-head attention layer: We benchmark a single multi-head attention layer over long sequences in a configuration similar to BERT-base: is , heads of size , and feed-forward dimension . Fig. 1 shows the results when setting to and the query chunk size to , which was shown to provide a good time-memory trade-off in §2.2.
We observe that top- attention has the same device-memory usage as the Performer (top) for sequences as long as tokens, while being as fast as vanilla attention, and even faster than Performer on inputs of length up to . With vanilla attention, we cannot fit even a single multi-head attention layer over a sequence of more than tokens, while top- uses less than GiB of memory over sequences of length . Lastly, we observe improvement in both time and memory when comparing top- attention to query chunking over vanilla attention, where using top- leads to a memory reduction for sequences of length .
Feed-forward layer : While considerable effort has been dedicated to devising efficient models for long contexts, a large feed-forward dimension is useful for knowledge-intensive tasks such as open-domain QA , and efforts have been made to reduce its complexity . We benchmark the resource usage of top- attention at a single feed-forward layer for different feed-forward dimensions using batch size and input length , which results in queries per batch.
Top- attention (Figure 3b), for and query chunk size , dramatically improves device-memory usage compared to vanilla attention: it allowed us to use a feed-forward dimension with GiB, while vanilla attention uses the same amount of memory with a feed-forward dimension . Fitting a linear curve to the memory usage of vanilla attention and top- attention, we estimate that top- attention can handle feed-forward dimension compared to for vanilla attention on a GiB machine. Moreover, comparing top- attention to query chunking, we again observe a improvement in memory usage when the number of keys is . Lastly, we observe only a minor slowdown in top- attention compared to vanilla attention.
12-layer model: We benchmark a -layer model to examine the cumulative utility of not caching in all layers compared to the Performer. We use the same architecture as BERT-base with batch size and vary the input length. We use a Transformer decoder with top- attention and chunk size at the self-attention layers, and simple query chunking with chunk size at the feed-forward layers.
We easily fit a -long input on a GiB GPU, improving memory consumption by more than compared to vanilla Transformer and compared to Performer. Moreover, top- attention outperforms query chunking in terms of both memory and runtime. As top- attention targets memory consumption but not runtime, a current limitation is that runtime, unlike Performer, is still quadratic. Thus, running multi-layer models on long sequences is reasonable in a fine-tuning or zero-shot inference setup, but further work is required for training from scratch 12-layer models over large datasets that contain long sequences.
Overall, our benchmarking results over multi-head attention, feed-forward, and multi-layer Transformer establish top- attention as a strong baseline for future work on efficient Transformers that dramatically improves memory consumption. Next, we evaluate top- attention on downstream tasks and show that top- attention can be used as a drop-in replacement for vanilla attention without additional pre-training, which can allow resource-constrained research groups experiment with Transformers over long sequences or models with a large feed-forward dimension.
Experimental Evaluation of Top-kk Attention
Having established top- attention as a memory efficient alternative to vanilla attention, we now show that, even for small values of , top- attention provides a high-quality approximation of vanilla attention, both at the multi-head attention and feed-forward layers. We empirically show this in a wide range of setups including (a) training from scratch on tasks that require handling long-range dependencies (§4.1) and on language modeling (§4.2), (b) fine-tuning pre-trained language models (T5) on multiple QA datasets (§4.5), and (c) performing zero-shot inference using pre-trained language models (UnifiedQA) without any training (§4.3).
Long Range Arena is a recently established benchmark for evaluating the ability of Transformer variants to handle long sequences. It comprises of multiple text classification tasks with inputs containing thousands of tokens (Table 1). In ListOps , given a sequence of operations on single-digit integers, the model predicts a single-digit solution modeled as -way classification. IMDb movie reviews is a character-level binary sentiment classification task. Lastly, in the ACL Anthology Network (AAN) task, a character-level model classifies if there is a citation between a pair of papers.
For each task, we downloaded and directly used the vanilla Transformer code offered by the authors and compared the performance before and after replacing the multi-head attention layers with top- attention, using identical hyperparameters for both cases (details in §A.1). https://github.com/google-research/long-range-arena
Test accuracy measured at the training checkpoint with the highest accuracy on the development set is reported in Table 1 and the learning curves on the development and test sets are shown in Fig. 4. On IMDb and AAN, the performance of top- is comparable or better than vanilla attention. For ListOps, there is a minor drop in performance (1.5 points), but learning curves (Figure 4a) exhibit similar behaviour.
Thus, top- attention, even for as small as % of the number of keys, results in a performance very similar to that of vanilla attention. This shows that an exact and sparse top- solution is a high-quality approximation for vanilla attention at multi-head attention layers.
2 Language Modeling
We further ascertain the findings of §4.1 via language modeling on WikiText-103 using a -layer Transformer decoder with M parameters. Using an input length of , we trained two models with vanilla and top- attentions at the self-attention layers, obtaining test perplexity scores of and respectively, slightly better in case of top- (details in §A.3).
3 Zero-shot Inference with UnifiedQA
We have established that the performance of top- attention is comparable to vanilla attention when training the model from scratch. In this set-up, several recently-proposed approaches have also reported competitive performances . Now, we consider a different and more practical setup, where the starting point is using an already pre-trained language model . As such models were trained using vanilla attention, replacing it with a new attention variant typically requires a corrective pre-training stage to allow the model weights to adjust to the new variant, which can be expensive for large models. For example, have shown that using random features without corrective pre-training leads to high error rates in a language modeling task. Moreover, as explained in §2.1, most past methods are incompatible with feed-forward layers. In the subsequent experiments we show that it is possible to replace vanilla with top- attention, at multi-head attention and feed-forward layers, and perform inference and fine-tuning without any need for such correction.
First, we compare the performance of UnifiedQA before and after replacing its feed-forward layers with our implementation of top- attention and directly performing inference on 12 different question answering (QA) datasets without any training. UnifiedQA is a T5-based model with B parameters , fine-tuned on a weighted mixture of QA datasets. The 12 datasets include diverse domains, such as science questions, factoid questions over Wikipedia, commonsense questions, etc. Details regarding the datasets and metrics can be found in §A.2.
Table 2 shows the results for increasing values of , where the feed-forward dimension of the model is . We observe that already when and , i.e., less than % of the number of keys, performance is comparable to vanilla Transformer. When (% of the number of keys), performance is equal or better than vanilla Transformer on all tasks. This highlights the plug-and-play property of top- attention, which can be used without any additional training.
4 Zero-shot Inference with BERT
To verify that the plug-and-play property of top- attention also holds at self-attention layers, we downloaded a BERT-large-uncased-whole-word-masking checkpoint already fine-tuned on SQuAD v1 and evaluated its performance on the development set before and after replacing its self-attention layers with top- attention. For as low as (% of input length), we only saw a minor decrease in the exact match scores (). Moreover, to empirically verify that dense approximations of vanilla attention (Performer, RFA, etc) indeed require corrective pre-training, we repeated the measurement using Performer attention with features, obtaining a score of .
5 T5 Finetuning
Having established the plug-and-play property of top- attention in zero-shot inference (§4.3, §4.4), we now show the effectiveness of top- attention when fine-tuning a model, and that there are no unforeseen issues stemming from training under high sprasity. Here, we use T5-base rather than T5-11B and evaluate on five QA datasets (and not 12) due to computational constraints.
Similar to §4.3, we replace the feed-forward layers of T5-base, which has feed-forward dimension , with our implementation of top- attention and fine-tuned on multiple QA datasets. As summarized in Table 3, we found that the performance of top- attention was again comparable to vanilla attention on BoolQ, CommonsenseQA and ROPES with a minor loss in performance on MCTest () and OpenbookQA ().
To summarize, our experiments in §4.1-§4.5 demonstrated that the performance of vanilla attention and top- attention is comparable at both multi-head attention (§4.1, §4.4) and feed-forward layers in multiple set-ups including training from scratch (§4.1, §4.2), fine-tuning (§4.5) and zero-shot inference (§4.3, §4.4), while dramatically improving memory usage, as shown in §3.
Discussion
Related work Our work follows a long line of works on efficient Transformers (see §1). Our method employs three main ideas: (a) computing the top- attention scores for each query (b) grouping the queries into chunks and processing these sequentially (c) caching only a part of the activations for the backward pass. Top- operation was used at self-attention layers by to show improved model performance, attributed to the removal of irrelevant information in the context. We use it to reduce the resource usage of multi-head attention and feed-forward layers. Processing query chunks sequentially was also used in Reformer as activations are not cached. But in that case, by replacing vanilla residual connections in the Transformer with reversible connections . Similar to the explanation provided in §2.2, these require an extra implicit forward pass during the backward pass and do not provide the compute and memory savings we get from our top- specific backward pass (§2.2). Secondly, replacing residual connections with reversible ones changes the function computed by the model and would require corrective pre-training to be used with BERT, T5, etc (§4.3-§4.5).
Limitations and future work As our method requires computing inner products of all queries and keys, it has a quadratic compute requirement. As seen in our pseudo-code (§2.2), there are four matrix products (Lines 8, 15, 21, 25) involving a large sparse matrix and a small dense one. Our current implementation does not leverage this sparsity and hence is as slow as vanilla attention. While future devices might allow faster sparse-dense products, in the immediate future, one can leverage block-sparse kernels which have been successfully used for such products .
Conclusion In this work, we proposed a memory-efficient and accurate sparse approximation of the primary sub-layers of a Transformer, benchmarked the resulting resource savings, and verified its quality and unique advantages, on a wide range of downstream tasks and evaluation set-ups.
Acknowledgments and Disclosure of Funding
We thank Achille Fokoue, Maria Chang and Avi Sil for helpful discussions. Most of our experiments were conducted on IBM’s Cognitive Computing Cluster with additional resources from Hybrid Cloud Infrastructure Research. This research was partially supported by (1) The Yandex Initiative for Machine Learning, and the European Research Council (ERC) under the European Union Horizons 2020 research and innovation programme (grant ERC DELPHI 802800), and (2) the IBM AI Residency program.
References
Appendix A Supplemental Material
ListOps aims to diagnose the capability of modelling hierarchically structured data. Given a sequence of operations on single-digit integers, the model predicts the solution, also a single-digit integer modeled as a -way classification. Character-level text classification with the IMDb movie review dataset is a binary sentiment classification task. In the character-level document retrieval with the ACL Anthology Network (AAN) , the model classifies if there is a citation between a pair of papers.
We used the code and pre-processed data provided by the authors of Long Range Arena and default model configurations. For each task, we used identical hyperparameters for vanilla and top- attentions (Table 4) and used at most two Nvidia A100 for each run.
Notation BSZ: effective batch size, SQL: input sequence length, LR: learning rate, WRM: linear LR warm-up steps, STEP: number of gradient updates, EFQ: evaluated every these many steps, NL: number of layers in encoder/decoder, HS: hidden size, FF: feed-forward dimension, NH: number of heads, VOC: vocabulary size, DRP: dropout rate, CLIP: maximum gradient norm.
A.2 Details of UnifiedQA inference & T5 finetuning
We used Hugging Face’s Transformers library for these experiments. Authors of UnifiedQA collected and pre-processed several QA datasets into a common format: “QUESTION \n CHOICES \n CONTEXT’’. We downloaded this data by following the instructions provided by the authors https://github.com/allenai/unifiedqa and used it for the UnifiedQA inference experiments (§4.3). Some statistics are shown in Table 5. Longer inputs were truncated to 512 tokens.
Given an instance from the pre-processed data, we computed the exact match score of a prediction with respect to the list of provided answers via the SQuAD v1 evaluation script .
For the T5 experiments (§4.5), we used a slightly different input format. Given an instance in the UnifiedQA format, we formed the modified instance as “question: QUESTION context:
A.3 Language Modeling on WikiText-103
WikiText- is a language modeling task based on English Wikipedia. We used the language modeling framework provided by Faiseq https://github.com/pytorch/fairseq and hyperparameters in Table 7. The details of Adam optimizer are =, =, weight-decay: , CLIP: none, LR schedule: inverse square root. During evaluation on test set, dataset is chunked into segments of length and perplexity is computed over each segment normally without access to other segments.
A.4 Benchmarking details
Benchmarking (§3) was done in PyTorch 1.8.1. For each run, we sampled a batch of random 32-bit input vectors and a backward pass was performed using the mean of the output elements as the loss. The part of code that was timed was enclosed within torch.cuda.synchronize() to ensure all CUDA threads finished. Memory usage was measured using torch.cuda.max_memory_reserved(). On Nvidia A100, any internal casting to TF32 was explicitly disabled.
We considered the option of performing matrix products involving large sparse matrices (Lines 8, 15, 21, 25 in our pseudo-code (§2.2)) by representing them in torch.sparse_coo_tensor format and using the torch.sparse framework to explicitly leverage the sparsity. Unfortunately, we could not obtain encouraging results even for 1% of number of keys (Figure 5) and plan on experimenting with block-sparse kernels in the near future.