The Devil in Linear Transformer
Zhen Qin, XiaoDong Han, Weixuan Sun, Dongxu Li, Lingpeng Kong, Nick Barnes, Yiran Zhong
Introduction
Transformer models show great performance on a wide range of natural language processing and computer vision tasks Qin et al. (2022); Sun et al. (2022b); Cheng et al. (2022a, b); Zhou et al. (2022). One issue of the vanilla transformer model lies in its quadratic space-time complexity with respect to the input length. Various prior works attempt to alleviate this inefficiency Zaheer et al. (2020); Beltagy et al. (2020); Tay et al. (2020a); Kitaev et al. (2020); Child et al. (2019); Liu et al. (2022); Sun et al. (2022b). In this work, we focus on a particular subset of these methods, known as kernel-based linear transformers Choromanski et al. (2020); Wang et al. (2020); Katharopoulos et al. (2020); Peng et al. (2020); Qin et al. (2022) considering their desirable linear space-time complexity.
Despite their space-time efficiency, linear transformers are not always in favor for practical adoption, largely due to the degraded performance than the vanilla model. To address this issue, we take a close look at existing kernel-based linear transformers and identify two deficiencies that lead to such a performance gap.
Unbounded gradients. Most existing linear transformers inherit attention formulation from the vanilla transformer, which scales attention scores to ensure they are bounded within $$. However, we theoretically show that such a scaling strategy renders unbounded gradients for linear transformer models. As a result, the unbounded gradients empirically lead to unstable convergence as our preliminary experiments suggest. Attention dilution. Previous works Titsias (2016); Jang et al. (2016); Gao and Pavel (2017); Qin et al. (2022); Sun et al. (2022b, a) suggest that in vanilla transformer, softmax attention maps tend to be local. In contrast, as shown in Fig 2, we observe that linear transformers often trivially distribute attention scores over the entire sequence even in early layers. Due to this issue, which we refer as attention dilution, important local information is less well preserved in linear models, resulting in inferior performance. This negative impact of attention dilution is also evidenced by the performance drop in our controlled experiments if partly replacing vanilla attention in transformer layers with linear attention ones.
To mitigate these issues, we propose a linear transformer model, called TransNormer, which shows better performance than vanilla transformer on a wide range of task while being significantly faster during runtime, as shown in Fig. 1.
To avoid the unbounded gradients, we introduce NormAttention, which gets rid of scaling over attention matrices while appending an additional normalization only after the attention layer. The choice of the normalization operator is unrestricted, for example, LayerNorm Ba et al. (2016) or RMSNorm Zhang and Sennrich (2019) both serve the purpose. We show empirical results demonstrating that with NormAttention, the gradients are more stable during training, which in turn leads to more consistent convergence.
To alleviate the attention dilution issue, we modify the vanilla attention and allow each token to only attend to its neighbouring tokens, resulting in a diagonal attention. To mimic the behaviors on local semantics of the vanilla transformer, we employ the diagonal attention on early layers while using NormAttention for later ones. In this way, we encourage the model to capture both local and global language context. Note that our diagonal attention can be efficiently computed such that the overall linear space-time complexity of TransNormer is preserved.
We perform extensive experiments on standard tasks, where TransNormer demonstrates lower language modeling perplexities on WikiText-103 and overall higher text classification accuracy on GLUE than vanilla model and other competing methods. In addition, on the challenging Long-Range Arena benchmark, TransNormer also shows favorable results while being faster and more scalable with longer inputs during both training and inference time.
Background and related work
Pattern based methods Zaheer et al. (2020); Beltagy et al. (2020); Tay et al. (2020a); Kitaev et al. (2020); Child et al. (2019) sparsify the attention calculation with handcrafted or learnable masking patterns. Kernel-based methods adopt kernel functions to decompose softmax attention, which reduces the theoretical space-time complexity to linear. In this paper, we refer the kernel-based variants as linear transformers for simplicity.
In the kernel-based methods Choromanski et al. (2020); Katharopoulos et al. (2020); Peng et al. (2020); Qin et al. (2022); Zheng et al. (2022); Wang et al. (2020), a kernel function maps queries and keys to their hidden representations. Then the output of the linear attention can be rewritten as:
where the product of keys and values are computed to avoid the quadratic matrix. Existing methods mainly differ in the design of kernel functions. For example, Choromanski et al. (2020) and Katharopoulos et al. (2020) adopt activation function to process query and key. Wang et al. (2020) assumes attention matrices are low-rank. Peng et al. (2020) and Zheng et al. (2022) approximate softmax under constrained theoretical bounds. Qin et al. (2022) propose a linear alternative to the attention based on empirical properties of the softmax function.
These methods focus on either approximating or altering the softmax operator while preserving its properties. Compared with the vanilla transformer, these methods often trade performance for efficiency, usually resulting in worse task performance. In this paper, we argue that there are two essential reasons leading to such a performance gap, discussed in detail as follows.
The devil in linear attention
In this section, we motivate the design principles of TransNormer by providing theoretical evidence for the unbounded gradients, and empirical results showing the adverse influence of attention dilution.
Few work on linear transformers analyzes their gradients during training. Our first key observation is that kernel-based linear attention suffer from unbounded gradients, causing unstable convergence during training. In the following, we highlight the main theoretical results while referring readers to Appendix D for the full derivation.
Vanilla and linear attention differ mainly in their computation of token-wise similarities Note that is not directly computed in linear attention, but can still be represented in this unified form, see Appendix D for more detailed derivation. In vanilla attention, is computed as:
while for linear attentions, can be decomposed using a kernel function , such that:
Given the above definitions, the gradients of the attention matrix is derived as:
Therefore, for the vanilla attention, the partial derivative is:
andA detailed proof of the upper bound can be found at Appendix B.
Since can be arbitrarily large, the gradient of linear attention has no upper bound. On the other hand, we can also show that the gradient of linear attention has no lower boundThe proof can be found in Appendix C.:
The unbounded gradients lead to less stable optimization and worse convergence results in our preliminary studies.
2 Attention dilution
It is a known property of vanilla attention to emphasize on neighbouring tokens Titsias (2016); Qin et al. (2022). However, this property does not directly inherit to the linear transformer variants.
To quantify the attention dilution issue, we introduce a metric called locally accumulated attention score, which measures how much attention scores are distributed within the local neighbourhood of a particular token.
For an input sequence of length , consider a local neighbourhood centering around token of total length , with the ratio relative to the total input, the locally accumulated attention score for token is defined as . A higher score indicates the particular attention layer concentrates on the local neighbourhood, while a lower score tends to indicate the issue of attention dilution, where scores are distributed more evenly to local and distant tokens. For example, means that that 40% of the neighbors around ’th token contribute 60% of the attention score.
In Fig. 2 (a), we compare locally accumulated attention scores (y-axis) for vanilla transformer and linear transformer, with varying sizes of neighbourhood by ratio (x-axis). We show the average score over each position across the entire sequence. It can be seen that the area under the vanilla model curve is significantly larger than that of the linear model. This provides evidence that the vanilla attention is more concentrated locally, while the linear transformer suffers from the issue of attention dilution. This is further qualitatively supported by Fig. 2 (b), where the attention maps for vanilla model are more concentrated than the linear model.
Method
Based on the aforementioned observations, we propose a new linear transformer network called TransNormer that addresses the above two limitations of current linear transformers. The overall architecture is shown in Fig. 3.
Vanilla attention suffers less in attention dilution while linear attention is more efficient and scalable on longer sequences. This motivate us to design a method that exploits the best of the both worlds by using these mechanisms in combined.
Specifically, our network consists of two types of attention: DiagAttention for the early stage of the model and NormAttention for the later stage. The former addresses the attention dilution issue and the later aims to stabilize training gradients. Note that by properly reshaping the inputs, the diagonal attention can be efficiently computed in linear space-time, thus preserving the overall linear complexity.
2 NormAttention
As proved in Sec. 3, the scaling operation, i.e., the denominator in Eq. 4, in the linear transformers hinder the optimization due to the unbounded gradients. To solve this issue, we propose to remove the scaling operation in the linear transformers. However, as shown in Table. 1, directly removing the scaling operation leads to critical performance drop since the attention map becomes unbounded in the forward pass. Therefore, an alternative is required to bound both attention maps during forward and their gradients during backward passes in linear attentions.
Our proposed solution is simple yet effective. Given a linear attention, the attention without scaling can be formulated as:
We empirically find that we can apply an arbitrary normalization on this attention to bound it, which leads to our NormAttention as:
It can be proved that the gradients of NormAttention is bounded byThe full derivation can be found in Appendix D.:
where is the loss function, is the small constant that used in RMSNorm, is the embedding dimension and
To demonstrate the gradients stability of the NormAttention, we compare the relative standard deviation of gradients during each training iterations to other linear transformers and vanilla transformer. Specifically, we train our model for 50k iterations with RoBERTa architecture on the WikiText103 Merity et al. (2017) and obtain the relative standard deviation of all iterations’ gradients. As shown in Table 2, existing linear methods Choromanski et al. (2020); Katharopoulos et al. (2020) have substantially higher deviations compared to vanilla attention, which leads to inferior results. The NormAttention produces more stable gradients, which validates the effectiveness of our method.
3 DiagAttention
To better understand the design principles, we show in Table 3 that by replacing partial layers of linear transformers with vanilla attention, the performance on language modeling is evidently improved. The results also suggest that capturing more local information in early layers are more helpful than otherwise.
To this end, we leverage none-overlapped block-based strategy to reduce the space-time complexity of the vanilla attention. Based on the observation in Fig. 2, we utilize a strict diagonal blocked pattern to constraint the attention in a certain range. Since the attentions are calculated inside each block, the computation complexity of our diagonal attention is , where is sequence length , is the block size and is feature dimension. When , the complexity scales linearly respect to the sequence length . In subsequent sections, we use DiagAttention to refer to Diagonal attention.
We empirically find that applying DiagAttention to the later stages of a model hurts the performance as shown in Table. 9. It indicates that the model requires a global field of view in the later layers, which also justifies our choices of NormAttention in later layers of TransNormer.
Experiments
In this section, we compare our method to other linear transformers and the vanilla transformer on autoregressive language modeling, bidirectional language modeling as well as the Long Range Arena benchmark (Tay et al., 2020b). We also provide an extensive ablation study to vindicate our choice in designing the TransNormer.
We validate our method on two variants of the TransNormer. The TransNormer T1 uses the ReLA attention (Zhang et al., 2021) in the DiagAttention and the elu as the activation function in the NormAttention. The TransNormer T2 uses the attention (Vaswani et al., 2017) in the DiagAttention and the 1+elu as the activation function in the NormAttention.
For experiments, we first study the autoregressive language modeling on WikiText-103 (Merity et al., 2017) in section 5.2. Then in section 5.2 we test our method on bidirectional language modeling, which is pre-trained on WikiText-103 (Merity et al., 2017) and then fine-tuned on several downstream tasks from the GLUE benchmark (Wang et al., 2018). Finally, we test TransNormer on the Long-Range Arena benchmark (Tay et al., 2020b) to evaluate its ability in modeling long-range dependencies and efficiencies in section 5.2.
We implement our models in the Fairseq framework (Ott et al., 2019) and train them on 8 V100 GPUS. We use the same training configuration for all competitors and we list detailed hyper-parameters in Appendix F. We choose the FLASH-quad, FLASH (Hua et al., 2022), Transformer-LS (Zhu et al., 2021), Performer (Choromanski et al., 2020), 1+elu (Katharopoulos et al., 2020) as our main competing methods.
For the autoregressive language modeling, we use 6 decoder layers (10 layers for the FlASH/FLASH-quad) as our base model and all models are trained on the WikiText-103 dataset (Merity et al., 2017) for 100K steps with a learning rate of . We use the perplexity (PPL) as the evaluation metric.
For the bidirectional language modeling, we choose the RoBERTa base (Liu et al., 2019) for all methods. It consists of 12 encoder layers (24 layers for the FLASH and FLASH-quad to match the number of parameters). All models are pre-trained on the WikiText-103 (Merity et al., 2017) for 50K steps with lr=0.005 and fine-tuned on the GLUE dataset (Wang et al., 2018). We use different learning rates among 1e-5, 3e-5, 6e-5, 1e-4 and choosing the best result after fine-tuning for 3 epochs.
For the Long-Range Arena benchmark, to make sure it reflect the practical speed in Pytorch platform, we re-implement the benchmark in Pytorch. We adopt the same configuration from the Skyformer Chen et al. (2021) and make sure all models have a similar parameter size. We use the same training hyper parameters for all models as well.
2 Results
We report the results in Table 4. It can be found that both TransNormer variants get comparable or better perplexity to the vanilla attention and outperform all existing linear models with a clear margin. For example, compared to previous state-of-the-art linear methods on validation setHua et al. (2022) and test setZhu et al. (2021), TransNormer T2 achieves substantially lower perplexity by 2.31 and 1.58 respectively. It demonstrates the effectiveness of our method in causal models.
We show our bidirectional results on the GLUE benchmark in Table. 5. Our method achieves superior performance to all the competing methods in average. On three tasks, i.e., SST-2, MRPC, CoLA, TransNormer reports comprehensively better results than all competing linear methods, such as 4.62 higher on CoLA. Further, one of our variants i.e., TransNormer T1, even outperforms the vanilla attention with a notable margin. It proves the effectiveness of our method in bidirectional language modeling.
The results before the transformer Long-short (abbr. LS) are taken from the Skyformer Chen et al. (2021). As shown in Table. 6, we achieve either first or second places across all five tasks. In terms of overall results, both TransNormer variants (T1,T2) outperform all other competing methods including vanilla transformer Vaswani et al. (2017), which validates our capability to encode long sequences.
3 Speed comparison
We compare the training and inference speed of the TransNormer with other methods. For a fair and comprehensive comparison, we follow exactly the same configurations of the SkyformerChen et al. (2021) and report step per second under different sequence lengths. Timing is conducted on a Nvidia A6000 GPU with 48G GPU memory. Table. 7 suggests that the vanilla transformer is substantially slow and exhausts GPU memory with sequence longer than 3k. Compared to other efficient transformers, our TransNormer achieves faster speed with comparable GPU memory footprints, while competing efficient methods all report worse results compared to our TransNormer. For instance, compared to FLASH-quad Hua et al. (2022) that achieves previous best linear results on both autoregressive and bidirectional benchmarks, our model performs over 300% faster during training and 150% faster during inference.
4 Ablation study
In this section, we justify our design choice of the TransNormer, including , the selection of the FFN module, and the size of the attention block in DiagAttention. We use the PPL from the Roberta pre-training stage as our evaluation metric.
As aforementioned, we empirically choose the first 6 layers as the early stage of the model and the rest as the later stage. We provide the designing ground for this choice in Table. 8. It can be also observed that either choosing the DiagAttention or NormAttention for the entire model will lead to inferior performance. We also provide the ablation results of swapping the order of the DiagAttention and the NormAttention in Table. 9. Using DiagAttention in the early stage achieves significantly better results than using it on later stage. It further proves our claim that the early stage focuses on neighbouring tokens while the later stage needs long-range attentions.
We ablate the selection of the FFN modules in Table. 10. Compared with the traditional FFN (Vaswani et al., 2017), the GLU (Shazeer, 2020) achieves better results.
From the Table. 11, we observe clear performance improvements with increased block sizes. However, since the complexity of the DiagAttention is , larger block size leads to heavier computational overhead. We choose a block size as 64 as a trade-off between performance and computational cost.
Finally, we study the effect that whether we should use both attentions in one layer. In particular, we compare either to 1) use DiagAttention and NormAttention sequentially in a layer with different orders; or to 2) use them in parallel in each attention layer and then concatenate their embedding output. Table. 12 shows that we should not use these attentions sequentially within a layer and apply them in parallel will double the computation complexities without improving the performance.
Conclusion
In this paper, we identified two key issues that cause the inferior performance of existing linear transformer models: 1) unbounded gradients; 2) attention dilution. For the former issue, we proposed a new NormAttention to stabilize the training gradients. For the latter, we develop DiagAttention to force the model concentrate attention in neighbouring tokens. The resultant model TransNormer marries the strength of the vanilla transformers and the linear transformers, outperforming competing linear transformers on both autoregressive and bidirectional language modeling, text classification tasks and the challenging Long-range arena benchmark.
Limitations
In this paper, we identified two main issues of current linear transformers and provided a comprehensive analysis in natural language processing tasks. However, with the booming development of vision transformers, whether they share the same issues of linear NLP transformers is yet to be discovered. We will validate our method on the linear vision transformers in our future work.
Ethics Statement
The proposed technique is beneficial to develop large-scale environment-friendly language models by reducing computing resource demand. Corpus used to train the model is from public web sources, which may contain biased, explicit or improper content. Further assessment and regulation have to be in-place before deploying the model in practice.
References
Appendix A Mathematical Notations
We use bold uppercase letters for matrices(), bold lowercase letters for vectors(), and lowercase letters for scalars(). We represent all vectors as column vectors and denote the th row of matrix by or . We use to denote the norm and to denote the Frobenius norm of the matrix and the vector.
Appendix B Proof of gradients’ upper bound
In this part, we will proof the bound in (8) and (10), all we need to prove is:
We adopt the theorem that geometric mean is bounded by arithmetic mean, i.e.,
We take to complete the proof. The first bound can be proven by:
For the second bound, we first use the fact that:
Appendix C Proof of Proposition 3.1
and kernel function , letWe assume that the image of contains vectors arbitrary close to , which is a common case in kernel function.:
Let , then , so .
Appendix D Analyze the gradient of each method
In this section, let’s consider a one-layer Transformer, for a multi-layer Transformer, we can prove our conclusion using induction.
We begin this section by introducing some mathematical notations.
In the subsequent discussion, we define gradient as:
where stands for loss function, is a parameter matrix.
The mapping has the following property:
D.2 Gradient analysis
Before we get started, we have the following propositions. The proof can be found in Appendix D.3.
D.2.2 Vanilla/Linear attention
According to (30), we can discuss vanilla and linear attention under one formula:
According to (9), in vanilla attention, we have:
where in vanilla attention and in linear attention.
On the other hand, according to Appendix C, in linear attention, there exist , such that:
Let , then . This means that the gradient in linear attention is unbounded.
D.2.3 NormAttention
We first define the second-moment of ’th row of :
Then is as follows:
Notice that we have the following upper bound:
Let’s summarize the previous results. In vanilla attention, we have:
In linear attention, there exist , such that:
So is bounded in vanilla attention and NormAttention, while it’s unbounded in linear attention. This makes the training of linear transformer unstable.
D.3 Proof of the proposition
The forward pass of the model isXAttention stands for vanilla/norm attention.:
So the gradient passed to XAttention module is bounded, i.e., . ∎
Appendix E Experiment configs
In this section, we will introduce detailed training hyperparameters. We introduce the configurations for autoregressive/bidirectional language model in table F. For LRA benchmark, we use the same configuration as Skyformer, which use 2-layer transformer model with 64 hidden dimensions, 2 attention heads, 85 GLU dimensions, Swish as GLU activation function. For batch size and learning rate , we use 16,1e-4 for Text Classification, 32,1e-4 for ListOps, 16,2e-4 for Document Retrieval, 128,2e-4 for Pathfinder, 256,1e-4 for Image Classification, the same as Skyformer.
Appendix F Pseudocode for visualization.
In this section, we provide pseudo codes for the 4th column of Figure 2 in Python: