GMAT: Global Memory Augmentation for Transformers

Ankit Gupta, Jonathan Berant

Introduction

The Transformer architecture has been widely successful in achieving state-of-the-art performance on a wide range of natural language processing (NLP) tasks, including machine translation , language modeling , question-answering , and many more. In particular, Transformers pre-trained on large amounts of text with a language modeling (LM) objective, have become the de-facto standard in NLP, exhibiting surprising amounts of linguistic and world knowledge . Moreover, Transformers with tailored attention patterns have been successfully used to replace convolutions in computer vision , and have also been useful in music generation , symbolic mathematics and other modalities.

One of the most powerful features in Transformers is (pairwise) self-attention, where all positions in an input sequence aggregate information from the entire sequence in parallel. However, this requires computing a similarity score for all pairs of positions simultaneously, leading to a Ω(L2)\Omega(L^{2}) memory requirement for length LL sequences, which is prohibitively expensive for long sequences. To alleviate this issue, several sparsifications of vanilla self-attention have been recently proposed; each restricting the number of positions that a given position can attend to . For example, in BlockBERT , the sequence is split into LM\frac{L}{M} chunks of length MM and positions in chunk ii only attend to positions in chunk σ(i)\sigma(i) for some pre-determined permutation σ\sigma, thereby having a O(M⋅L)O(M\cdot L) memory requirement. The Reformer uses locality-sensitive hashing (LSH) to arrange similar vectors close to one another and then chunks them. Each chunk then attends to only a couple of chunks leading to a O(M⋅L)O(M\cdot L) memory requirement. While such sparsifications often lead to performance that is comparable to vanilla Transformers, they have some undesirable consequences:

A position can require many layers to accumulate information from the entire input, and thus struggle when aggregating global information is necessary. For example, in §3.1, we show that a LSH Transformer (Reformer without reversible connections) struggles on a simple tagging task that requires information from the entire sequence.

Most sparsification schemes pose an inductive bias based on the locality of natural language, by restricting a position to only attend to its nearby tokens. While this is often a reasonable assumption, it raises several concerns. First, it is trivial to construct examples where locality is violated. For example, vanilla attention is invariant to input permutation, and thus can handle tasks such as “word deshuffling”, where a randomly shuffled sequence needs to be mapped to the original order. On the other hand, any locality-based inductive bias would be detrimental. Second, progress in natural language understanding has led to increasing interest in handling global dependencies in long documents and even entire books , where a locality-based inductive bias is sub-optimal.

In this work, we propose Global Memory Augmentation for Transformers (GMAT). We augment sparse variants of the Transformer with a small global memory which is read and updated by all the positions using vanilla attention. Specifically, we prefix every input sequence (of length LL) with a list of MM memory tokens. At each multi-head attention layer, for each head (Figure 1a), the LL tokensFor brevity, we use the word token to refer both to the input token and its contextualized representation interchangeably. of the main sequence attend to other tokens of the main sequence using any sparse variant of attention, whereas they attend to the MM memory tokens using vanilla dense attention. Moreover, the MM memory tokens attend to all M+LM+L tokens using vanilla attention. This results in a O(M⋅(L+M))O(M\cdot(L+M)) memory overhead which is manageable for MM ≪L\ll L. Because the number of parameters in Transformers does not depend on the length of the input (modulo learned positional embeddings), the number of parameters grows by only a negligible M⋅EM\cdot E parameters, for an embedding size EE.

We propose also to use GMAT for sequence compression (Figure 1c). After encoding an input sequence with NcN_{c} GMAT layers, we discard the vectors corresponding to the main sequence XX, and keep only the global memory vectors, which are now a compressed representation of the entire input. The memory vectors are then processed and decompressed using NdN_{d} layers back to the original input length. The sequence can now be stored using only M(≪L)M(\ll L) vectors, and decompression is done with a small number (Nd)(N_{d}) of GMAT layers.

We evaluate GMAT on a wide range of tasks and show: (a) large improvements on synthetic tasks where global reasoning is required, (b) it improves masked langauge modeling (MLM) accuracy (used in Transformer pre-training), (c) improvements on two reading comprehension (RC) tasks, and last (d) moderate reductions in MLM and RC performance when using GMAT for compression.

To summarize, GMAT is a simple extension of the Transformers architecture, that can be seamlessly combined with any sparse attention variant. We show GMAT is useful for reducing memory requirements as well as for sequence compression, and demonstrate performance enhancements on a wide range of tasks. Our code and data can be downloaded from https://github.com/ag1988/gmat.

Global-Memory Augmented Transformers

A Transformer is a stack of layers each consisting of sub-layers such as multi-head attention, feed-forward, etc. Its contextualizing component is the multi-head attention defined as follows.

In multi-head attention, instead of computing a single attention output with dmodeld_{\text{model}} dimensional keys, queries, and values, these are linearly projected down in parallel hh times to d=dmodel/hd=d_{\text{model}}/h dimensions, using different learned projection matrices. Attention is applied to each of the hh new queries, keys and values, yielding dd dimensional outputs which are concatenated and again projected to obtain the dmodeld_{\text{model}}-dimensional output.

The attention function (Eq. 1) requires the computation of QKTQK^{T} containing LQ⋅LKL_{Q}\cdot L_{K} entries and can be expensive for long sequences. To alleviate this issue, sparse attention variants relax this requirement and compute only a few entries of QKTQK^{T}, masking out the rest. For a binary maskThe sparsity of BB can be leveraged via customized implementations of matrix product . B∈{0,−∞}LQ×LKB\in\{0,-\infty\}^{L_{Q}\times L_{K}},

Global Memory

As explained in §1, sparse attention variants have some undesirable properties. To remedy this, we augment such models with a small global memory which is read and updated by all the positions using vanilla attention (Figure 1a). Specifically, we prefix every token sequence XX (of length LL) with a sequence of MM memory tokens [m1m_{1}],…, [mMm_{M}]. At each multi-head attention layer of the model, at each head, the LL representations XX corresponding to the tokens of the main sequence attend to the other positions in XX using any sparse attention variant, but attend to the representations of all memory tokens XMX_{M} normally (Eq. 3). Moreover, the memory tokens attend to all the M+LM+L tokens normally.

This results in a O(M⋅(L+M))O(M\cdot(L+M)) memory overhead (manageable for MM ≪L\ll L). Moreover, this does not add any parameters to the model, except for a negligible M⋅EM\cdot E parameters used to embed the MM new memory tokens with an embedding size of EE.

Chunked self-attention

To explore the limits of GMAT and highlight its ability to contextualize over multiple fragments via a memory, we work with chunked self-attention (Figure 1b), a simple sparsification method. In C×kC\times k chunked self-attention (Figure 1b), a sequence of length C⋅kC\cdot k is partitioned into kk contiguous chunks of length CC. Each token within a given chunk uses vanilla (multi-head) attention to attend to tokens in its chunk in addition to the global memory but does not attend to other chunks. Hence, chunks interact with each other only via the memory. Without memory, training with C×kC\times k attention is equivalent to training with vanilla Transformer over a length-CC sequence. While more complex sparsification schemes are possible, this setup focuses on the ability to contextualize disjoint segments through the global memory only. Note that a single model can be used with different values of CC and kk, as Transformer models are invariant to the length of their input, aside for positional embeddings, which we handle below. We use the notation (C×k,M)(C\times k,M) to denote a chunked self-attention model where the input sequence XX has length C⋅kC\cdot k with a global memory of size MM.

Positional Embeddings

As attention is invariant to order, it is important to supply the model with positional information corresponding to the individual tokens. Rather than have a distinct learnable vector for each position , we represent a position pp as a tuple (q,r)(q,r) where r=p (mod 512)r=p\ (\text{mod }512), and q=⌊p/512⌋q={\lfloor p/512\rfloor}. Each 0≤r<5120\leq r<512 and 0≤q<640\leq q<64 has a distinct learnable vector.These particular values allow us to initialize the vectors for rr with the learned 512 positional embeddings of pre-trained LMs such as BERT. The positional embedding of pp is represented by the sum of the vectors corresponding to qq and rr and allows us to model positions up to 2152^{15}. Memory tokens have a fixed position, and thus positional embeddings are used only for the main sequence XX.

1 Sequence Compression

Contextualized word representations improve performance compared to fixed word embeddings such as GloVe . Unfortunately, some large models that compute contextualized representations do not fit on popular GPUs and need specialized hardware . Instead of leaving this computation to the users, an appealing option might be to release pre-computed contextualized representations, at least for popular benchmarks, similar to word embeddings . However, storing a vector for each position in a large corpus is expensive. A second natural motivation for GMAT is for sequence compression, i.e. using a small memory to represent a long sequence. This can dramatically reduce the overhead in storing pre-computing contextualized representations for large corpora.

Consider a NN-layer GMAT with a memory of size MM. We apply the NN model layers in 3 steps as shown in Figure 1c. Given a length LL input sequence, let WW and PP denote its word and positional embeddings. First, the bottom NcN_{c} layers are applied for compressing the entire information of the input sequence into just MM memory vectors XM(c)X_{M}^{(c)}. The next NmN_{m} layers are then applied only on the MM memory vectors resulting in richer representations XM(m)X_{M}^{(m)}. This length MM sequence is restored to the original length M+LM+L by concatenating the positional embeddings PP and, finally, the remaining Nd=N−Nc−NmN_{d}=N-N_{c}-N_{m} layers are applied for decompressing the information packed in XM(m)X_{M}^{(m)} into the final representations. Here, the positional embeddings PP act as queries for restoring the original input information from XM(m)X_{M}^{(m)}.

The M(≪L)M(\ll L) vectors XM(m)X_{M}^{(m)} can be efficiently serialized for later use and are decompressed using minimal post-processing into representations of length LL. In §4.3 and §4.4, we show that using the decompressed representations leads to only a small performance degradation on masked language modeling and reading comprehension. We use positional embeddings PP in the decompression step and not the contextualized representation X(c)X^{(c)} (Figure 1c), since we want the output to depend only on MM vectors instead of M+LM+L.

In contemporaneous work , the authors demonstrate the effectiveness of a sparse sliding window attention pattern combined with a global memory on multiple natural understanding tasks. In contrast to our work, the contribution of the global memory component is not evaluated, and is designed on a task-by-task basis. Moreover, attention scores for the memory and the main sequence are computed using separate sets of projection matrices, thereby introducing new parameters. In comparison, we demonstrate the utility of a global memory for various sparse attention patterns on synthetic, NLP and compression tasks, introducing a minimal set of new parameters.

Global Reasoning on Synthetic Data

We consider two synthetic datasets, where global reasoning over the entire input is required.

Recently proposed sparse attention schemes exploit the locality of natural language (§1). One exception is locality sensitive hashing (LSH) attention . While in LSH attention, tokens are not bound to attend only within their proximity, it does limit the number of positions a token can attend to at each head. We examine the utility of adding GMAT to a LSH Transformer , that is, a Transformer that uses LSH attention, for a sequence tagging task.

2 Numerical Reasoning Over Text

Having shown the utility of GMAT on a combinatorial task, we now transition towards language tasks. We create a pseudo-textual task that requires global mathematical reasoning by generating examples from a rule-based generator from . The generator generates examples that include a passage PP, describing a sequence of related events, and a question QQ that requires aggregating information from various parts of PP and performing numerical reasoning to arrive at a numeric answer AA (see Table 2).

Following , we train a generative encoder-decoder model, GenBERT, on the generated data after replacing the encoder self-attention with chunked self-attention (§2), and compare the performance before and after GMAT augmentation (see training and data details in §A.1). We summarize the results in Table 3. Compared to vanilla attention over the entire input (i.e., (140×1,0)(140\times 1,0)), chunking the encoder input into 2 chunks significantly reduced the performance (i.e., (70×2,0)(70\times 2,0)), in line with the global nature of the task. Surprisingly, adding a global memory of size 30 reduced accuracy even further. We hypothesize this is due to the strict weight-tying strategy employed by GenBERT, where the parameters of the Transformer encoder and decoder are tied, leading to underfitting. To handle that, we untie the parameters of the projection matrices that update the memory representations XMX_{M} in all attention heads (Eq. 3, right), initializing them randomly. This separates the attention heads that update the main sequence from the heads that update the memory. In this setup, accuracy improved substantially, almost recovering the performance of vanilla attention.

Masked Language Modeling

One of the biggest success stories of Transformers is as an architecture for pre-training LMs. We now investigate pre-training GMAT with a masked language modeling objective, as a memory-efficient replacement for models such as BERT . Past work has shown strong correlation between performance on the MLM task and that on downstream applications . For our experiments, we use the BERT-base architecture after making the modifications described in §2.

We form examples by sampling sequences of length LL from English Wikipedia and the PG19 dataset, and replacing sub-words with the [MASK] token following the procedure in (details in §A.2). The model is trained to maximize the log probability of the masked out tokens. We evaluate the error of the model as the fraction of tokens predicted incorrectly, and the MLM “perplexity” as the reciprocal of the geometric mean of probabilities of all masked out tokens.Equivalently, the natural exponential of the average loss over the development set. PG19 contains 29K long books, and is thus likely to benefit from modeling long context, while in Wikipedia most articles are short and can fit into the 512 word pieces that are the input of BERT. We experiment with training a Transformer from scratch, as well as initializing with BERT-base.

As shown in Figure 3a, we train 3 models on Wikipedia. The setting (512×1,0)(512\times 1,0) corresponds to standard MLM training on instances of length 512 without global memory. Similarly, (1024×1,0)(1024\times 1,0) denotes training with vanilla attention over a context of size 1024, incurring a large memory penalty.We do not train in the (2048×1,0)(2048\times 1,0) setting due to memory constraints. Lastly, in (512×4,64)(512\times 4,64), a 2048-long context is chunked into four 512-chunks that have to interact via a global memory of size 64. Increasing the context size to 1024 improves MLM performance on Wikipedia (Table 4). Using global memory improves sample complexity and performance compared to training on 512-long instances, albeit only moderately. Thus, the chunks are able to leverage global memory to exchange information, alleviating the context fragmentation problem .

2 BERT Initialization

To show that GMAT can be easily integrated into existing pre-trained LMs, we take a pre-trained BERT-base model, and further train it using GMAT. Because BERT was pre-trained on Wikipedia, improving performance on Wikipedia itself could be difficult, as it already has high confidence on tokens from this corpus. Hence, we also train on PG19, which was not part of BERT’s training data.

Table 5 summarizes the results. On Wikipedia, increasing the context size to 1024 provides a significant improvement (Figure 3b), but global memory does not improve performance compared to standard MLM training on 512-long instances. However, on PG19 (Figure 3c) using global memory substantially improves perplexity from 4.4→4.354.4\rightarrow 4.35, closing roughly half the gap from using a context of size 10241024, which obtains an MLM perplexity of 4.34.3. This hints that the lack of improvement on the Wikipedia data might be due to the fact that BERT was pre-trained on Wikipedia.

The above results indicate that disjoint text segments can exchange useful information via global memory. However, because natural language has a locality bias, the utility of memory diminishes as the chunk length CC increases. To determine the efficacy of GMAT when CC is small, where contextualization should highly depend on the memory, we experiment with chunks of size 8. As expected, without access to a reasonably-large surrounding context, the model (8×64,0)(8\times 64,0) fails to predict masked tokens (Table 5). Interestingly, a small global memory of size 6464 significantly improves performance (53.11→32.9853.11\rightarrow 32.98 error, 20.68→4.9420.68\rightarrow 4.94 perplexity), reaching performance that is close to (512×4,0)(512\times 4,0). We further evaluate the pre-trained GMAT models on reading comprehension tasks in §4.4.

3 Sequence Compression

We turn to sequence compression, where our goal is to compress a sequence of length LL into MM vectors that can be saved and later decompressed back into a sequence of length LL, with minimal drop in performance. Using the setup described in §2.1, we use NcN_{c} compression layers, followed by Nd=N−NcN_{d}=N-N_{c} decompression layers, and train with the same data and MLM objective as above on Wikipedia. As shown in Table 6, we found that Nc=9N_{c}=9 outperforms Nc=3N_{c}=3 (which also happens to be well-aligned with our need for a small number of decompression layers). Compared to a model without compression, we observe a moderate degradation in performance (29.163→32.9829.163\rightarrow 32.98 error, and 3.953→5.0173.953\rightarrow 5.017 MLM perplexity), showing that a global memory of size just 64 provides a compact and useful representation for the entire sequence of length 512.

4 Reading Comprehension Performance

While MLM performance is known to correlate well with downstream applications , we take Wikipedia-based GMAT models trained with the MLM objective in §4.2 and §4.3 and further fine-tune them on reading comprehension (RC) tasks.

HotpotQA

To investigate the utility of GMAT for long-range reasoning, we fine-tune our models on HotpotQA , an RC dataset focusing on multi-paragraph reasoning. In HotpotQA, examples comprise of a question QQ, 2 gold paragraphs G1,G2G_{1},G_{2} required for answering QQ, and 8 distractor paragraphs. Each gold paragraph contains supporting facts: the sentences relevant for answering QQ.

Conclusion

In this work, we proposed GMAT, a simple extension to the Transformer architecture that allows a better trade-off between compute and performance and can be naturally used for sequence compression. Our approach can be seamlessly integrated with the increasingly-popular sparse Transformer variants. We show GMAT (a) leads to performance improvements on a wide range of tasks, and (b) can be used to compress long sequences by factor of 8×8\times with only a small degradation in performance.

Broader Impact

Transformers have become a popular architecture for sequence processing and generation in natural language processing and outside of it. The goal of this paper it to reduce the memory requirement and thereby allow for longer sequences to be processed. Moreover, our compression technique can facilitate the use of pre-computed contextualized representations, allowing users access to an approximation of these representations even if they cannot compute the representations from scratch themselves. As such, we consider a positive impact of this work to be the ability of more users with constraints on their computational resources to use the Transformer architecture and its pre-trained representations. Moreover, being able to process long documents can open the door to new applications in natural language processing, such as multiple-document understanding, and perhaps also processing of sequences outside of NLP, for example in Biology. As Transformers are becoming ubiquitous in machine learning, naturally any negative impact that can be attributed to Transformers (fake news generation, classifiers in sensitive domains such as the justice system and healthcare) are also inherited by our approach, and perhaps enhanced when long sequences need to be processed.

Acknowledgments and Disclosure of Funding

We thank Shimi Salant, Yoav Goldberg and Mor Geva for helpful discussions and constructive suggestions. This research was partially supported by The Israel Science Foundation grant 942/16, 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).

References

Appendix A Supplemental Material

For creating the textual synthetic data for the generative QA task of §3.2, we used the data generation set-up of (§4.2 of ). Default templates and vocabulary were used to create passages containing 5 sentences. While instantiating the templates, the probability of sampling from one of the previously used values was set to 0.999 to promote inter-sentence dependencies. This gave us 629906/15K train/dev passage-question-answer triples. Among these, we only kept the samples where the answer was a number not appearing in the passage, and discarded the rest. This gave us 223067 training and 5146 evaluation instances.

We only kept the decoder/generative head of GenBERT (§3 of ) and allowed the decoder to attend to all the encoder outputs in the cross-attention layers. As the weights of the encoder and decoder are tied, we used segment ids 0 for the encoder input sequence and 1 for the decoder inputs.

A.2 Data for Masked LM task

The instances for the MLM task (§4) were formed separately using 5.2M pages from English Wikipedia (October 2017 dump) and the training set of PG19 dataset containing ∼\sim29K books from Project Gutenberg . For each dataset, after appending a special symbol at the end of each document, the documents were arranged in a random order and concatenated into a single long text which was then tokenized into a list of tokens. Depending upon the input length LL of the experiment (512/1024/etc) this list was chunked into full length L−2L-2 sequences which were then masked randomly following and enclosed within [CLS] and [SEP] tokens. For each dataset, the first 2.55B tokens (i.e. 510×5510\times 5M) were used to form the respective training set, next 10.20M tokens (510×20510\times 20K) the dev set and the rest were discarded.

A.3 Finetuning on HotpotQA

Given the question QQ and 10 arranged paragraphs PiP_{i}’s, each PiP_{i} is extended by prefixing it with its title. Moreover, to handle yes/no questions, a special string yes no is also prefixed. The context DD is formed by simply concatenating the resulting paragraphs. Following , given a window/chunk PP from tokenized DD, the corresponding instance is formed as [CLS] Q [SEP] P [SEP].

Supporting Facts Tagging Task (SF): Besides the standard span extraction loss, we also include another task using the supporting facts supervision. Contextualized representations of the model are linearly projected to 2 scores (for 0/1) per token and normalized to obtain log-probabilities. For an input, loss is computed as negative log probability of the correct tag averaged over the positions. As supporting facts positions are fewer, log-probabilities are weighted according to the respective class (0/1) size.

A.4 Hyperparameters

For all our experiments, we used an older version of Hugging Face’s Transformers library . For convenience, we denote the training hyperparameters using the following abbreviations, INS: number of training instances, BSZ: number of instances in a batch, ISZ: instance size, SQL: final input sequence length, LR: learning rate, WRM: linear LR warm-up proportion, EP: number of epochs, STP: number of optimizer steps, GAC: gradient accumulation steps, POSq: whether (y/n) qq part is included in positional embeddings defined in §2.

The hyperparameters for majority tagging are in Table 12, for GenBERT finetuning in Table 13, for MLM trainings in Table 10, for SQuAD finetuning in Table 11 and for HotpotQA finetuning in Table 14.