$O(n)$ Connections are Expressive Enough: Universal Approximability of Sparse Transformers
Chulhee Yun, Yin-Wen Chang, Srinadh Bhojanapalli, Ankit Singh Rawat, Sashank J. Reddi, Sanjiv Kumar
Introduction
Transformer networks and their variants have played a key role in the recent advancement of the state of the art in many natural language processing tasks, such as machine translation , language modeling , and question answering . The key component of these networks is the self-attention layer , which updates the embeddings of the input tokens based on their context. Naturally, the self-attention layer also plays the key role in the analysis of Transformers ; for example, Yun et al. show that Transformers can approximate any continuous sequence-to-sequence functions (i.e., universal approximation), by proving that self-attention layers can compute contextual mappings of the input embeddings.
On the other hand, the self-attention layer is also the main bottleneck in scaling these models. It involves computation of pairwise inner products between input tokens, which results in quadratic computational complexity in the length of the input sequence . To mitigate this issue, researchers have developed methods to sparsify the pairwise interactions/connections in self-attention layers to reduce the computational complexity and/or improve model interpretability, and have shown successful empirical results on tasks with long sequence lengths . For example, Child et al. propose sparse Transformers for sequence generation. One of the sparsity patterns considered in is the Strided pattern, where the sparse attention layers alternate between two patterns: each token attends to only i) local neighbors, and then ii) one after every tokens in a strided manner. By choosing , they propose sparse attention layers with connections and show improvements on both speed and performance over the dense Transformer.
In the existing results, the rule of thumb for designing sparsity patterns (e.g., Strided) is connectivity; the intuition is that if each token can attend to the other tokens in multiple “hops,” then the resulting sparse Transformers do not lose much expressive power. However, there has been no formal justification for this intuition. How does sparsifying the interaction in the self-attention layers affect the model’s expressive power and ability to learn? What are the sparsity levels at which the model still retains its rich expressive power, and how is it affected by the sparsity pattern? Such fundamental questions about sparse attention models still remain unanswered.
In this paper, we take the first step towards a theoretical understanding of sparse Transformers.
We propose a unified framework to analyze sparse Transformers, which generalizes the existing approaches that sparsify attention layers (§ 3.1).
We propose a set of intuitive conditions on the sparsity pattern (Assumption 1) and the probability map (Assumption 2). Then, in Theorem 1, we show that Sparse Transformers, of fixed width and arbitrary depth, satisfying these conditions are universal approximators of any continuous sequence-to-sequence functions for any given fixed sequence length (§ 3.2 and § 3.3).
We next show some examples of existing sparse Transformers that satisfy these conditions, and hence have universal approximability (§ 3.4). Surprisingly, we show that there are sparse Transformers with only connections per self-attention layer (instead of ) that have enough expressive power to approximate arbitrary continuous functions (Corollary 2).
We report experimental results on standard NLP tasks using sparse Transformers, comparing different sparsity patterns/levels (§ 5).
Preliminaries and related works
In this section, we summarize the notation we will use throughout the paper, give a brief overview of Transformers, and then discuss existing efforts to sparsify the self-attention mechanism.
2 Transformers and their universal approximation power
3 Sparse Transformers
The first category reduces computation by making sparse in a pre-determined manner. Each token in the sequence only attends to a fixed smaller set of other tokens instead of the whole sequence . In some papers, auxiliary tokens are added to improve connectivity between existing tokens while maintaining sparsity . One drawback of these approaches is that the sparsity pattern is independent of input, so it cannot adapt to the data. To remedy this issue, proposes to learn local attention span from data. In a concurrent paper, Zaheer et al. propose the BigBird sparsity pattern which falls into this category. For BigBird, the authors show its theoretical properties such as universal approximation and Turing completeness, as well as its superior empirical performance. We note that our paper focuses on universal approximation for a broader class of sparse Transformers, by proposing a unifying framework to analyze them.
The second category studies making sparse after the full has been computed . Here, the focus is not on the computational gain via sparsity, because the full score matrix has to be computed first; rather, the goal here is to make attention layers more interpretable, as well as to improve performance. This line of works modifies in (1a) to other probability maps, by using top- elements or adopting sparser variants such as sparselin-gen or -entmax . Compared to the first category, this approach has an advantage that sparsity patterns are adaptive to data.
The last category attempts to get the best of both worlds. This line of works tries to learn sparsity patterns from data using extra components predicting the connection between tokens, e.g., -means clustering , LSTM , or locality-sensitive hashing . This way, one can adaptively determine the sparsity patterns before computing the score matrix. However, the drawback of this approach is that one needs extra computation to train/run these additional components, which may be expensive.
Universal approximation theorem for sparse Transformers
In this section, we derive a unifying framework to study sparse Transformers. We then propose a set of conditions on the sparse self-attention layers, and prove that the sparse Transformers satisfying theses conditions are universal approximators of any continuous sequence-to-sequence functions. Finally, we show some examples of existing sparse Transformers that satisfy these conditions.
We modify the Transformer block in (1) to the following sparse Transformer block ():
where the sets , for and , define the sparsity patterns (formally defined below), which are indexed by . Moreover, the parameter dimensions stay the same as in (1).
Note that there are three main modifications from the dense Transformer.
(Sparsity patterns) Note that denotes the -th column of the -th sparse attention head. Unlike dense Transformers, the inner product of the -th query vector is taken only with , the key vectors of tokens in the set . Hence, instead of all tokens, the -th token computes attention scores with only tokens in . For , we refer to the collection of the index sets , or simply , as a sparsity pattern. As a result, is a linear combination of columns in , rather than the whole sequence.
(Probability map) After computing the attention score matrix, the dense Transformer (1) uses the softmax operator to get a column stochastic matrix. In the sparse Transformers, we generalize to . The probability map is any map that takes a matrix as input and outputs a column stochastic matrix.
As a sanity check, by choosing , for all , and , we recover the dense Transformer (1). Note also that the sparse Transformer formulation covers the first and second categories of existing results discussed in § 2.3. The first category corresponds to choosing a predetermined sparsity pattern(s) , while setting . The second category corresponds to opting for a probability map other than softmax , while maintaining for all .
In this paper, we assume for simplicity that all sparse attention heads in a single layer have identical sparsity patterns . However, since our result only requires two sparse attention heads per layer (as we will see in Theorem 1), our result can be easily extended to the case that allows multiple sparsity patterns in a single layer.
Similar to in § 2.2, we define the class of functions represented by sparse Transformers. We hide the dependence of this class on the sparsity patterns and probability map to simplify the notation.
2 Conditions on sparsity patterns and probability map
In this section, we define a set of conditions on the sparsity patterns and the probability map that ensures that the sparse Transformer universally approximate the function class (cf. § 2.2).
For and the index sets , we define a sequence of sets in a recursive way:
The set is the set of all tokens that the -th token can directly/indirectly attend to, after sparse attention layers with sparsity patterns cycling through . We now state our conditions on sparsity patterns.
The sparsity patterns satisfy the following:
For all and , we have .
There exists a permutation such that, for all , .
Assumption 1.1 is equivalent to saying that every token always attends to itself. Assumption 1.2 requires that there is a chain of direct connections that covers all tokens; note that the set is the set of all tokens that the -th token directly attends to. To elaborate more about the chain, consider a directed graph with vertices corresponding to the tokens. For any , we add a directed edge . Given a graph constructed this way, Assumption 1.2 requires that the graph has a Hamiltonian path . Assumption 1.3 requires that after sparse attention layers, every token can attend to all the other tokens, either directly or indirectly.
As we discuss in § 3.4, the statements in Assumption 1 are natural enough to be satisfied by many existing sparsity patterns studied in the literature. In fact, Assumption 1.3 is necessary for universal approximation. If , , and , then the first token never attends to the second, so this sparse Transformer cannot approximate a function whose first output token is dependent on both input tokens. The other two assumptions are required in parts of our proof, which involve “propagating information” over all the tokens in a sequential manner.
We now state the assumption on the probability map . For this, we define to be the hardmax operator, which outputs the one-hot representation of the entry for each column of the input matrix. Since is a column-wise operator that outputs a column-stochastic matrix, we state the assumption for the operation of on a single column.
For any and , such that, for any column input satisfying (where ), we have and .
Assumption 2 requires that, for inputs that have some margin between the unique maximum entry and the other entries, can closely approximate the behavior of the hardmax operator by scaling its input by a positive factor . This assumption is satisfied by softmax and other sparse variants such as sparselin-gen and -entmax, as we show in § B of the supplementary material.
It is straightforward to check that the dense Transformer, which corresponds to , , and in our framework, satisfies both Assumptions 1 and 2.
3 Sparse Transformers are universal approximators
The key justifying intuition for adopting sparse attention layers is that, if each token can attend to the other tokens in multiple hopsNote that this corresponds to our Assumption 1.3., then these models do not lose too much expressive power. However, turning this intuition into a rigorous analysis is not straightforward. Moreover, recent results show that limited width can render universal approximation impossible even with arbitrary depth , highlighting the challenges in analyzing sparse (limited “width”) Transformers.
We now state our main theorem, which shows that if the sparsity patterns and the probability map satisfy Assumptions 1 and 2, sparse Transformers with attention heads of size , and hidden layer width are universal approximators of continuous sequence-to-sequence functions on any compact domain (recall that denotes the class of such continuous functions).
Consider any , and the class of sparse Transformers (cf. (3.1)) with the underlying sparse attention layers satisfying Assumptions 1 and 2. Then, for any and , there exists a function such that
As discussed earlier, dense Transformers satisfy Assumptions 1 and 2, which means that Theorem 1 subsumes the existing result for dense Transformers. We note that the required , , and in Theorem 1 are independent of , , or the sparsity patterns. We provide a high-level proof sketch of Theorem 1 in § 4.1. There, we also discuss how many layers are sufficient for -approximation of , and show that Theorem 1 requires only times more self-attention layers than Yun et al. .
We would like to emphasize that Theorem 1 provides the first formal evidence that well-designed sparse attention layers do not limit Transformer’s universal approximation power. In § 3.4, we show a surprising fact that some existing sparse self-attention layers with only connections (as opposed to in regular self-attention layers) retain enough expressive power to approximate . Combined with the number of layers analyzed in § 4.1, this means that our analysis reduces the connections per layer from to , with only times more attention layers. This advantage of sparse Transformers over their dense counterpart becomes even stronger with increasing sequence length , providing a theoretical support for the adoption of sparsity for tasks with long sequence lengths.
On a final note, Theorem 1 views the sequence length as a fixed constant. Hence, our result does not contradict a recent paper by Hahn which studies the limitation of Transformers for varying . Also, our analysis applies to the encoder part of the Transformer network .
4 Analysis of existing sparse Transformers
By Theorem 1, any sparse Transformer that satisfies our Assumptions 1 and 2 has universal approximation ability. In this section, we give some examples of such sparse Transformers.
Child et al. propose two kinds of -step sparsity patterns (i.e., ) for sequence generation tasks, namely Strided and Fixed patterns. We consider the extension of their auto-regressive patterns (i.e., attending only to past tokens) to the whole sequence. In the Strided pattern, a token first attends to its neighbors and then attends to one token after every tokens in a strided manner. The sparsity pattern for the -th token reads
In the Fixed pattern, we divide the token into segments of length . A token in a segment has access to other tokens in the same segment, and then the last tokens of the other segments:
The Strided and Fixed patterns satisfy both Assumption 1 and 2 for all values of . Specifically, Assumption 1.3 holds with , because any token can directly/indirectly access all the tokens in two hops. As for Assumption 1.2, the identity permutation suffices to satisfy the assumption for both patterns. By choosing , sparse Transformers with the Strided and Fixed patterns achieve universal approximation power with connections per attention layer.
Guo et al. consider the Star sparsity pattern where they add an auxiliary relay token that attends to all the tokens, and the other tokens attend only to neighboring tokens and the relay token. There is only one sparsity pattern, so . The Star sparsity pattern can be written as
where . For any fixed , this sparse Transformer has connections per attention layer, and it satisfies both assumptions. Specifically, Assumption 1.2 is satisfied with the identity permutation, i.e., for . Since any token can access other tokens within two hops, Assumption 1.3 is satisfied with . This demonstrates that connections per layer suffice for sparse attention layers to have universal approximation power. One can similarly check that the sliding window sparsity patterns with/without global attention, proposed in Longformer , also satisfy the assumptions with connections. For the BigBird sparsity pattern , it is also straightforward to check that a combination of its window attention and global attention satisfies Assumption 1 with connections. We state this interesting observation as a corollary below.
There exist sparse Transformers with connections per self-attention layer that are universal approximators in the sense of Theorem 1.
Recall that another line of results that replaces softmax with sparse variants also fits into our formulation, with and . As we show in § B, these alternative ’s satisfy Assumption 2. Thus, by Theorem 1, these models also have the universal approximation property.
Proof sketch and discussion
Step 2. We then approximate with a sparse Transformer network with a slightly modified architecture. In this architecture, we replace in the feed-forward layer with any piecewise linear activation , where denotes the class of (possibly discontinuous) piecewise linear functions with three pieces. We also replace in the sparse attention layer with the hardmax operator. We refer to the function class represented by the modified sparse Transformer as . By a careful construction, Lemma 3 shows that any can be exactly represented by the modified Transformer. To this end, we first carefully choose the positional embedding . We then quantize the inputs using feed-forward layers (Lemma 6), construct a contextual mapping using self-attention layers to map the quantized inputs to unique “ids” (Lemma 7), and then construct a value mapping with feed-forward layers to map the ids to desired output values (Lemma 8). See § D and § E in the supplementary material for details.
Step 3. The final step is to approximate the function with a sparse Transformer . This is done by approximating and with and , respectively, while carefully bounding the accumulation of errors introduced by the approximation. See § F in the supplementary material for the details.
For in Lemma 3, there exists such that .
Combining these three steps, we establish that . ∎
How many layers are sufficient? In § D, Lemmas 6–8 show that we need sparse Transformer blocks (2) for quantization, for the contextual mapping, and for the value mapping. Recall that is from (2), is from Assumption 1, and is from Step 1 above. In comparison, § C of shows that the dense counterpart requires , , and Transformer blocks (1) for the three corresponding lemmas. Note two observations: 1) The value mapping dominates the depth, and its depth requirements are identical for the two cases; and 2) For contextual mappings (where the attention layers are used), we need roughly times more layers for sparse models. Recall from § 3.4 that is usually a small constant. These observations mean that sparse Transformers can achieve universal approximation using depth of the same order in , and as the dense Transformers.
2 Key challenges in the proof
While the high level outline of the proof is similar to the one for dense Transformers , the proof in crucially relies on having all connections for computing attention in each layer, which we do not have in sparse Transformers. The sparsity in attention mechanism and the choice of general probability map pose nontrivial challenges in the proof. We highlight the key differences below.
Establishing the Step 2 of the dense result relies on constructing a contextual mapping using attention layers. A contextual mapping is a function that maps tokens in different sequences to unique values, thereby allowing Transformers to distinguish the same token appearing in different contexts. A crucial ingredient in the construction of such a mapping is a shift operation implemented with two attention heads in an attention layer. This shift operation involves each token taking the maximum and minimum over the entire sequence, which obviously cannot be done with sparse Transformers as it would require each token to attend to all the other tokens in the sequence. We circumvent this issue by carefully choosing the positional embedding dependent on (cf. Assumption 1.2), and ensuring that a similar shift operation is applied in a desired order even under sparsity.
As the final phase of the contextual mapping in , a single attention layer shifts the entire sequence by the maximum over the sequence. Again, this cannot be directly implemented due to sparsity. Using Assumption 1.3, we instead prove that by stacking sparse layers, one can successfully implement a similar operation that shifts the entire sequence by the maximum over the whole sequence, up to some controlled errors. This way, we overcome the difficulties posed by the sparsity and construct a new version of contextual mappings. The details can be found in § E.2 of the supplementary material.
Moreover, the proof of Step 3 in uses the simple fact that softmax can approximate hardmax arbitrarily closely. Since we do not restrict ourselves to softmax and generalize the probability map, a more careful argument is required. Since there are many layers in the network , it turns out that approximating it with an original sparse Transformer in requires carefully controlling the approximation errors accumulated over layers. The proof of Lemma 4 in § F of the supplementary material shows that this is indeed possible by utilizing Assumption 2.
Experiments
We now present our experimental study comparing different design and implementation choices, including sparsity patterns and levels, on four tasks: i) a synthetic copying task, ii) language modeling, iii) translation, and iv) GLUE tasks. Our goal is to understand the effect of such choices while employing sparse Transformers to the tasks with small sequence lengths, complementing the existing results for sparse Transformers on long sequence tasks.
We consider four sparsity patterns: Strided (4), Fixed (5), Star (6) and Random. The first three patterns are proposed in and ; we test them for different values of . In case of the Random pattern, given a sparsity level, we make connections uniformly at random. Following , Strided and Fixed patterns are tested for three different head configurations: i) Sequential, where the sparse attention layers alternate between and , as described in the previous sections; ii) Union, where all sparse attention layers use the sparsity pattern ; and iii) Multihead, where half of the attention heads in every attention layer use and the other half use . Note that, given the same sequence length, Union is less sparse than the other two configurations. Thus, to ensure fair comparisons, we compare different configurations based on their sparsity levels.
We use maximum sequence length 256 in all our experiments, except 128 for GLUE tasks. For the copying task, we experiment with only one sparse Transformer block (cf. Eq (2)), with varying numbers of attention layers with attention heads. For language modeling and translation, we use the Tensor2Tensor framework and employ 12-block and 6-block (respectively) Transformers with attention heads per block. For GLUE tasks, we experiment with the model. For more details of the setup, see § G of the supplementary material.
2 Results
Copying task. We consider a synthetic copying task proposed in , where the input sequence has the format , where is a 127 length sequence of symbols in $$. The models have to predict (copy) the second part, given the first half of the input. This task tests the ability of sparse Transformers to communicate the information. Table 1 presents the results for this task. Except for the Star and Random patterns, we can see that the networks learn to copy the sequences with four sparse attention layers. One possible explanation for the bad performance of Star is that, except for the relay token, it only attends to local neighbors while the task requires to copy distant tokens.
Language modeling. We conduct the language modeling experiments on the One Billion Word Benchmark which has almost one billion tokens and a vocabulary of more than 800K unique tokens. In Figure 1a, we plot the perplexity against the sparsity level. We observe that the Strided pattern and the Star achieve the best performance across all sparsity levels. For both the Strided and Fixed patterns, the Union configuration shows the best performance.
Translation. For the translation task, we train the model on WMT18 English-Czech (en-cs) dataset and test it on the Newstest 2015 dataset. We plot the BLEU score against the sparsity level in Figure 1b. We apply the same sparsity pattern to both the encoder and the decoder. The Strided and Fixed patterns with Union configuration show the best scores, which are similar to the dense attention. The Union configuration is also the least sensitive to the sparsity levels.
GLUE Tasks. We experiment with the model and report results on two sentence-pair classification tasks: MNLI (Figure 2a) and XNLI (Figure 2b). We plot the average accuracy of three runs on the dev set against the sparsity level. Additional results of the CoLA and MRPC tasks are reported in § H of the supplementary material.
In all tasks, the Random pattern performs worse than the deterministic patterns, demonstrating the need for a careful design of sparsity patterns. Overall, our experiments suggest that the design of the optimal sparsity patterns is heavily dependent on specific tasks. For example, the Star pattern shows the best performance on the language modeling task, while having trouble with copying, translation, and BERT experiments. Among the three head configurations tested for Strided and Fixed, the Union performs the best in language modeling and translation but suffers in BERT tasks. In translation experiments, we see an interesting trend that the performance of Multihead configuration improves as sparsity increases. We conjecture that this is due to the fact that in Strided and Fixed, we have and (cf. Eqs (4) and (5)), so the sparsest choice of is the one with the best “balance” between and .
Conclusion
Recently, sparse Transformers have received a lot of attention as they enable more efficient/faster attention mechanisms for the tasks with very long sequence lengths. We take an initial step to provide a theoretical understanding of these models. We provide a unifying framework that captures existing sparse attention models, and prove a universal approximation theorem for sparse Transformers which holds under intuitive conditions on sparsity patterns and probability maps. We also carry out experiments comparing different sparsity patterns and levels on standard NLP tasks. We hope that this work will shed light on the understanding of sparsity in attention layers, and provide guidance for the design of sparse attention models.
Broader Impact
This work studies theoretical aspects of a class of widely used neural network models in NLP and related areas. Since we do not propose a new method nor a new dataset, we expect that the impact of this work on ethical aspects and future societal consequences will be small, if any. Other than that, this work brings new insights into the sparsity in attention models, hence may make an impact on the study of faster and more efficient NLP models.
Acknowledgments and Disclosure of Funding
CY acknowledges partial support as a graduate Research Assistant from the NSF Grant (CAREER 1846088). CY also acknowledges Korea Foundation for Advanced Studies for their support.
References
Appendix A Outline and notation
The supplementary material is organized as follows. First, § B proves that the softmax operator as well as its sparse versions indeed satisfy Assumption 2. Next, § C provides formal statements of Step 1 in the proof sketch (§ 4.1). The outline of proof of Lemma 3 (Step 2 in the proof sketch) is presented in § D, followed by a separate section (§ E) proving the three key sublemmas in the proof. The proof of Step 3, Lemma 4, is given in § F. Lastly, § G and § H present the detailed setup of our experiments and additional experiment results, respectively.
Appendix B Sparse probability maps satisfy Assumption 2
In this section, we show that the softmax operator as well as the probability maps used to replace softmax in the existing approaches, namely softmax with only top- inputs , sparselin-gen , and -entmax , all satisfy Assumption 2. We restate the assumption for reader’s convenience: See 2 As in the assumption, we only consider the operation of these probability maps on a single vector, as they are applied column-wise. For each of the probability maps, we will show that for any and , we can choose that satisfies the conditions of Assumption 2.
We assume without loss of generality that the entry of is in decreasing order, where the first two entries satisfy . For any such and any , our aim is to show the existence of such that . Then, follows.
Now, since for , note that
Since is an increasing function in , one can increase sufficiently large to make it greater than .
The same argument holds for the softmax with top- inputs, used in . By the assumption on , entries are the top components. Thus,
can be satisfied by choosing large enough .
B.2 Sparselin-gen
We now consider the case where is sparselin-gen , which was used to sparsify the attention score matrices in . Given a regularization parameter , the sparselin-gen used in is defined as
Now, assume without loss of generality that the entry of is in decreasing order, where the first two entries satisfy . For any such and any , our aim is to show the existence of such that . This is done by choosing . To see this, notice that if ’s are in decreasing order, then are also in decreasing order. Now consider
If , then for all , and . If , then
B.3 α𝛼\alpha-entmax
Next, we consider the case where is -entmax , which was used to sparsify the attention score matrices in . Given a parameter , the -entmax is defined as
where is the probability simplex and is the Tsallis continuous family of entropies
As shown in , the solution of -entmax is equal to softmax if , and otherwise () it is given in the form
Again, assume without loss of generality that the entry of is in decreasing order, where the first two entries satisfy . For any such and any , our aim is to show the existence of such that . This is done by choosing .
Note that due to our choice of . Then, we will show that with such a , must hold. For the sake of contradiction, suppose not: . Then, by monotonicity of , we have . This means
in particular, we have . However, recall that , which implies . This results in
thus contradicting . Therefore, must hold.
Appendix C Details of the Step 1 in the proof sketch (§ 4.1)
We start by formally defining the function class .
For any and , there exists a small enough such that there exists such that .
Note that for any we have , so we have
Appendix D Proof of Lemma 3 (Step 2 in § 4.1)
In this section, we describe in further details how modified sparse Transformers (the class ) are able to exactly express arbitrary piecewise constant functions in . We show that we can compute a contextual mapping of the entire input sequences without relying on dense self-attention layers. The token-wise feed-forward layers then transform these contextual mappings to the desired output sequence.
Choose the positional embedding according to in Assumption 1.2. After addition, each column of the input are in disjoint intervals.
Given the input , a series of modified feed-forward layers quantizes it so that each entry of the quantized input has a value in (Lemma 6).
Next, a series of modified sparse self-attention layers takes the quantized input and implement a contextual mapping such that, for different quantized input sequences and , all the elements in and are distinct (Lemma 7).
Finally, a series of modified feed-forward layers maps each element in the context id to the desired output value of at the input (Lemma 8).
We defer the proofs of Lemmas 6, 7, and 8 to a separate section: see § E.
Before discussing the details of each step, we note that although a Transformer network stacks self-attention and feed-forward layers in an alternate manner, we can use a series of arbitrary number of the same layers, thanks to skip connections. The outline of the proof is similar to , but key component in their proof called selective shift operation relies on the fact that each token can attend to the entire sequence; this is not true in sparse Transformers, which poses a nontrivial challenge. We overcome this issue by a more careful construction of the positional embedding and sparse self-attention layers.
Recall from Assumption 1.2 that there exists a permutation such that for all , is one of the tokens that the -th token directly attends to. Using this permutation , we choose the columns of positional embedding in the following way:
As a result, the -th column of will be in the range , and similarly for . This means that the entries corresponding to different tokens lie be in disjoint intervals of the form , where .
D.2 Quantization by feed-forward layers
Note from the previous step that each entry of must be in . Next, we quantize this interval of input using to a set of -grid points . This allows us to deal with finite set of values, which proves useful in the later stages of the proof. The next lemma shows that the quantization can be carried out using a seried of the modified feed-forward layers.
D.3 Contextual mapping by sparse self-attention layers
This contextual mapping maps each unique sequence/context into different context ids, enabling the network to distinguish the same token appearing in different sequences.
D.4 Value mapping by feed-forward layers
D.5 Finishing the proof
Appendix E Proof of Lemmas 6, 7, and 8
The proof goes as follows. Using token-wise feed-forward layers, we implement the quantization function that quantizes the first row of the input. Then we stack another layers to quantize the second row, and so on.
For the first row, we add layers of the following form, for .
E.2 Proof of Lemma 7
One can consider a sparse self-attention layer that consists of two such heads, with :
The -th entry of reads
This means that for input columns satisfying only, shifts up the first entry of by the difference of maximum and minimum values of over the sparsity pattern , while leaving other columns intact. By choosing and properly, we can selectively modify certain columns without touching other columns; we refer to this operation as the sparse selective shift operation, and we will see later that this is indeed the key ingredient of our proof.
In fact, this operation is a sparse version of the selective shift operation used in . Since is usually only a small subset of , one cannot calculate the maximum and minimum of over the whole sequence, as done in . Instead, we use Assumption 1.2 and a more careful choice of to get around the restriction posed by sparsity.
The -th entry of reads
So, for each column , the all-max-shift operation shifts up the first entry of by the maximum value of over the sparsity pattern . Unlike the selective shift operation, the all-max-shift operation is applied to all the columns.
Recall that the any input to this step is in
E.2.2 Construction of layers
Given these preliminaries, we now describe our construction of . Recall from Assumption 1.2 that the permutation satisfies for . From this, for we let be any index such that . For simplicity of notation, let for and .
Next, starting from , we want to sequentially stack sparse selective shift operations
in increasing order of . That is, we want to add sparse attention layers with sparsity patterns that apply the selective shift operation to each possible value of . Recall that the sparsity patterns have to cycle from to , so we have to place other remaining sparsity patterns (whose indices are not ) in between the layers. This can be done by setting all the other sparse attention layers to be the identity. This way, we stack a total of sparse attention layers for , another for , and so on, up to .
After these layers, we further stack all-max-shift operations. For , we add all-max-shift operations of the form
E.2.3 Selective shift operations
First consider the first layers. Omitting layers that are identity, they are essentially selective shift operations for . Since is the set of possible values of , these layers perform selective shift operation on the -th column without changing the other columns. Each possible value of undergoes one and only shift operation (by the corresponding layer with ), by which the -th entry of the input is updated.
Recall by Assumption 1.2 that , and that and are the maximum and minimum over the whole sequence (see (8)). By Assumption 1.1 we also have . Since both and are in , the maximum and minimum value of ’s over are and , respectively. Therefore, the -th entry of the input matrix is shifted up as follows:
Let be the -th column after the shift operation has shifted to . Then, define
Note that because
which is true. Therefore, becomes the new maximum among the current values , and the new minimum element is .
We now consider the next layers, which are essentially for . They apply the shift operation to the -th column. Since we have , the shift operation similarly yields
We can also show , because
So after this operation and are the new maximum and minimum over the updated sequence .
The same process continues. The next layers shifts the -th columns and results in which is greater than . After the first layers, all columns except -th column have been shifted, resulting in satisfying
Let us denote the output of the -th layer as .
is one-to-one. Recall that for each column , the map is one-to-one. Also, permutation of columns is one-to-one, which implies that it suffices to show that the map is one-to-one.
Suppose we have two sequences and that map to the same value of . Then,
Suppose . Since they both lie inside , we have
Note that all the terms other than are of “coarser resolution.” For example, the first term
in the summation can only take values , so it can never cancel the difference and make the sum zero. This implies that must hold.
Next, suppose . Since we have ,
E.2.4 All-max-shift operations
Next, we explain the operation of the all-max-shift layers. Recall from Assumption 1.3 that any token can attend to all the other tokens after steps, either directly or indirectly. Also recall from the last subsection that the input to the first all-max-shift layer is , and the maximum entry of is , the unique id for input . From the statement of Lemma 7, the output after the all-max-shift operations for input is denoted as . In this subsection, we show that through all-max-shift operations, the maximum will propagate to all tokens and be a “dominant” term, which determines the interval that lies in. As a result, we can show Properties 7.1 and 7.2 of at the end.
Note that the unique id has the following upper bound:
where we used . A similar bound
also holds from a similar derivation. Next, recall from Assumption 1.3 the definitions
and that there exists such that, for all , . Finally, the following inequality will be useful throughout: for any integer ,
Let us now describe the operation that the all-max-shift layers , , carry out.
The input to the first all-max-shift layer is . Let the output of the layer be . Recall that consists of values , which are all strictly greater than 0 and strictly less than (by (10)). So, for each column , the layer update reads
where . After the update, is “dominated” by , meaning that for any ,
This is because the minimum gap between different values of is at least , and we have
so if , that solely determines the order because cannot reverse it. Also, by the definition of , for any index set we have
If , we move on to the second layer.
At the second all-max-shift, we have sparsity patterns . Let us the output of this layer as . For each column , the layer update reads
where . If we look at the update more closely, we can apply (13) and get
Again, the last term dominates the rest of the terms in , because the minimum gap between different values of is at least , and
The last inequality holds due to inequality (12), because
If , we move on to the third layer, which outputs . Similarly, we can show that is dominated by because the rest of the terms in is strictly upper-bounded
which can then be shown to be smaller than :
The last inequality is due to the fact that for , which can derived from (12). Repeating this process, after all layers we get , and is dominated by
This is because the remaining terms in can be strictly upper-bounded
which is then dominated by the smallest difference possible in :
The last inequality used , derived from (12).
E.2.5 Verifying Properties 7.1 and 7.2
After these all-max-shift operations, we define the output of the last all-max-shift layers to be the output of the function for input , i.e., .
This is because anything added by the all-max-shift operations is an integer multiple of , and for all . Recall that is the input matrix for the first max-shift operation, and that the components of are , which were shown to be distinct by (9). Since produce distinct outputs for a operation, they themselves have to distinct. This proves Property 7.1.
Also, by the “domination” argument in the previous subsection, the output has the property that for any column, lies inside an interval determined by , the unique id for the input :
E.3 Proof of Lemma 8
To prove this lemma, we implement a token-wise function that maps
This layer updates any column of its input that satisfies , without modifying any other columns that are out of this range.
Appendix F Proof of Lemma 4 (Step 3 in § 4.1)
In this section, we describe how the modified sparse Transformer network constructed in Lemma 3 can be approximated with an original sparse Transformer network . Recall that is a “modified” sparse Transformer network, which employ the hardmax operators in place of operators in sparse self-attention layers and piecewise linear activations instead of s in feed-forward layers. The goal of this lemma is to approximate the function with a standard sparse Transformer with accuracy . As the construction of consists of three steps, we will approximate each of them step by step. The whole intuition behind the proof is that as long as we are considering approximation, we can approximate and as closely as we want with and s, respectively. However, as the proof will show, controlling the aggregated error over layers is not a trivial job.
We first consider approximating from Lemma 6 with a standard feed-forward layer counterpart, . Recall from § E.1 that the modified feed-forward layers used in are of the form
for and . Note that the activation can be closely approximated by three s:
where . Note that except for an interval , and by shrinking this interval can be made arbitrarily small. Consider approximating the layers (14) with standard feed-forward layers, by replacing with its approximation . Let the resulting function be .
Then, it is easy to check that holds if all coordinates of are in the intervals of the form for some ; i.e., the intervals in which perfectly approximates . The Lebesgue measure of the set of such inputs is
Let us now consider approximating the contextual mapping in Lemma 7, constructed using the hardmax operators, with the standard sparse self-attention layers employing operator. We will call the approximation . Recall that satisfies Assumption 2: See 2 This means that can closely approximate in the sense that whenever the input vector to the operator has a maximum element by some margin , then the -th component of the output is close to , while the other components of are close to .
Recall that consists of two parts. The first part is a composition of sparse selective shift operations, and the second is a composition of all-max-shift operations. We will first examine how “errors” are introduced when is replaced with in both operations, discuss how the errors accumulate, and show how to choose the right and to control the errors in the approximation .
Recall that the key component in both the selective shift operation and all-max-shift operation is the sparse attention head , which computes its -th column as the following:
Now suppose we replaced with satisfying Assumption 2. Suppose each entry in differs at least by , which is true in the construction of . We choose and some , and corresponding . Then, replace with and define
If , it is easy to check that satisfies
Similarly, if , we have
Now consider the approximate sparse selective shift operator , implemented with . For , we define
For any column satisfying , we have
and for any column satisfying , we get
Recall that for the hardmax version, we had
From this observation, the approximation error of the selective shift operator on the -th entry of the output can be bounded as follows:
where we used for simplicity.
Next, we examine the approximation error of the all-max-shift operation introduced by replacement of with . Let us define the approximate all-max-shift operation :
From (15), we can check that the approximation error of the all-max-shift operation is bounded as
Given these approximation error bounds of single operations, we now analyze the accumulation of errors through multiple layers. We first consider the first self-attention layers in . Recall that they consist of selective shift layers for and identity layers. A natural way to approximate these layers with standard self-attention layers is to use approximate layers , with sufficiently large . As we have seen above, there is no error introduced by except for the first row. Thus, we will analyze the approximation error of for the first row only.
Let us remind the readers how the first selective shift operation (done by the first layers) originally worked in . The input to is , and we define and . Recall from Eqs. (7) and (8) in § E.2 that
and , so will undergo the selective shift by one of the self-attention layers, which updates the -th entry of the input. Let be the updated value of the column and . The new sequence satisfies
where the strict upper bound on is from Eq. (11).
In case of the approximation , we have seen that the error depends on the gap between maximum and minimum of ’s, and this gap may grow larger as error accumulates; in the worst case, it may grow exponentially. To see this, suppose and are the maximum and minimum value of ’s, and they go through a selective shift operation, but they do not belong to the range of the operation . Then, and will be updated to and , which are bounded by
showing that the gap may grow exponentially in the worst case:
In the error-less case (), for any input sequence , the maximum possible difference between maximum and minimum of is bounded above by , and after one selective shift operation was done on the -th column, the difference is then bounded by . Therefore, the worst-case possible error introduced by is bounded above by the sum of the worst-case errors calculated assuming that we started off with max-min difference . Using this observation, the error on each first-row entry of the sequence after the first layers is bounded above by
where a factor of is introduced because when the selective shift operation is applied to the -th column, it may introduce an error which is twice the magnitude of the error introduced to the other columns. We want to make (16) smaller than . By Assumption 2, we can always choose that satisfies the assumption for
Using such , we can control the total accumulated error by the first selective shift operations below :
Therefore, after the first selective shift layers, the accumulated error for each entry of the first row is at most .
We can also apply similar arguments to the remaining selective shift layers. For example, for the -th set of selective shift layers where the operation is done on -th column of the input, the gap between the maximum and the minimum, including the accumulated error from previous layers, is bounded above by . Therefore, for this set of layers, the maximum accumulated error is bounded by
So, choosing that satisfies Assumption 2 for and , we can control the accumulated error introduced by the layers below :
In total, the accumulated error by the first layers, which correspond to the selective shift operation part of the construction, is at most .
For all-max-shift operations, we approximate the hardmax all-max-shift operations with its -counterparts, . We can similarly bound the accumulated error in the all-max-shift operations. Recall from § E.2 that during the whole series of all-max-shift operations, the maximum entry in the sequence is upper-bounded by and minimum entry is lower-bounded by . Therefore, the gap between the max and min elements, taking into consideration the errors from selective shift operations, is bounded from above by . Then, using a similar argument as the select shift operation layers, the maximum error is bounded above by
and we want to make it smaller than . By Assumption 2, we can always choose that satisfies the assumption for
Using such , we can control the total accumulated error by the first selective shift operations below :
We now consider the approximation of the value mapping with standard feed-forward layers. In , we implemented the function with layers of the form
Since the output of contextual mapping and its approximation differ in only the first row and by , one can approximate each layer in by replacing with an approximation , implementable with four ’s:
Hence, using , we have
F.4 Finishing the proof
One can make close enough to 1 so that the second term is less than . This makes , hence finishing the proof.
Appendix G Experimental setup
We generated the synthetic dataset for the copying task. The input sequence to the copying task has the format , where is a 127 length sequence of symbols randomly sampled from the range of $$. The training set contains 100K sequences, while the testing set contains 10K sequences.
We implement the copying task as a masked-LM style prediction task by masking all the tokens in the second half of the sequence. For the test examples, each masked token is predicted independently. For the results reported in § 5, we experiment with bidirectional models, where each token can attend to both previous and future tokens.
The maximum sequence length is , and we use embedding dimension . The model has 1 to 4 attention layers with attention heads of size , followed by a feed-forward hidden layer of size . We train the model with the AdamW optimizer with weight decay and no dropout. We train the model using 3,000 warmup steps and a total of 500K training steps. The learning rate is . We use the batch size 1,024 on 8 TPUv3 chips.
For all sparsity patterns other than the Random pattern, we choose the segment length to be 16 for all patterns. This segment length results in the sparsest level for the Strided and Fixed patterns. In Table 1, we include the sparsity level as a reference. For this task, we report the prediction accuracy for all the tokens.
G.2 Language modeling
For the language modeling task, we train on the One Billion Word Benchmark which contains almost one billion tokens and a vocabulary of more than 800K tokens.
We use the Transformer model in the Tensor2Tensor framework . We use a 12-block (cf. (2)) Transformer, with embedding dimension , maximum sequence length , number of heads , head size , and feed-forward hidden layer size . Since language modeling task is auto-regressive (attending to only past tokens) in nature, we evaluate the (sparse) attention score matrices and mask them to be an upper-triangular matrix. We train the model with the Adafactor with weight decay. We train the model using 10K warmup steps and a total of 240K steps. We use the batch size 4,096 on 8 TPUv2 chips.
G.3 Translation
For the translation task, we train on the WMT18 en-cs datasets (Europarl v7, Common Crawl corpus, News Commentary v13, and CzEng), with a total of 15M pairs of sentences, and test on the newstest2015 en-cs dataset, with 2,656 pairs.
We use the encoder-decoder architecture and apply the sparse attention on both encoder and decoder. We use the Transformer model in the Tensor2Tensor framework and the same setup as the language modeling task, except for having 6 blocks in the Transformer networks, with head size and having autoregressive patterns only in decoders.
For this task, we report the cased BLEU score.
G.4 GLUE tasks
For the GLUE tasks, we use the pre-training and fine-tuning framework . Following Devlin et al. we first pre-train a model for 450K steps on the BooksCorpus (800M words) and the English Wikipedia datasets (2,500M words). We later finetune the model on data from each task separately. For each setting, we use the same sparsity pattern and head configuration in both the pre-training and the fine-tuning stages. The sequence length is in both stages.
We report the average accuracy of three runs on the dev set for all tasks. For each setting, we pre-train a model and run fine-tuning three times.
Appendix H Additional experimental results
We report additional experimental results in this section.
We include the results for the copying task using auto-regressive (unidirectional) models as in LM, where each token can only attend to previous tokens, in Table 2. In this case, the Star pattern cannot attend to the last replay token. Indeed, the Star pattern shows better performance when the model is bidirectional (cf. Table 1).
H.2 Translation
We present experimental results of the translation tasks on the WMT English-German and German-English datasets in Figure 3. We train on WMT18 (Europarl v7, Common Crawl corpus and News Commentary v13) and test on newstest 2015 datasets. The figures show similar trends to the results on the WMT en-cs dataset in Figure 1b.
H.3 GLUE tasks
Figure 4 presents the results comparing the sparsity patterns and the head configurations on the CoLA and MRPC tasks using the model. CoLA is a single-sentence classification task, asking if a sentence is a grammatical English sentence. MRPC is a sentence-pair classification task, where each example is a pair of sentences and the label indicates whether the sentences are semantically equivalent.