Big Bird: Transformers for Longer Sequences

Manzil Zaheer, Guru Guruganesh, Avinava Dubey, Joshua Ainslie, Chris Alberti, Santiago Ontanon, Philip Pham, Anirudh Ravula, Qifan Wang, Li Yang, Amr Ahmed

Introduction

Models based on Transformers , such as BERT , are wildly successful for a wide variety of Natural Language Processing (NLP) tasks and consequently are mainstay of modern NLP research. Their versatility and robustness are the primary drivers behind the wide-scale adoption of Transformers. The model is easily adapted for a diverse range of sequence based tasks – as a seq2seq model for translation , summarization , generation , etc. or as a standalone encoders for sentiment analysis , POS tagging , machine reading comprehension , etc. – and it is known to vastly outperform previous sequence models like LSTM . The key innovation in Transformers is the introduction of a self-attention mechanism, which can be evaluated in parallel for each token of the input sequence, eliminating the sequential dependency in recurrent neural networks, like LSTM. This parallelism enables Transformers to leverage the full power of modern SIMD hardware accelerators like GPUs/TPUs, thereby facilitating training of NLP models on datasets of unprecedented size. This ability to train on large scale data has led to surfacing of models like BERT and T5 , which pretrain transformers on large general purpose corpora and transfer the knowledge to down-stream task. The pretraining has led to significant improvement in low data regime downstream tasks as well as tasks with sufficient data and thus have been a major force behind the ubiquity of transformers in contemporary NLP.

The self-attention mechanism overcomes constraints of RNNs (namely the sequential nature of RNN) by allowing each token in the input sequence to attend independently to every other token in the sequence. This design choice has several interesting repercussions. In particular, the full self-attention have computational and memory requirement that is quadratic in the sequence length. We note that while the corpus can be large, the sequence length, which provides the context in many applications is very limited. Using commonly available current hardware and model sizes, this requirement translates to roughly being able to handle input sequences of length 512 tokens. This reduces its direct applicability to tasks that require larger context, like QA , document classification, etc.

However, while we know that self-attention and Transformers are useful, our theoretical understanding is rudimentary. What aspects of the self-attention model are necessary for its performance? What can we say about the expressivity of Transformers and similar models? Apriori, it was not even clear from the design if the proposed self-attention mechanism was as effective as RNNs. For example, the self-attention does not even obey sequence order as it is permutation equivariant. This concern has been partially resolved, as Yun et al. showed that transformers are expressive enough to capture all continuous sequence to sequence functions with a compact domain. Meanwhile, Pérez et al. showed that the full transformer is Turing Complete (i.e. can simulate a full Turing machine). Two natural questions arise: Can we achieve the empirical benefits of a fully quadratic self-attention scheme using fewer inner-products? Do these sparse attention mechanisms preserve the expressivity and flexibility of the original network?

In this paper, we address both the above questions and produce a sparse attention mechanism that improves performance on a multitude of tasks that require long contexts. We systematically develop BigBird, an attention mechanism whose complexity is linear in the number of tokens (Sec. 2). We take inspiration from graph sparsification methods and understand where the proof for expressiveness of Transformers breaks down when full-attention is relaxed to form the proposed attention pattern. This understanding helped us develop BigBird, which is theoretically as expressive and also empirically useful. In particular, our BigBird consists of three main part:

A set of gg global tokens attending on all parts of the sequence.

All tokens attending to a set of ww local neighboring tokens.

All tokens attending to a set of rr random tokens.

This leads to a high performing attention mechanism scaling to much longer sequence lengths (8x).

To summarize, our main contributions are:

BigBird satisfies all the known theoretical properties of full transformer (Sec. 3). In particular, we show that adding extra tokens allows one to express all continuous sequence to sequence functions with only O(n)O(n)-inner products. Furthermore, we show that under standard assumptions regarding precision, BigBird is Turing complete.

Empirically, we show that the extended context modelled by BigBird benefits variety of NLP tasks. We achieve state of the art results for question answering and document summarization on a number of different datasets. Summary of these results are presented in Sec. 4.

Lastly, we introduce a novel application of attention based models where long contexts are beneficial: extracting contextual representations of genomics sequences like DNA. With longer masked LM pretraining, BigBird improves performance on downstream tasks such as promoter-region and chromatin profile prediction (Sec. 5).

There have been a number of interesting attempts, that were aimed at alleviating the quadratic dependency of Transformers, which can broadly categorized into two directions. First line of work embraces the length limitation and develops method around it. Simplest methods in this category just employ sliding window , but in general most work fits in the following general paradigm: using some other mechanism select a smaller subset of relevant contexts to feed in the transformer and optionally iterate, i.e. call transformer block multiple time with different contexts each time. Most prominently, SpanBERT , ORQA , REALM , RAG have achieved strong performance for different tasks. However, it is worth noting that these methods often require significant engineering efforts (like back prop through large scale nearest neighbor search) and are hard to train.

Second line of work questions if full attention is essential and have tried to come up with approaches that do not require full attention, thereby reducing the memory and computation requirements. Prominently, Dai et al. , Sukhbaatar et al. , Rae et al. have proposed auto-regresive models that work well for left-to-right language modeling but suffer in tasks which require bidirectional context. Child et al. proposed a sparse model that reduces the complexity to O(nn)O(n\sqrt{n}), Kitaev et al. further reduced the complexity to O(nlog⁡(n))O(n\log(n)) by using LSH to compute nearest neighbors. Ye et al. proposed binary partitions of the data where as Qiu et al. reduced complexity by using block sparsity. Recently, Longformer [beltagy2020longformer] introduced a localized sliding window based mask with few global mask to reduce computation and extended BERT to longer sequence based tasks. Finally, our work is closely related to and built on the work of Extended Transformers Construction [ainslie2020etc]. This work was designed to encode structure in text for transformers. The idea of global tokens was used extensively by them to achieve their goals. Our theoretical work can be seen as providing a justification for the success of these models as well. It is important to note that most of the aforementioned methods are heuristic based and empirically are not as versatile and robust as the original transformer, i.e. the same architecture do not attain SoTA on multiple standard benchmarks. (There is one exception of Longformer which we include in all our comparisons, see Sec. E.3 for a more detailed comparison). Moreover, these approximations do not come with theoretical guarantees.

BigBird Architecture

The second viewpoint which inspired the creation of BigBird is that most contexts within NLP and computational biology have data which displays a great deal of locality of reference. In this phenomenon, a great deal of information about a token can be derived from its neighboring tokens. Most pertinently, clark2019does investigated self-attention models in NLP tasks and concluded that that neighboring inner-products are extremely important. The concept of locality, proximity of tokens in linguistic structure, also forms the basis of various linguistic theories such as transformational-generative grammar. In the terminology of graph theory, clustering coefficient is a measure of locality of connectivity, and is high when the graph contains many cliques or near-cliques (subgraphs that are almost fully interconnected). Simple Erdős-Rényi random graphs do not have a high clustering coefficient [sussman2017clusteringcoeff], but a class of random graphs, known as small world graphs, exhibit high clustering coefficient [watts1998collective]. A particular model introduced by watts1998collective is of high relevance to us as it achieves a good balance between average shortest path and the notion of locality. The generative process of their model is as follows: Construct a regular ring lattice, a graph with nn nodes each connected to ww neighbors, ww/2 on each side.

In other words we begin with a sliding window on the nodes. Then a random subset (kk%) of all connections is replaced with a random connection. The other (100 - kk)% local connections are retained. However, deleting such random edges might be inefficient on modern hardware, so we retain it, which will not affect its properties. In summary, to capture these local structures in the context, in BigBird, we define a sliding window attention, so that during self attention of width ww, query at location ii attends from i−w/2i-w/2 to i+w/2i+w/2 keys. In our notation, A(i,i−w/2:i+w/2)=1A(i,i-w/2:i+w/2)=1 (see Fig. 1(b)). As an initial sanity check, we performed basic experiments to test whether these intuitions are sufficient in getting performance close to BERT like models, while keeping attention linear in the number of tokens. We found that random blocks and local window were insufficient in capturing all the context necessary to compete with the performance of BERT.

The final piece of BigBird is inspired from our theoretical analysis (Sec. 3), which is critical for empirical performance. More specifically, our theory utilizes the importance of “global tokens” (tokens that attend to all tokens in the sequence and to whom all tokens attend to (see Fig. 1(c)). These global tokens can be defined in two ways:

BigBird-itc: In internal transformer construction (itc), we make some existing tokens “global”, which attend over the entire sequence. Concretely, we choose a subset GG of indices (with g:=∣G∣g:=|G|), such that A(i,:)=1A(i,:)=1 and A(:,i)=1A(:,i)=1 for all i∈Gi\in G.

BigBird-etc: In extended transformer construction (etc), we include additional “global” tokens such as CLS. Concretely, we add gg global tokens that attend to all existing tokens. In our notation, this corresponds to creating a new matrix B∈(N+g)×(N+g)B\in^{(N+g)\times(N+g)} by adding gg rows to matrix AA, such that B(i,:)=1B(i,:)=1, and B(:,i)=1B(:,i)=1 for all i∈{1,2,…g}i\in\{1,2,\ldots g\}, and B(g+i,g+j)=A(i,j)∀ i,j∈{1,…,N}B(g+i,g+j)=A(i,j)\forall\ i,j\in\{1,\ldots,N\}. This adds extra location to store context and as we will see in the experiments improves performance.

The final attention mechanism for BigBird (Fig. 1(d)) has all three of these properties: queries attend to rr random keys, each query attends to w/2w/2 tokens to the left of its location and w/2w/2 to the right of its location and they contain gg global tokens (The global tokens can be from existing tokens or extra added tokens). We provide implementation details in App. D.

Theoretical Results about Sparse Attention Mechanism

In this section, we will show that that sparse attention mechanisms are as powerful and expressive as full-attention mechanisms in two respects. First, we show that when sparse attention mechanisms are used in a standalone encoder (such as BERT), they are Universal Approximators of sequence to sequence functions in the style of Yun19. We note that this property was also explored theoretically in contemporary work yun2020on. Second, unlike [yun2020on], we further show that sparse encoder-decoder transformers are Turing Complete (assuming the same conditions defined in [Perez19]). Complementing the above positive results, we also show that moving to a sparse-attention mechanism incurs a cost, i.e. there is no free lunch. In Sec. 3.4, we show lower bounds by exhibiting a natural task where any sufficiently sparse mechanism will require polynomially more layers.

The complete Transformer encoder stack is nothing but the repeated application of a single-layer encoder (with independent parameters). We denote class of such Transformer encoders stack, defined using generalized encoder (Sec. 2), by TDH,m,q\mathcal{T}_{D}^{H,m,q} which consists of HH-heads with head size mm and qq is the hidden layer size of the output network, and the attention layer is defined by the directed graph DD.

2 Universal Approximators

The star-graph SS centered at is the graph defined on {0,…,n}\{0,\dots,n\}. The neighborhood of all vertices ii is N(i)={0,i}N(i)=\{0,i\} for i∈{1…n}i\in\{1\dots n\} and N(0)={1,…n}N(0)=\{1,\dots n\}.

Our main theorem is that the sparse attention mechanism defined by any graph containing SS is a universal approximator:

Given 1<p<∞1<p<\infty and ϵ>0\epsilon>0, for any f∈FCDf\in\mathcal{F}_{CD}, there exists a transformer with sparse-attention, g∈TDH,m,qg\in\mathcal{T}_{D}^{H,m,q} such that dp(f,g)≤ϵd_{p}(f,g)\leq\epsilon where DD is any graph containing star graph SS.

To prove the theorem, we will follow the standard proof structure outlined in [Yun19].

Step 2: Approximate piece-wise constant functions by modified transformers. This is the key step of the proof where the self-attention mechanism is used to generate a contextual-mapping of the input. Informally, a contextual mapping is a unique code for the pair consisting of a matrix (X,xi)({\bm{X}},{\bm{x}}_{i}) and a column. Its uniqueness allows the Feed forward layers to use each code to map it to a unique output column.

The main technical challenge is computing the contextual mapping using only sparse attention mechanism. This was done in [Yun19] using a “selective” shift operator which shift up entries that are in a specific interval. Key to their proof was the fact that the shift, was exactly the range of the largest entry to the smallest entry.

Creating a contextual mapping with a sparse attention mechanism is quite a challenge. In particular, because each query only attends to a few keys, it is not at all clear that sufficient information can be corralled to make a contextual embedding of the entire matrix. To get around this, we develop a sparse shift operator which shifts the entries of the matrices if they lie in a certain range. The exact amount of the shift is controlled by the directed sparse attention graphg DD. The second key ingredient is the use of additional global token. By carefully applying the operator to a set of chosen ranges, we will show that each column will contain a unique mapping of the full mapping. Therefore, we can augment the loss of inner-products in the self attention mechanism by using multiple layers and an auxiliary global token.

Step 3: Approximate modified transformers by original Transformers: The final step is to approximate the modified transformers by the original transformer which uses ReLU and softmax.

3 Turing Completeness

Transformers are a very general class. In the original paper of vaswani2017attention, they were used in both an encoder and a decoder. While the previous section outlined how powerful just the encoders were, another natural question is to ask what the additional power of both a decoder along with an encoder is? Perez19 showed that the full transformer based on a quadratic attention mechanism is Turing Complete. This result makes one unrealistic assumption, which is that the model works on arbitrary precision model. Of course, this is necessary as otherwise, Transformers are bounded finite state machines and cannot be Turing Complete.

It is natural to ask if the full attention mechanism is necessary. Or can a sparse attention mechanism also be used to simulate any Turing Machine? We show that this is indeed the case: we can use a sparse encoder and sparse decoder to simulate any Turing Machine.

To use the sparse attention mechanism in the transformer architecture, we need to define a suitable modification where each token only reacts to previous tokens. Unlike the case for BERT, where the entire attention mechanism is applied once, in full transformers, the sparse attention mechanism at decoder side is used token by token. Secondly the work of Perez19, uses each token as a representation of the tape history and uses the full attention to move and retrieve the correct tape symbol. Most of the construction of Perez19 goes through for sparse attentions, except for their addressing scheme to point back in history (Lemma B.4 in [Perez19]). We show how to simulate this using a sparse attention mechanism and defer the details to App. B.

4 Limitations

Task 1. Given nn unit vectors {u1,…,un}\{u_{1},\dots,u_{n}\}, find f(u1,…,un)→(u1∗,…,un∗)f(u_{1},\dots,u_{n})\to(u_{1^{*}},\dots,u_{n^{*}}) where for a fixed j∈[n]j\in[n], we define j∗=arg max⁡k∥uk−uj∥22j^{*}=\operatorname*{arg\,max}_{k}\|u_{k}-u_{j}\|_{2}^{2}.

Finding vectors that are furthest apart boils down to minimize inner product search in case of unit vectors. For a full-attention mechanism with appropriate query and keys, this task is very easy as we can evaluate all pair-wise inner products.

The impossibility for sparse-attention follows from hardness results stemming from Orthogonal Vector Conjecture(OVC) [abboud2014consequences, abboud2015tight, backurs2015edit, williams2005new]. The OVC is a widely used assumption in fine-grained complexity. Informally, it states that one cannot determine if the minimum inner product among nn boolean vectors is in subquadratic time. In App. C, we show a reduction using OVC to show that if a transformer g∈TDH=1,m=2d,q=0g\in\mathcal{T}_{D}^{H=1,m=2d,q=0} for any sparse directed graph DD can evaluate the Task 11, it can solve the orthogonal vector problem.

We give a formal proof of this fact in App. C.

Experiments: Natural Language Processing

In this section our goal is to showcase benefits of modeling longer input sequence for NLP tasks, for which we select three representative tasks. We begin with basic masked language modeling (MLM; devlin2018bert) to check if better contextual representations can be learnt by utilizing longer contiguous sequences. Next, we consider QA with supporting evidence, for which capability to handle longer sequence would allow us to retrieve more evidence using crude systems like TF-IDF/BM25. Finally, we tackle long document classification where discriminating information may not be located in first 512 tokens. Below we summarize the results for BigBird using sequence length 4096code available at http://goo.gle/bigbird-transformer, while we defer all other setup details including computational resources, batch size, step size, to App. E.

We follow [devlin2018bert, liu2019roberta] to create base and large versions of BigBird and pretrain it using MLM objective. This task involves predicting a random subset of tokens which have been masked out. We use four standard data-sets for pretraining (listed in Sec. E.1, Tab. 10), warm-starting from the public RoBERTa checkpointhttps://github.com/pytorch/fairseq/tree/master/examples/roberta. We compare performance in predicting the masked out tokens in terms of bits per character, following [beltagy2020longformer]. As seen in Sec. E.1, Tab. 10, both BigBird and Longformer perform better than limited length RoBERTa, with BigBird-etc performing the best. We note that we trained our models on a reasonable 16GB16GB memory/chip with batch size of 32-64. Our memory efficiency is due to efficient blocking and sparsity structure of the sparse attention mechanism described in Sec. 2.

We considered following four challenging datasets:

Natural Questions [kwiatkowski2019natural]: For the given question, find a short span of answer (SA) from the given evidences as well highlight the paragraph from the given evidences containing information about the correct answer (LA).

HotpotQA-distractor [yang2018hotpotqa]: Similar to natural questions, it requires finding the answer (Ans) as well as the supporting facts (Sup) over different documents needed for multi-hop reasoning from the given evidences.

TriviaQA-wiki [JoshiTriviaQA2017]: We need to provide an answer for the given question using provided Wikipedia evidence, however, the answer might not be present in the given evidence. On a smaller verified subset of question, the given evidence is guaranteed to contain the answer. Nevertheless, we model the answer as span selection problem in this case as well.

WikiHop [welbl2018constructing]: Chose correct option from multiple-choice questions (MCQ), by aggregating information spread across multiple documents given in the evidences.

As these tasks are very competitive, multiple highly engineered systems have been designed specific each dataset confirming to respective output formats. For a fair comparison, we had to use some additional regularization for training BigBird, details of which are provided in Sec. E.2 along with exact architecture description. We experiment using the base sized model and select the best configuration on the development set for each dataset (as reported in Tab. 2). We can see that BigBird-etc, with expanded global tokens consistently outperforms all other models. Thus, we chose this configuration to train a large sized model to be used for evaluation on the hidden test set.

In Tab. 3, we compare BigBird-etc model to top-3 entries from the leaderboard excluding BigBird. One can clearly see the importance of using longer context as both Longformer and BigBird outperform models with smaller contexts. Also, it is worth noting that BigBird submission is a single model, whereas the other top-3 entries for Natural Questions are ensembles, which might explain the slightly lower accuracy in exact answer phrase selection.

We experiment on datasets of different lengths and contents, specifically various document classification and GLUE tasks. Following BERT, we used one layer with cross entropy loss on top of the first [CLS] token. We see that gains of using BigBird are more significant when we have longer documents and fewer training examples. For instance, using base sized model, BigBird improves state-of-the-art for Arxiv dataset by about 5%\bm{5\%} points. On Patents dataset, there is improvement over using simple BERT/RoBERTa, but given the large size of training data the improvement over SoTA (which is not BERT based) is not significant. Note that this performance gain is not seen for much smaller IMDb dataset. Along with experimental setup detail, we present detailed results in Sec. E.4 which show competitive performance.

1 Encoder-Decoder Tasks

For an encoder-decoder setup, one can easily see that both suffer from quadratic complexity due to the full self attention. We focus on introducing the sparse attention mechanism of BigBird only at the encoder side. This is because, in practical generative applications, the length of output sequence is typically small as compared to the input. For example for text summarization, we see in realistic scenarios (c.f. Sec. E.5 Tab. 18) that the median output sequence length is ∼200\sim 200 where as the input sequence’s median length is >3000>3000. For such applications, it is more efficient to use sparse attention mechanism for the encoder and full self-attention for the decoder.

Document summarization is a task of creating a short and accurate summary of a text document. We used three long document datasets for testing our model details of which are mention in Tab. 18. In this paper we focus on abstractive summarization of long documents where using a longer contextual encoder should improve performance. The reasons are two fold: First, the salient content can be evenly distributed in the long document, not just in first 512 tokens, and this is by design in the BigPatents dataset [sharma2019bigpatent]. Second, longer documents exhibit a richer discourse structure and summaries are considerably more abstractive, thereby observing more context helps. As has been pointed out recently [rothe2019leveraging, zhang2019pegasus], pretraining helps in generative tasks, we warm start from our general purpose MLM pretraining on base-sized models as well as utilizing state-of-the-art summarization specific pretraining from Pegasus [zhang2019pegasus] on large-sized models. The results of training BigBird sparse encoder along with full decoder on these long document datasets are presented in Tab. 4. We can clearly see modeling longer context brings significant improvement. Along with hyperparameters, we also present results on shorter but more widespread datasets in Sec. E.5, which show that using sparse attention does not hamper performance either.

Experiments: Genomics

There has been a recent upsurge in using deep learning for genomics data [tampuu2019viraminer, zhang2019ncnet, busia2019deep], which has resulted in improved performance on several biologically-significant tasks such as promoter site prediction [oubounyt2019deepromoter], methylation analysis [levy2020methylnet], predicting functional effects of non-coding variant [zhou2015predicting], etc. These approaches consume DNA sequence fragments as inputs, and therefore we believe longer input sequence handling capability of BigBird would be beneficial as many functional effects in DNA are highly non-local [buldyrev1995long]. Furthermore, taking inspiration from NLP, we learn powerful contextual representations for DNA fragments utilizing abundant unlabeled data (e.g. human reference genome, Saccharomyces Genome Database) via MLM pretraining. Next, we showcase that our long input BigBird along with the proposed pretraining significantly improves performances in two downstream tasks. Detailed experimental setup for the two tasks are provided in App. F.

As explored in liang2012segmenting, instead of operating on base pairs, we propose to first segment DNA into tokens so as to further increase the context length (App. F, Fig. 7). In particular, we build a byte-pair encoding [kudo2018sentencepiece] table for the DNA sequence of size 32K, with each token representing 8.78 base pairs on average. We learn contextual representation of these token on the human reference genome (GRCh37)https://www.ncbi.nlm.nih.gov/assembly/GCF_000001405.13/ using MLM objective. We then report the bits per character (BPC) on a held-out set in Tab. 5. We find that attention based contextual representation of DNA does improve BPC, which is further improved by using longer context.

Promoter is a DNA region typically located upstream of the gene, which is the site of transcription initiation. Multiple methods have been proposed to identify the promoter regions in a given DNA sequence [yang2017exploiting, lin2017identifying, bharanikumar2018promoterpredict, xiao2019ipsw, oubounyt2019deepromoter], as it is an important first step in understanding gene regulation. The corresponding machine learning task is to classify a given DNA fragment as promoter or non-promoter sequence. We use the dataset compiled by oubounyt2019deepromoter which was built from Eukaryotic Promoter Database (EPDnew) [dreos2013epd] https://epd.epfl.ch/human/human_database.php?db=human. We finetuned the pretrained BigBird model from above, using the training data and report F1 on test dataset. We compare our results to the previously reported best method in Tab. 6. We see that BigBird achieve nearly perfect accuracy with a 5%5\% jump from the previous best reported accuracy.

Non-coding regions of DNA do not code for proteins. Majority of diseases and other trait associated single-nucleotide polymorphism are correlated to non-coding genomic variations [zhou2015predicting, khurana2016role]. Thus, understanding the functional effects of non-coding regions of DNA is a very important task. An important step in this process, as defined by zhou2015predicting, is to predict large-scale chromatin-profiling from non-coding genomic sequence. To this effect, DeepSea [zhou2015predicting], compiled 919 chromatin-profile of 2.4M non-coding variants from Encyclopedia of DNA Elements (ENCODE)https://www.encodeproject.org/ and Roadmap Epigenomics projectshttp://www.roadmapepigenomics.org/. The corresponding ML task is to predict, for a given non-coding region of DNA, these 919 chromatin-profile including 690690 transcription factors (TF) binding profiles for 160160 different TFs, 125125 DNase I sensitivity (DHS) profiles and 104104 histone-mark (HM) profiles. We jointly learn 919 binary classifiers to predict these functional effects from sequence of DNA fragments. On held-out chromosomes, we compare AUC with the baselines in Tab. 7 and see that we significantly improve on performance on the harder task HM, which is known to have longer-range correlations [gates2017histone] than others.

Conclusion

We propose BigBird: a sparse attention mechanism that is linear in the number of tokens. BigBird satisfies a number of theoretical results: it is a universal approximator of sequence to sequence functions and is also Turing complete. Theoretically, we use the power of extra global tokens preserve the expressive powers of the model. We complement these results by showing that moving to sparse attention mechanism do incur a cost. Empirically, BigBird gives state-of-the-art performance on a number of NLP tasks such as question answering and long document classification. We further introduce attention based contextual language model for DNA and fine-tune it for down stream tasks such as promoter region prediction and predicting effects of non-coding variants.

References

Appendix A Universal Approximators

An attention mechanism Attn that takes in the sequence X{\bm{X}} and returns sequence (a1,...,an)({\bm{a}}_{1},...,{\bm{a}}_{n}) of the same length and dimensionality; and

Then ii-th output vector of Enc⁡(X)\operatorname{Enc}({\bm{X}}) is computed as follows:

Now it remains to define Attn and OO which we do next.

where N(i)N(i) denote the out-neighbors set of node ii in DD. In other words, the set of arcs (directed edges) in DD represents the set of inner products that our attention mechanism will consider. Also recall that σ\sigma is a scoring function such as softmax or hardmax.

Lastly, we define the output fully connected network as follows:

Additional Notation We introduce a few pieces of additional notation that will be useful. Let [a,b)δ={a,a+δ,…,a+⌊b−aδ⌋⋅δ}[a,b)_{\delta}=\{a,a+\delta,\dots,a+\lfloor\frac{b-a}{\delta}\rfloor\cdot\delta\}. Therefore, [0,1)δ={0,δ,2δ,…,(1−δ)}[0,1)_{\delta}=\{0,\delta,2\delta,\dots,(1-\delta)\}. We use 1[E]\mathbf{1}[\mathcal{E}] to denote the indicator variable; it is 11 if the event E\mathcal{E} occurs and otherwise.

A.2 Proof

In this section, we will present the full proof of theorem 1. The proof will contain three parts. The first and the third part will largely follow standard techniques. The main innovation lies is in the second part.

First, we consider a suitable partition of the region (0,1)(0,1) into a grid of granularity δ\delta, which we denote by GδG_{\delta}. We do this using Lemma 8 from Yun19, which we restate for completeness:

For any given f∈FCDf\in\mathcal{F}_{CD} and 1≤p≤∞1\leq p\leq\infty, there exists a δ>0\delta>0 such that there exists a piece-wise constant function fˉ\bar{f} with dp(f,fˉ)≤ϵ3d_{p}(f,\bar{f})\leq\frac{\epsilon}{3}. Concretely, fˉ\bar{f} is defined as

Since transformers can learn a positional embedding EE, without any loss of generality, we can consider the translated function. In particular, define

We will try to approximate g(X)=f(X−E)g(X)=f(X-E) where gg is defined on the domain d×[δ−d,δ−d+1]d×⋯×[δ−(n−1)d,δ−(n−1)d+1]d^{d}\times[\delta^{-d},\delta^{-d}+1]^{d}\times\dots\times[\delta^{-(n-1)d},\delta^{-(n-1)d}+1]^{d}. To do so, we will apply a suitable modification of Lemma 1, which will consider the discretized grid

A.2.2 Contextual Mappings and Sparse Attention Mechanisms

The main idea in this section is the use of contextual mapping to enable Transformers to compute any discretized function. A contextual mapping is an unique encoding of each tuple (X,xi)(X,x_{i}) where X∈GδEX\in\mathbf{G}^{E}_{\delta}, and each column xi∈[δ−(i−1)d,δ−(i−1)d+1)δdx_{i}\in[\delta^{-(i-1)d},\delta^{-(i-1)d}+1)^{d}_{\delta} for all i∈[n]i\in[n]. We restate the definition adapted to our setting below

For any P∈GδEP\in\mathbf{G}^{E}_{\delta}, q(P)q(P) contains distinct entries.

For any two P,P′∈GδEP,P^{\prime}\in\mathbf{G}^{E}_{\delta} with P≠P′P\neq P^{\prime}, all entries of q(P)q(P) and q(P′)q(P^{\prime}) are distinct.

The key technical novelty of the proof is computing a contextual mapping using only the sparse attention mechanism. We create a “selective shift” operator which only shifts entries of a vector that lie in a certain range. We will use this shift operator strategically to ensure that we attain a contextual mapping at the end of the process. The lemma below, which is based on parts of the proof of Lemma 6 of [Yun19], states that we can implement a suitable “selective” shift operator using a sparse attention mechanism.

Note that e1∈Rd+1e_{1}\in R^{d+1} denotes (1,0,…,0)(1,0,\dots,0).

Consider the function , which can be implemented by a sparse attention mechanism :

This is because the Key, Query and Value functions are simply affine transformations of XX.

Given any graph DD, the above function will evaluate to the following:

The following lemma, which is the heart of the proof, uses the above selective shift operators to construct contextual mappings.

To successfully encode the entire context in each token, we will interleave the shift operator to target the original columns 1,…,n1,\dots,n and to target the global column . After a column ii is targeted, its inner product with uu will encode the entire context of the first ii columns. Next, we will shift the global token to take this context into account. This can be subsequently used by the remaining columns.

For i∈{0,1,…,n}i\in\{0,1,\dots,n\}, we will use lil_{i} to denote the innerproducts ⟨u,xi⟩\left\langle u,x_{i}\right\rangle at the beginning. For fi=⟨u,xi⟩f_{i}=\left\langle u,x_{i}\right\rangle after the ithi^{th} column has changed for i∈{1,…,n}i\in\{1,\dots,n\} and we will use f0kf_{0}^{k} to denote ⟨u,x0⟩\left\langle u,x_{0}\right\rangle after the kthk^{th} phase. We need to distinguish the global token further as it’s inner product will change in each phase. Initially, given X∈GδEX\in\mathbf{G}^{E}_{\delta}, the following are true:

Note that all lil_{i} ordered in distinct buckets l1<l2<⋯<ln<l0l_{1}<l_{2}<\dots<l_{n}<l_{0}.

We do this in phases indexed from i∈{1,…,n}i\in\{1,\dots,n\}. Each phase consists of two distinct parts: The low shift operation: These operation will be of the form

for values v∈[δ−id),δ−(i+1)d)δv\in[\delta^{-id}),\delta^{-(i+1)d})_{\delta}. The range is chosen so that only lil_{i} will be in the range and no other ljl_{j} j≠ij\neq i is in the range. This will shift exactly the ithi^{th} column xix_{i} so that the new inner product fi=⟨u,xi⟩f_{i}=\left\langle u,x_{i}\right\rangle is substantially larger than lil_{i}. Furthermore, no other column of XX will be affected. The high shift operation: These operation will be of the form

Finally, we define the following constants for all k∈{0,1,…,n}k\in\{0,1,\dots,n\}.

After each kk phases, we will maintain the following invariants:

The order of the inner products after kthk^{th} phase is

The case k=0k=0, is trivial as we simply set S0=δ−(n+1)dS_{0}=\delta^{-(n+1)d}, T0=δ−(n+1)⋅d+δT_{0}=\delta^{-(n+1)\cdot d}+\delta.

The previous lemma shows that we can compute a contextual mapping using only sparse transforms. We now use the following lemma to show that we can use a contextual mapping and feed-forward layers to accurately map to the desired output of the function fˉ\bar{f}.

A.2.3 Approximating modified Transformers by Transformers

The previous section assumed we used Transformers that used hardmax operator σH\sigma_{H} and activations functions belonging to the set Φ\Phi. This is without loss of generality as following lemma shows.

For each g∈Tˉ2,1,1g\in\bar{\mathcal{T}}^{2,1,1} and 1≤p≤∞1\leq p\leq\infty, ∃g∈T2,1,4\exists g\in\mathcal{T}^{2,1,4} such that dp(g,gˉ)≤ϵ/3d_{p}(g,\bar{g})\leq\epsilon/3

Combining the above lemma with the Lemma 3, we get our main result:

Let 1≤p≤∞1\leq p\leq\infty and ϵ>0\epsilon>0, there exists a transformer network g∈TD2,1,4g\in\mathcal{T}_{D}^{2,1,4} which achieves a ratio of dp(f,g)≤ϵd_{p}(f,g)\leq\epsilon where DD is the sparse graph.

Since the sparsity graph associated with BigBird contains a star network, we know that it can express any continuous function from a compact domain.

We would like to note that, contemporary work done by yun2020on, also parallelly explored the ability of sparse transformers with linear connections to capture sequence-to-sequence functions on the compact domain.

Appendix B Turing Completeness

In this section, we will extend our results to the setting of Perez19. Our exposition will largely use their proof structure but we will make a few changes. We repeat some of the lemmas with the amendments to make the exposition self-contained.

An attention mechanism Attn that takes in the sequence Yj{\bm{Y}}_{j} and returns sequence (p1,...,pj)({\bm{p}}_{1},...,{\bm{p}}_{j}) of the same length and dimensionality;

Then ii-th output vector of Dec⁡(Yj;Ke,Ve)\operatorname{Dec}({\bm{Y}}_{j};{\bm{K}}^{\textbf{e}},{\bm{V}}^{\textbf{e}}) is computed as follows:

\textscAttnD\textsc{Attn}_{D} and OO are as defined in Sec. A.1 and it remains to define CrossAttn. The ithi^{\textrm{th}} output vector of multi-head cross-attention attention is given by

We will use the same setup of Turning Machine that was used by Perez19 (see section B.4). Given a Turing Machine M=(Q,Σ,δ,qinit,F)M=(Q,\Sigma,\delta,q_{init},F), we use the following notation

B.2 Details of the Simulation

In this section, we give more details on the architecture of the encoder and decoder needed to implement our simulation strategy.

Given the Turing machine MM, we will show that a transformer with an appropriate encoder and decoder TD\mathcal{T}_{D} can simulate each step of MM’s execution. Our simulation strategy will mostly follow Perez19, except we will use a sparse attention mechanism. The main idea is to maintain the current Turing machine state q(j)q^{(j)} and symbol under the head s(j)s^{(j)} as part of the decoder sequence Y{\bm{Y}} for all time step jj so that we can always simulate the corresponding Turing machine transition δ(q(j),s(j))=(q(j),v(j),m(j))\delta(q^{(j)},s^{(j)})=(q^{(j)},v^{(j)},m^{(j)}). The key difference will rise in Lemma B.4 of Perez19, where full attention is used to select the appropriate symbol from tape history in one step. To accomplish the same task with sparse attention, we will exploit the associative property of max and break down the symbol selection over multiple steps. Thus, unlike Perez19 one decoding step of our sparse transformer TD\mathcal{T}_{D} does not correspond to one step of the Turing machine MM. In particular, we will have two type of steps: compute step corresponding to update of MM’s state and intermediate steps corresponding to aggregating the max (which in turn is used for symbol selection). Let ii denote the step of TD\mathcal{T}_{D} and g(i)g(i) denote the step of MM being simulated at step ii of the decoder. At each decoding step we want to maintain the current Turing machine state qg(i)q^{g(i)} and symbol under the sg(i)s^{g(i)} in yi{\bm{y}}_{i}. For roughly O(i)O(\sqrt{i}) intermediate steps the state will remain the same, while we aggregate information about relevant past output symbols through sparse attention. To maintain the same state for intermediate steps, we introduce an extra switching layer (Sec. B.2.3). Finally, at the next compute step we will make the transition to new state qg(i)+1q^{g(i)+1}, new head movement mg(i)m^{g(i)}, and new output symbol vg(i)v^{g(i)} to be written. Thereby we are able to completely simulate the given Turing machine MM. As a result, we can prove the following main theorem:

There exists a sparse attention mechanism using O(n)O(n) inner products such that the resulting class of Transformer Networks using this sparse attention mechanism is Turing Complete.

Encoder

As [Perez19], we use the same trivial single layer encoder where resulting K(e){\bm{K}}^{(e)} contains position embedding and V(e){\bm{V}}^{(e)} contains one-hot symbol representation.

Decoder

This graph can be seen as a special case of BigBird where first type of edges are realizations of random and second type of edges correspond to locality. Also note that this graph satisfies the left-to-right constraint of decoder, i.e. no node attends to a node in the future.

where g(i)=⌊−1+1+8i2⌋g(i)=\left\lfloor\frac{-1+\sqrt{1+8i}}{2}\right\rfloor and h(i)=g(i+1)−g(i)h(i)=g(i+1)-g(i). Note that h(i)h(i) reduces to a binary indicator variable 1{−1+1+8i2=⌊−1+1+8i2⌋}\mathbf{1}\left\{\frac{-1+\sqrt{1+8i}}{2}=\left\lfloor\frac{-1+\sqrt{1+8i}}{2}\right\rfloor\right\}.

Induction Setup

We next show how to construct the decoder layers to produce the sequence of outputs y1,y2,…{\bm{y}}_{1},{\bm{y}}_{2},\ldots, where yi{\bm{y}}_{i} is given by:

That is, at step ii of our sparse decoder yi{\bm{y}}_{i}, it will contain the information about the state of the turing machine MM at time g(i)g(i), the symbol under the head of MM at time g(i)g(i), and the current location of head of MM at time g(i)g(i). We also have a placeholder symbol ww and placeholder scalars u1,u2,u3u_{1},u_{2},u_{3}, whose role will be clear from our construction.

We consider as the starting vector for the decoder the vector

We assume that the start head is at c(0)=0c^{(0)}=0, the initial state is q(0)=qinitq^{(0)}=q_{\text{init}}, and s(0)=#s^{(0)}=\# as we initialize from clean tape. We show the correctness of our construction by an inductive argument: we describe the architecture piece by piece and at the same time will show for every r≥0r\geq 0 , our architecture constructs yr+1{\bm{y}}_{r+1} from the previous vectors (y0,…,yr)({\bm{y}}_{0},\ldots,{\bm{y}}_{r}).

Thus, assume that y1,…,yr{\bm{y}}_{1},\ldots,{\bm{y}}_{r} satisfy the properties stated above. Since we are using positional encodings, the actual input for the first layer of the decoder is the sequence

We denote by y‾i\overline{{\bm{y}}}_{i} the vector yi{\bm{y}}_{i} plus its positional encoding. Thus we have ∀ 1≤i≤r\forall\ 1\leq i\leq r that

B.2.1 Layer 1: Simulate Transition Function

In this layer, we use the cross-attention between encoder and decoder to access the input string and a feed-forward network to simulate the transition function of MM. The first self attention in Eq. 5 is not used in this layer and we just produce the identity. This identity function is achieved by setting all queries, keys, values to be 0 everywhere plus the residual connection. Thus, we have pi1=y‾i{{\bm{p}}}^{1}_{i}=\overline{{\bm{y}}}_{i}.

Since pi1{\bm{p}}^{1}_{i} is of the form [A‾,…,A‾,1,g(i)+1,A‾,…,A‾][\underline{\phantom{A}},\ldots,\underline{\phantom{A}},1,g(i)+1,\underline{\phantom{A}},\ldots,\underline{\phantom{A}}], we know by Lemma B.1 of Perez19 that if we use pi1{\bm{p}}^{1}_{i} to attend over the encoder we obtain

where α\alpha and β\beta are as defined in Eq. (21) of [Perez19]. Thus in Eq. 4 we finally produce the vector ai1{\bm{a}}^{1}_{i} given by

As the final piece of the first decoder layer we use a function O1(⋅)O_{1}(\cdot) (Eq. 3) that satisfies the following lemma.

That is, function O1(⋅)O_{1}(\cdot) simulates transition δ(qg(i),sg(i))\delta(q^{g(i)},s^{g(i)}) to construct ⟦ qg(i)+1 ⟧\llbracket\ q^{g(i)+1}\ \rrbracket, ⟦ vg(i) ⟧\llbracket\ v^{g(i)}\ \rrbracket, and mg(i)m^{g(i)} besides some other linear transformations.

Thus, finally the output of the first decoder layer is

B.2.2 Layer 2: Finding Head Node

In this layer, we only use the feed-forward network to evaluate the next location of the head. The self-attention and cross-attention are set to be the identity function, so ai2=pi2=zi1{\bm{a}}_{i}^{2}={\bm{p}}_{i}^{2}={\bm{z}}_{i}^{1}. Recall that cg(i)c^{g(i)} is the cell to which MM is pointing to at time g(i)g(i), and that it satisfies the following recursion cg(i)+1=cg(i)+mg(i)c^{g(i)+1}=c^{g(i)}+m^{g(i)}, which can be expanded to see that that cg(i)+1=m(0)+m(1)+⋯+mg(i)c^{g(i)+1}=m^{(0)}+m^{(1)}+\cdots+m^{g(i)}. Its not difficult to see that a two layer network with non-linearity can compute cg(i)+1/(g(i)+1)c^{g(i)+1}/(g(i)+1) and cg(i)/(g(i)+1)c^{g(i)}/(g(i)+1) from cg(i)c^{g(i)}, mg(i)m^{g(i)}, and 1/(g(i)+1)1/(g(i)+1) using the relation cg(i)+1=cg(i)+mg(i)c^{g(i)+1}=c^{g(i)}+m^{g(i)}. At the end of layer 2, we obtain

B.2.3 Layer 3: Distinguishing Node Type

This is an additional layer (not present in the work of [Perez19]), where we propagate computations in our sparse graph. In particular, we will use this layer to “compute” or accumulate state in intermediate nodes. We make this clear below. The self-attention and cross-attention are all set to be the identity function, so ai3=pi3=zi2{\bm{a}}_{i}^{3}={\bm{p}}_{i}^{3}={\bm{z}}_{i}^{2}. In this layer, we only use the dense attention layers to select the newly computed states or to continue with previous states. Using idea similar to Lemma B.6 of [Perez19], we can construct a dense network such that

The negatives are generated to offset results from skip connection. We utilize such network to switch Turing machine state and position embedding for intermediate steps to the values received from previous time step and do nothing for compute nodes. We use h(i)h(i) as the flipping bit bb. Thus, at end of layer 3, we obtain

where we used h(i)h(i) for selecting old states. In particular,

We copy the input state and head position as is for intermediate nodes. We do not need to transition to next Turing machine states in these nodes.

To preserve the symbol under the head for intermediate nodes, we copy the previous symbol to α\alpha location and set β=g(i)+1\beta=g(i)+1, as the symbol at α\alpha location will be copied as the symbol under head for next transformer step by the final transformation layer if β=g(i)+1\beta=g(i)+1. Thus, we correctly preserve the previous symbol under head as Turing machine does not transition these nodes. For compute nodes, things happen as usual.

Finally for the intermediate nodes, we copy the position embedding corresponding to current best symbol ww, which is stored in u1,u2,u3u_{1},u_{2},u_{3}. For compute node, we let the position embedding correspond to current Turing machine step.

For further simplification note that g(i+1)=g(i)g(i+1)=g(i) if h(i)=0h(i)=0 else g(i)+1g(i)+1 when h(i)=1h(i)=1. With this fact, we can conclude that q^(i)=qg(i+1)\hat{q}^{(i)}=q^{g(i+1)} and c^(i)=cg(i+1)\hat{c}^{(i)}=c^{g(i+1)}. Thus, we can write,

B.2.4 Layer 4: Finding next symbol on tape

We use similar query, key, value functions as used for full attention by [Perez19] ∀i\forall i:

It is clear that the three functions are linear transformations and thus they can be defined by feed-forward networks. Notice that the query vector is always formed using current time step position embedding, whereas key and value vectors are formed using copied over entries for intermediate nodes and using current entries only for compute node.

Perez19 find the desired vl(j+1)v^{l(j+1)} as vm(j)v^{m(j)} using full attention, where

Note the minimization is only over Turing machine steps, i.e. over compute nodes in our case. We show below that we can estimates m(j)m(j) by parts using sparse attention mechanism. The main idea is just to notice that minimization problem min⁡m∈{0,...,t}χtj\min_{m\in\{0,...,t\}}\chi_{t}^{j} can be expressed as min⁡{⋯min⁡{min⁡{χ0j,χ1j},χ2j},...,χtj}\min\{\cdots\min\{\min\{\chi^{j}_{0},\chi^{j}_{1}\},\chi^{j}_{2}\},...,\chi^{j}_{t}\} by the associativity property.

By definition of our graph DD, at every intermediate node ii of the form j(j+1)/2+kj(j+1)/2+k, i.e. where k>0k>0, g(i)=jg(i)=j and h(i)=0h(i)=0, we will attend over node k(k+1)/2k(k+1)/2 and best till now copied from i−1i-1. The node k(k+1)/2k(k+1)/2 is never an intermediate node as h(k(k+1)/2)=1h(k(k+1)/2)=1 for all kk and in fact corresponds to Turing machine’s step kk. This will help us select the key and value corresponding to min between node k(k+1)/2k(k+1)/2 and i−1i-1. In other words, at node ii of the form j(j+1)/2+kj(j+1)/2+k we would have evaluated m(k)m(k) and corresponding value selected:

The cross-attention and feed-forward network are set to be identity, so zi4=ai4=pi4{\bm{z}}_{i}^{4}={\bm{a}}_{i}^{4}={\bm{p}}_{i}^{4}.

B.2.5 Final transformation

We finish our construction by using the final transformation function F(⋅)F(\cdot) from the corresponding lemma from Perez19, with a slight modification.

The modification is to let w,u1,u2,u3w,u_{1},u_{2},u_{3} to pass through. This yields the desired input to transformer at next time step for both intermediate and compute node, thereby concluding our induction.

Appendix C Limitations

We consider the simple problem of finding the furthest vector for each vector in the given sequence of length nn and dimension d∈Ω(log⁡2n)d\in\Omega(\log^{2}n). The assumption on the dimension is mild , as in many situations the dimension d=768d=768 is actually comparable to the number of nn.

Finding vectors that are furthest apart boils down to minimizing inner product search in case of unit vectors. For a full-attention mechanism with appropriate query and keys, this task is very easy as we can evaluate all pair-wise inner products.

The impossibility for sparse-attention follows from hardness results stemming from Orthogonal Vector Conjecture (OVC) [abboud2015tight, abboud2014consequences, williams2005new, backurs2015edit], which is a widely used assumption in fine-grained complexity. Informally, it states that one cannot determine if the minimum inner product among nn Boolean vectors is in subquadratic time.

For every ϵ>0\epsilon>0, there is a c≥1c\geq 1 such that given nn Boolean vectors in dd dimension, cannot determine if there is a pair of orthogonal vectors in O(n2−ϵ)O(n^{2-\epsilon}) time on instances with d≥clog⁡nd\geq c\log n.

Using 1, we show a reduction to show that a transformer g∈TDH=O(d),m=O(d),q=O(d)g\in\mathcal{T}_{D}^{H=O(d),m=O(d),q=O(d)} for any sparse directed graph DD which completes Task 11 must require a superlinear number of layers.

We begin by providing an explicit construction of a single layer full self-attention that can evaluate Task 1.

Step 2 Construct query, key, value functions as follows:

Step 3 Let O(ai)=0O(a_{i})=0, then the output zi=[ui;ui∗]z_{i}=[u_{i};u_{i^{*}}] as desired.

To complete the argument, observe that it now only takes O(n)O(n) inner products to check if there is a pair of orthogonal vectors as we need only compare ⟨ui,ui∗⟩\left\langle u_{i},u_{i^{*}}\right\rangle.

Suppose we can solve Task 1 using a network g∈TDH=O(d),m=O(d),q=O(d)g\in\mathcal{T}_{D}^{H=O(d),m=O(d),q=O(d)} that has ll layers. Recall that all the computation we do in one layer is:

Appendix D Implementation details

We optimize the code for modern hardware. Hardware accelerators like GPUs and TPUs truly shine on coalesced memory operations which load blocks of contiguous bytes at once. Thus, its not very efficient to have small sporadic look-ups caused by a sliding window or random element queries. We alleviate this by “blockifying” the lookups.

Ideally, if the adjacency matrix AA described in Sec. 2 is sparse, one would hope this would be sufficient to speed up the implementation. Unfortunately, it is well known [gray2017gpu, yao2019balanced], that such sparse multiplications cannot be efficiently implemented in GPUs. GPUs have thousands of cores performing operations in parallel. Thus, we cannot efficiently perform the sparse matrix multiplication mentioned in section Sec. 2.

As a result we propose to first blockify the attention pattern i.e. we pack sets of query and keys together and then define attention on these blocks. It is easier to explain this process using the example shown in Fig. 3. Suppose, there are 1212 query and 1212 key vectors to attend to. Using a block size of 22, we split the query matrix into 12/2=612/2=6 blocks and similarly the key matrix into 12/2=612/2=6 blocks. Then the three different building components of BigBird are defined on the block matrix. In particular the three different components are:

Random attention: Each query block attends to rr random key blocks. In Fig. 3(a), r=1r=1 with block size 22. This implies that each query block of size 22 randomly attends to a key block of size 22.

Window local attention: While creating the block, we ensure that the number of query blocks and the number of key blocks are the same. This helps us in defining the block window attention. Every query block with index jj attends to key block with index j−(w−1)/2j-(w-1)/2 to j+(w−1)/2j+(w-1)/2, including key block jj. In Fig. 3(b), w=3w=3 with block size 22. It means that each query block jj (size 22 queries) attends to key block j−1,j,j+1j-1,j,j+1.

Global attention: Global attention remains the same as defined in Sec. 2, but we compute it in terms of blocks. In Fig. 3(c), g=1g=1 with block size 22. For BigBird-itc this implies that one query and key block, attend to everyone.

The resulting overall attention matrix is shown in Fig. 3(d). Unfortunately, simply trying to compute this attention score as multiplying arbitrary pairs of query and key vectors would require use of gather operation, which is inefficient. Upon closer examination of window and global attention, we observe that we can compute these attention scores without using a gather operation.

The resulting AA tensor of size ⌈n/b⌋×b×b\lceil n/b\rfloor\times b\times b can be reshaped to correspond to the block diagonal portion of the full attention pattern. Now to extend the attention from block diagonal to a window, i.e. where query block with index jj attends to key block with index j−(w−1)/2j-(w-1)/2 to j+(w−1)/2j+(w-1)/2, we make ww copies of the reshaped key tensor K′K^{\prime}. We “roll” each copy of key-block tensor incrementally along the first axis of length ⌈n/b⌉\lceil n/b\rceil as illustrated in Fig. 5. Multiplying these ww rolled key-block tensors with the query-block tensor would yield the desired window attention scores (Fig. 4(c)). Likewise the global component, we can always include the first gg blocks from key tensor corresponding to the global tokens. Finally, for the random attention, which is very small (r=3r=3 for all of our experiments), we resort to using gather ops (Fig. 4(d)). Also note by design, each query block attends to exactly rr random blocks.

Thus, the result of all the three components is basically a compact dense tensor K′′K^{\prime\prime} of size ⌈n/b⌉×(g+w+r)b×d\lceil n/b\rceil\times(g+w+r)b\times d as shown in Fig. 6. Computing the final attention score then just boils down to a dense tensor multiplication, at which TPU/GPU are very efficient. Specifically, we need to multiply Q′Q^{\prime} (size: ⌈n/b⌉×b×d\lceil n/b\rceil\times b\times d) and K′′K^{\prime\prime} (size: ⌈n/b⌉×(g+w+r)b×d\lceil n/b\rceil\times(g+w+r)b\times d) with a cost of O(n(g+w+r)bd)O(n(g+w+r)bd) to yield the desired attention score tensor of size ⌈n/b⌉×b×(g+w+r)b\lceil n/b\rceil\times b\times(g+w+r)b, which can be reshaped to obtain all the attention scores according to the BigBird pattern.

Appendix E NLP experiments details

We use four publicly available datasets Books [zhu2015aligning], CC-News [guu2020realm], Stories [trinh2018simple] and Wikipedia to pretrain BigBird. We borrow the sentencepiece vocabulary as RoBERTa (which is in turn borrowed from GPT2). We split any document longer than 40964096 into multiple documents and we join documents that were much smaller than 40964096. Following the original BERT training, we mask 15%15\% of tokens in these four datasets, and train to predict the mask. We warm start from RoBERTa’s checkpoint. We train two different models: BigBird-itc-base and BigBird-etc-base. The hyper-parameters for these two models are given in Tab. 8. In all experiments we use a learning rate warmup over the first 10,000 steps, and linear decay of the learning rate.

Similar to the norm, we trained a large version of model as well, which has 24 layers with 16 heads and hidden dimension of 1024. Following the observation from RoBERTa, we pretrain on a larger batch size of 2048 for this size. For BigBird-itc the block length was kept same as base size, but for BigBird-etc the block length was almost doubled to 169. All the remaining parameters were the same.

E.2 Question Answering

The detailed statistics of the four datasets used are given in Tab. 11. All the hyperparameters for BigBird, used for creating Tab. 2 are shown in Tab. 12 and those submitted to get Tab. 3 are shown in Tab. 13. We use two types of regularization in training:

We used a variant of contrastive predictive coding [oord2018representation] as a dual encoder model.

We use position embedding for itc and relative position encoding [shaw2018self] for etc.

Next, we will mention the dataset/task specific part of the model.

The data consists of each question with multiple evidence paragraphs. We filtered 16 QA where the answer was not in the given evidences. For BigBird-itc, we use first 128128 global tokens. For BigBird-etc, we have one global token for each question token, one for each evidence paragraph, and one for each sentence within the paragraph, for a maximum of 256256 global token. We use a dense layer on the output corresponding to global token of the evidence paragraph to predict whether its a supporting fact with a threshold over the output logits. The answer type (yes/no/span) is predicted with a single dense layer from the global CLS token. For span based answers, the spans are predicted with dense layers on the sequence with the distance between start and end positions to be no more than 30 words. The spans are ranked by sum of start and end logits.

Here also the data consists of question with supporting evidence, but in form of a single, potentially long, document and not multiple paragraphs. We largely follow the setup of [alberti2019bert]. For documents, that are longer than 4096, a sliding window approach is used with stride of 2048. We use CLS token at the beginning, followed by the question followed by a separator token followed by the document as input. For BigBird-itc, we make the first 128128 tokens as global. For BigBird-etc, we make a global token for CLS, question, and one token for each of the paragraphs. We train four predictors at the final layer to predict long answer start, long answer end, short answer start and short answer end respectively. Instead of independently predicting the start and end of answers we first predict the start and then predict the best end location beyond the start. For short answer, we limit the distance between start and end positions to be no more than 38 words. The answer type (null, yes, no, short, long) is predicted from CLS token output embedding. When the logit for a yes/no answer is higher than the logits for short, long or null answer, we replace the short answer with a corresponding yes/no text.

The data consists of question-answer pairs with Wikipedia articles as the “noisy” supporting evidence. We call them noisy because the given Wikipedia articles may or may not contain the answer. Moreover, the answer entities is not annotated to appropriate span in the article, rather all occurrences found using fuzzy string matching are listed. We use CLS token at the beginning, followed by the question followed by a separator token followed by the document as input. For BigBird-itc, we make the first 128128 tokens as global. For BigBird-etc, we make a global token for CLS, question, and one token for each sentence up to a maximum of 320 global tokens. Given the noisy nature of answer span, we follow clark2017simple for training. We use a dense layer on the sequence to predict the answer span for each article independently, with the distance between start and end positions to be no more than 16 words. For each article the span with maximum start logit + end logit is chosen. Then we normalize over all the documents associated with that question.

For each question in WikiHop, we are given upto 7979 candidates, and 6363 supporting paragraphs. In our BigBird-itc model, following beltagy2020longformer, we concatenate the answer and the question with special tokens, [q] Question [/q] [ans] Ans1 [/ans] …\ldots [ans] AnsN [/ans] along with the context. As the start of the text, always contains questions followed by answers, we make the first 128128 token attend globally. In BigBird-etc model, we do not need to insert special [ans], [/ans] etc. as we design global tokens appropriately. Along with global tokens for question, we have one per candidate answer up to a maximum of 430. Further, we linked answer tokens to their mentions using relative position label. Lastly, we use a dense layer that takes in the output vector corresponding to a candidate answer, and predicts a score for the current candidate to be the correct answer. We apply this dense layer to each candidate independently and the candidate with the best score is picked as our final answer.

It is worthwhile to note that explicitly designed attention connection in etc works slightly better, the random connection based itc is pretty competative.

E.3 Relationship to Contemporary Work

child2019generating introduced localized sliding window to reduce computation. A more recent version, which includes localized sliding windows and global tokens was introduced independently by Longofrmer[beltagy2020longformer]. Although BigBird contains additional random tokens, there are also differences in the way global and local tokens are realized. In particular even when there is no random token, as used to get SoTA in question answering, there are two key differences between Longformer and BigBird-etc (see [ainslie2020etc]):

We use global-local attention with relative position encodings enables it to better handle structured inputs

Unlike Longformer, we train the global tokens using CPC loss and learn their use during finetuning.

E.4 Classification

We experiment on datasets of different lengths and contents, as listed in Tab. 15. In particular, we look at sentiment analysis (IMDb [maas2011learning] and Yelp-5 [zhang2015character]) task and topic assignment (Arxiv [he2019long], Patents [lee2020patent], and Hyperpartisan [kiesel2019semeval]) task. Following BERT, we used one layer with cross entropy loss on top of the first [CLS] token from the BigBird encoder consuming 4096 tokens. We report the results of document classification experiments in Tab. 15. We compare against state-of-the-art (SoTA) methods for each dataset and plain RoBERTa model with 512 tokens truncation. In all experiments we use a learning rate warmup over the first 10% steps, and linear decay of the learning rate and detail list of remaining hyperparameters are provided in Tab. 14. For better quantitative evaluation, we compute the fraction of the dataset that exceeds 512 tokens, i.e. the length at which the document are often truncated. We see that gains of using BigBird are more significant when we have longer documents and fewer training examples. For instance, using base sized model, BigBird improves state-of-the-art for Arxiv dataset by about 5%\bm{5\%} points. On Patents dataset, there is improvement over using simple BERT/RoBERTa, but given the large size of training data the improvement over SoTA (which is not BERT based) is not significant. Note that this performance gain is not seen for much smaller IMDb dataset. Along with experimental setup detail, we present detailed results in Sec. E.4 which show competitive performance.

The General Language Understanding Evaluation (GLUE) benchmark [wang2018glue], test language models on 8 different natural language understanding tasks. We used the same training parameters as mentioned in https://github.com/pytorch/fairseq/blob/master/examples/roberta/README.glue.md. Our model parameters are b=64,g=2×b,w=3×b,r=3×bb=64,g=2\times b,w=3\times b,r=3\times b ( we used the BigBird-itc base model pretrained on MLM task). We compare the performance of BigBird to BERT, XLNet [yang2019xlnet] and RoBERTa in Tab. 16. We find that even on task that have a much smaller context, our performance is competitive to full attention models.

E.5 Summarization

As discussed in Sec. 4.1, given the small length of output sequence, we used sparse BigBird attention only for encoder, while keeping the full attention for decoder. The number of hidden layers, number of heads, and hidden dimension is same for encoder and decoder. The hyperparameters are detailed in Tab. 17. We summarize our result in Tab. 20. In all experiments, we use a learning rate warmup over the first 10,000 steps, and square root decay of the learning rate.

Following success of several recent works [rothe2019leveraging, liu2019roberta], we warm start our encoder-decoder BigBird transformer model with pretrained weights and the weights between encoder and decoder are shared. In particular, the query/key/value matrix of self-attention and all the feedforward layers are shared between encoder and decoder. The only variable that is initialized randomly is the encoder-decoder attention. For base sized model, we utilize our MLM pretrained model on 4096 sequence length from Sec. E.1, which is in turn initialized using the public RoBERTa checkpoint. For the large size model, we lift weight from the state-of-the-art Pegasus model [zhang2019pegasus], which is pretrained using an objective designed for summarization task.

To check if sparse attention causes significant degradation as compared to full attention, we further experiment on two shorter but popular datasets, where full attention can be used without significantly truncating the document. The statistics of these two datasets are in Tab. 19. We see that our performance is competitive, which shows that sparse attention can achieve similar performance to a full attention models.

Appendix F Genomics experiments details

In this section we provide details of the experimental setup for BigBird on genomics data.

We try to keep the experimental setup as close to a typical NLP pipeline. In this regard, we take human reference GRCh37https://www.ncbi.nlm.nih.gov/assembly/GCF_000001405.39 and convert it into documents D\mathcal{D}. Each document d∈Dd\in\mathcal{D} is a sequence of sentences, where each sentence is a sequence of fragments of DNA. We construct the documents as follows:

Start with empty document set D=∅D=\emptyset.

For each chromosome CC, repeat the following procedure 10 times.

Pick uniformly at random a starting point qq between base pairs 0 and 5000 from the 5’ end.

Pick uniformly at random ss a number between 50 and 100 to denote number of sentences per document.

Constructs a document dd containing ss sentences using consecutive base pairs (bps). The length of each sentence is chosen uniformly at random between 500-1000. Thus the resulting document has 25,00025,000 - 100,000100,000 bps.

By this procedure we end-up with approximately 450K450K documents.

Next we run sentencepiece [kudo2018sentencepiece] tokenization on the resulting documents. In particular, using 5 characters as the building blocks (four for bases - A, T, C, G and one for missing symbol N), we construct a byte pair encoding table of size 32k, with each token representing 8.78 base pairs on average.

Using the above constructed documents, we construct a dataset for two pretraining tasks following devlin2018bert:

Masked Language Model (MLM): In order to train a deep bidirectional representation, BERT training introduces the MLM task, where we simply mask out 15% of the input tokens at random, and then predict those masked tokens. We can simply replace such masked out of the tokens with a [MASK] placeholder, but it leads to a distribution mis-match for downstream tasks which will not have such placeholders. To mitigate with this issue, out of the 15% of the tokens selected for masking:

80% of the tokens are actually replaced with the token [MASK].

10% of the time tokens are replaced with a random token.

10% of the time tokens are left unchanged, but are still predicted at output.

We run this entire sequence through the BigBird transformer encoder and then predict corresponding to the masked positions, based on the context provided by the other non-masked tokens in the sequence.

Next Sentence Prediction (NSP): In order to understand relationship between two sequences, BERT training introduces the NSP task, where we predict if a given pair of sequences are contiguous or not. During training the model gets as input pairs of sequences separated by [SEP] token along with a [CLS] token at the start. Overall the input pattern is: [CLS] sequence A [SEP] sequence B [SEP]. For 50% of the time the second sequence comes from true sequence after the first one. Remaining 50% of the time it is a a random sequence from the full dataset. The model is then required to predict this relationship using the output corresponding to the [CLS] token, which is fed into a simple binary classification layer.

The sequence of steps is visually elaborated in Fig. 9. The model is trained with both MLM and NSP together. Training hyperparameter is provided in second columns of Tab. 21. In all experiments we use a learning rate warmup over the first 10,000 steps, and linear decay of the learning rate.

We additionally performed a simple ablation study to validate the hypothesis, that similar to NLP, having a larger context improves performance. We use MLM task described above to test how BigBird performed with sequences of different length. Accuracy on MLM task with increasing sequence length is shown in Fig. 8. Not only longer context improves final accuracy, it also leads to faster learning, as we have now more opportunities for masking.

F.2 Promoter Region Prediction

The promoter region plays an important role in transcription initiation and thus its recognition is an important area of interest in the field of bioinformatics. Following oubounyt2019deepromoter, we use datasets from Eukaryotic Promoter Database (EPDnew) [dreos2013epd], which contains 29,597 promoter region in the human genome. Around the transcription start site (TSS), we extract a sequence of 8000 bp (-5000 +3000 bp) from the human reference genome GRCh37. Since EPDnew uses newer GRCh38, we convert to GRCh37 coordinates using LiftOver [kent2002human].

Following oubounyt2019deepromoter for each promoter region example, a negative example (non-promoter sequences) with the same size of the positive one is constructed as follow: The positive sequence is divided into 20 subsequences. Then, 12 subsequences are picked randomly and substituted randomly. The remaining 8 subsequences are conserved. This process is illustrated in Figure 1 of [oubounyt2019deepromoter]. Applying this process to the positive set results in new non-promoter sequences with conserved parts from promoter sequences (the unchanged subsequences, 8 subsequences out of 20). These parameters enable generating a negative set that has 32 and 40% of its sequences containing conserved portions of promoter sequences.

We prefix and append each example with [CLS] and [SEP] token respectively. The output corresponding to the [CLS] token from BigBird transformer encoder is fed to a simple binary classification layer. We fine-tune the pretrained BigBird from Sec. F.1 using hyper-parameters described in Tab. 21. We note that high performance is not surprising due to the overlap in the nature of negative example generation and MLM pretraining.

F.3 Chromatin-Profile Prediction

The first step of sequence-based algorithmic framework for predicting non-coding effects is to build a model to predict, large scale chromatic profile [zhou2015predicting].

In this paper, we use the dataset provided in zhou2015predicting http://deepsea.princeton.edu/media/code/deepsea_train_bundle.v0.9.tar.gz, to train BigBird to predict the chromatic profile.

Each training sample consists of a 8,000-bp sequence from the human GRCh37 reference genome centered on each 200-bp bin and is paired with a label vector for 919 chromatin features. As before, we prefix and append each example with [CLS] and [SEP] token respectively. The output corresponding to the [CLS] token from BigBird transformer encoder is fed to a linear layer with 919 heads. Thus we jointly predict the 919 independent binary classification problems. We fine-tune the pretrained BigBird from Sec. F.1 using hyper-parameters described in Tab. 21. As the data is highly imbalanced data (way more negative examples than positive examples), we upweighted loss function for positive examples by factor of 8.

We used training and testing split provided by zhou2015predicting using chromosomes and strictly non-overlapping. Chromosome 8 and 9 were excluded from training to test chromatin feature prediction performances, and the rest of the autosomes were used for training and validation. 4,000 samples on chromosome 7 spanning the genomic coordinates 30,508,751–35,296,850 were used as the validation set.

As the predicted probability for each sequence in DeepSea zhou2015predicting was computed as the ensemble average of the probability predictions for the forward and complementary sequence pairs, we also predict using an ensemble of two BigBird model trained independently.