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-kk attention has multiple attractive properties:

Top-kk 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 4K4K (Figure 1, bottom-left). This allows us, e.g., to train a typical 12-layer Transformer decoder over 32K32K-long inputs on a 3030GiB GPU (Figure 3a).

Top-kk 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-kk 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-kk attention on a wide range of tasks and demonstrate its mentioned advantages. Training from scratch, we show top-kk 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-kk 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-kk 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-kk 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-kk 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 QK⊤QK^{\top} containing LQ⋅LKL_{Q}\cdot L_{K} entries and can be expensive for long sequences (LQL_{Q} and LKL_{K} are typically the sequence length). To alleviate this issue, sparse attention variants relax this requirement and compute only a few entries of QK⊤QK^{\top}, masking out the rest. For a binary mask B∈{0,−∞}LQ×LKB\in\{0,-\infty\}^{L_{Q}\times L_{K}},

The sparsity of BB 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 11-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-kk attention, where for each query, we mask out all but its kk largest dot products with the keys, that is, in each row of QK⊤QK^{\top} we only keep its kk 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 CC, the rows of QQ are grouped into LQ/CL_{Q}/C contiguous chunks of size CC and the attention function (Eq. 1, 3, 4) is computed using QC,K,VQ_{\mathcal{C}},K,V as inputs where QCQ_{\mathcal{C}} denotes the subset of QQ corresponding to chunk C\mathcal{C}.

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 LQL_{Q} queries is bounded by the memory required to process a single chunk. Therefore, modulo the storage required for Q,K,VQ,K,V and the outputs themselves, the peak memory usage reduces from Ω(LQ⋅LK)\Omega\left(L_{Q}\cdot L_{K}\right) to O(C⋅LK)O(C\cdot L_{K}) which is linear with respect to LKL_{K} for a fixed chunk size CC.

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 65,53665,536. We observe that chunk sizes 29,2102^{9},2^{10} yield a good trade-off between time and memory.

Input checkpointing

Taking inspiration from gradient checkpointing , we observe that if the inputs QC,K,VQ_{\mathcal{C}},K,V are available during the backward pass, we can re-compute oCo_{\mathcal{C}} and then use the produced intermediate activations to compute d(QC)d\left(Q_{\mathcal{C}}\right) from d(oC)d\left(o_{\mathcal{C}}\right). Once d(QC)d\left(Q_{\mathcal{C}}\right) 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 NN attention layers and fixed dd, this reduces the peak memory usage from Ω(LQ⋅LK⋅N)\Omega\left(L_{Q}\cdot L_{K}\cdot N\right) to O((LQ+LK)⋅N+C⋅LK)O\left((L_{Q}+L_{K})\cdot N+C\cdot L_{K}\right).

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-kk 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 50%50\%. We now describe how to further improve both compute and memory by combining query chunking and input checkpointing with top-kk attention.

Improving efficiency through top-kk attention

Benchmarking

In this section, we benchmark top-kk attention in terms of time and memory, and compare it to vanilla attention, query-chunking without the top-kk 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 1212-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 3030GiB. 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: dmodeld_{\text{model}} is 768768, 1212 heads of size 6464, and feed-forward dimension 30723072. Fig. 1 shows the results when setting kk to 128128 and the query chunk size to 10241024, which was shown to provide a good time-memory trade-off in §2.2.

We observe that top-kk attention has the same device-memory usage as the Performer (top) for sequences as long as 65K65K tokens, while being as fast as vanilla attention, and even faster than Performer on inputs of length up to 4K4K. With vanilla attention, we cannot fit even a single multi-head attention layer over a sequence of more than 10K10K tokens, while top-kk uses less than 1010GiB of memory over sequences of length 65K65K. Lastly, we observe improvement in both time and memory when comparing top-kk attention to query chunking over vanilla attention, where using top-kk leads to a 3×3\times memory reduction for sequences of length 65K65K.

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-kk attention at a single feed-forward layer for different feed-forward dimensions using batch size 512512 and input length 512512, which results in 2182^{18} queries per batch.

Top-kk attention (Figure 3b), for k=512k=512 and query chunk size 2142^{14}, dramatically improves device-memory usage compared to vanilla attention: it allowed us to use a feed-forward dimension 65K65K with 1111GiB, while vanilla attention uses the same amount of memory with a feed-forward dimension 2K2K. Fitting a linear curve to the memory usage of vanilla attention and top-kk attention, we estimate that top-kk attention can handle feed-forward dimension 205K205K compared to 7K7K for vanilla attention on a 3030GiB machine. Moreover, comparing top-kk attention to query chunking, we again observe a 3×3\times improvement in memory usage when the number of keys is 65K65K. Lastly, we observe only a minor slowdown in top-kk attention compared to vanilla attention.

12-layer model: We benchmark a 1212-layer model to examine the cumulative utility of not caching QK⊤QK^{\top} in all NN layers compared to the Performer. We use the same architecture as BERT-base with batch size 11 and vary the input length. We use a Transformer decoder with top-6464 attention and chunk size 1,0241,024 at the self-attention layers, and simple query chunking with chunk size 4,0964,096 at the feed-forward layers.

We easily fit a 32K32K-long input on a 3030GiB GPU, improving memory consumption by more than 8×8\times compared to vanilla Transformer and 2×2\times compared to Performer. Moreover, top-kk attention outperforms query chunking in terms of both memory and runtime. As top-kk 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-kk attention as a strong baseline for future work on efficient Transformers that dramatically improves memory consumption. Next, we evaluate top-kk attention on downstream tasks and show that top-kk 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-kk attention as a memory efficient alternative to vanilla attention, we now show that, even for small values of kk, top-kk 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 1010-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-128128 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-128128 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-kk attention, even for kk as small as 33% of the number of keys, results in a performance very similar to that of vanilla attention. This shows that an exact and sparse top-kk 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 66-layer Transformer decoder with 156156M parameters. Using an input length of 10241024, we trained two models with vanilla and top-6464 attentions at the self-attention layers, obtaining test perplexity scores of 30.9630.96 and 30.5130.51 respectively, slightly better in case of top-6464 (details in §A.3).

3 Zero-shot Inference with UnifiedQA

We have established that the performance of top-kk 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-kk 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-kk attention and directly performing inference on 12 different question answering (QA) datasets without any training. UnifiedQA is a T5-based model with 1111B 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 kk, where the feed-forward dimension of the model is 65,53665,536. We observe that already when k=256k=256 and k=512k=512, i.e., less than 11% of the number of keys, performance is comparable to vanilla Transformer. When k=4,096k=4,096 (66% 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-kk 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-kk 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-kk attention. For kk as low as 1616 (44% of input length), we only saw a minor decrease in the exact match scores (86.9→86.286.9\rightarrow 86.2). 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 256256 features, obtaining a score of 0.380.38.

5 T5 Finetuning

Having established the plug-and-play property of top-kk attention in zero-shot inference (§4.3, §4.4), we now show the effectiveness of top-kk 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 30723072, with our implementation of top-256256 attention and fine-tuned on multiple QA datasets. As summarized in Table 3, we found that the performance of top-256256 attention was again comparable to vanilla attention on BoolQ, CommonsenseQA and ROPES with a minor loss in performance on MCTest (81.2→79.481.2\rightarrow 79.4) and OpenbookQA (58.8→58.058.8\rightarrow 58.0).

To summarize, our experiments in §4.1-§4.5 demonstrated that the performance of vanilla attention and top-kk 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-kk 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-kk 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-kk 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 1010-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-kk 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: CHOICES \n CONTEXT”.

A.3 Language Modeling on WikiText-103

WikiText-103103 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 β1\beta_{1}=0.90.9, β2\beta_{2}=0.980.98, weight-decay: 0.010.01, CLIP: none, LR schedule: inverse square root. During evaluation on test set, dataset is chunked into segments of length 10241024 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 k=k= 1% of number of keys (Figure 5) and plan on experimenting with block-sparse kernels in the near future.