cosFormer: Rethinking Softmax in Attention
Zhen Qin, Weixuan Sun, Hui Deng, Dongxu Li, Yunshen Wei, Baohong Lv, Junjie Yan, Lingpeng Kong, Yiran Zhong
Introduction
With years of development, the transformer model (Vaswani et al., 2017) and its variants (Zaheer et al., 2020; Wang et al., 2020; Tay et al., 2020a) have been successfully adapted to three most popular artificial intelligence (AI) fields: i.e., natural language processing (Devlin et al., 2019; Liu et al., 2019), computer vision (Dosovitskiy et al., 2020; Carion et al., 2020; Liu et al., 2021) and audio processing (Schneider et al., 2019; Baevski et al., 2020). Compared with conventional recurrent (Hochreiter & Schmidhuber, 1997) and convolutional architectures (He et al., 2016), transformer-based architectures are generally more scalable to data volumes (Brown et al., 2020) and stronger in capturing global information with less inductive bias, thus excelling on many tasks.
Dot-product attention with softmax normalization is the cornerstone of the transformer to capture long-range dependencies. However, its quadratic space and time complexity with regard to the length of the sequence make its computational overhead prohibitive, especially for long inputs. To address this issue, numerous methods are proposed recently, such as the sparse attention matrix (Zaheer et al., 2020; Beltagy et al., 2020; Tay et al., 2020a; Kitaev et al., 2019; Child et al., 2019),low-rank representations (Wang et al., 2020) or kernel-based methods (Peng et al., 2020; Choromanski et al., 2020; Katharopoulos et al., 2020), among many others. These methods achieve reduced computational complexity with comparable performances when compared with the vanilla attention architecture on several selected tasks or corpus.
However, the improved efficiency is usually achieved via introducing additional yet often impractical assumptions on the attention matrix (Wang et al., 2020) or with valid approximation of softmax operation only within constrained theoretical bounds (Choromanski et al., 2020; Peng et al., 2020) Therefore, when their assumptions are unsatisfied or when approximation errors get accumulated, these methods may not always be advantageous over the vanilla architecture (Narang et al., 2021). Consequently, performance deficiencies in a broad application spectrum are often observed in these transformer variants, especially those with linear complexity. For example, the Performer (Choromanski et al., 2020), RFA (Peng et al., 2020) and Reformer (Kitaev et al., 2019) show less satisfactory performance on the GLUE benchmark (Wang et al., 2018) when compared with the vanilla architecture as suggested in our preliminary experiments (Tab. 2). Furthermore, many of these aforementioned methods are not applicable to casual attentions, which are critical for auto-regressive training. For example, techniques proposed in Linformer (Wang et al., 2020) and BigBird (Zaheer et al., 2020) are specific to cross attentions.
Since the softmax operator appears to be the main hurdle while efficient yet accurate approximation to softmax is difficult to achieve, one question naturally arises: “Can we replace the softmax operator with a linear function instead, while maintaining its key properties?”. By digging into the softmax attention, we find two key properties that affect its empirical performance: (i) elements in the attention matrix are non-negative (Tsai et al., 2019; Katharopoulos et al., 2020); (ii) the non-linear re-weighting scheme acts as a stabilizer for the attention weights (Titsias, 2016; Gao & Pavel, 2017; Jang et al., 2016). These findings reveal some new insights of the current approaches. For example, the linear transformer (Katharopoulos et al., 2020) achieves property (i) using an exponential linear unit (Clevert et al., 2016) activation function. However, due to lack of the re-weighting scheme, it underperforms other efficient transformer variants on the Long-Range Arena benchmark as shown in Figure 1 as well as the language modeling task (Table 2) based on our controlled experiments.
In this paper, we propose a new variant of linear transformer called cosFormer that satisfies both of the above properties. Specifically, we enforce the non-negative property by passing the features to a ReLU (Agarap, 2018) activation function before computing the similarity scores. In this way, we encourage the model to avoid aggregating negatively-correlated contextual information. Further, we adopt a cos re-weighting scheme to stabilize the attention weights. This helps the model to amplify local correlations, which usually contain more relevant information for natural language tasks. Thanks to the Ptolemy’s theorem, our attention can be exactly decomposed into a linear form. We perform extensive experiments on both autoregressive language models and bidirectional models on five public benchmarks, including WikiText-103 (Merity et al., 2017), GLUE (Wang et al., 2018), IMDB (Maas et al., 2011), AMAZON (Ni et al., 2019) and Long-Range Arena benchmark (Tay et al., 2020b). Our model shows much better inference speed and smaller memory footprint, while achieving on par performance with the vanilla transformer. It is noteworthy that our method ranks 1 on the Long-Range Arena benchmark, showing favorable performance than other competitors, which well demonstrates its strong capacity in modeling long sequence inputs.
Our Method
In this section, we provide technique details of our linear transformer called cosFormer . The key insight of the cosFormer is to replace the non-decomposable non-linear softmax operation by a linear operation with decomposable non-linear re-weighting mechanism. Our model is applicable to both casual and cross attentions with a linear time and space complexity with regard to the input sequence length, thus exhibiting strong capacity in modeling long-range dependencies.
where is a feedforward network that contains a residual connection; is the self-attention function that computes the attention matrix , which has quadratic space and time complexity with respect to , thus becoming the computation bottleneck of on long inputs.
where measures the similarity between queries. If , the Eq. 2 becomes the dot-product attention with softmax normalization. In this case, the space and time complexity to compute one row of the output is . Therefore, the total space and time complexity for computing grows quadratically with respect to the input length.
2 Linearization of Self-attention
According to Eq. 2, we can select any similarity functions to compute the attention matrix. In order to maintain a linear computation budget, one solution is to adopt a decomposable similarity function such that:
where is a kernel function that maps the queries and keys to their hidden representations. Then one can rewrite Eq. 2 in the form of kernel functions as:
After that, attention operation in linear complexity is achieved via the matrix product property:
As aforementioned, the key to the linear attentions is to find a decomposable similarity function that generalizes well to different tasks. Most existing linear transformers are trying to find an unbiased estimation of the softmax attention. For example, RFA (Peng et al., 2020) approximates the softmax operation with random feature maps using theorem of random fourier features (Rahimi & Recht, 2008) and the Performer (Choromanski et al., 2020) utilizes positive random features to approximate it. However, we empirically find that these methods are sensitive to the selection of sampling rate and becomes unstable if the sampling rate gets too high. Also, to accommodate recency bias, gating mechanisms are employed to better exploit more recent context.
3 Analysis of Softmax Attention
In this work, we empirically identify two key properties of the softmax operation that may play important roles for its performance: 1) it ensures all values in the attention matrix to be non-negative; 2) it provides a non-linear re-weighting mechanism to concentrates the distribution of attention connections and stabilizes the training(Titsias, 2016; Gao & Pavel, 2017; Jang et al., 2016).
4 cosFormer
Based on the observations above, we propose our model cosFormer , which discards entirely the softmax normalization while still features the non-negativity and re-weighting mechanism. Our cosFormer consists two main components: a linear projection kernel and a cos-Based Re-weighting mechanism. Below we describe details of each components:
Recall the general form of the attention in Eq. 2, let us define a linear similarity as:
Based on Eq. 4, we rearrange the order of dot-product and obtain the formulation of the proposed attention in linear complexity as:
cos-Based Re-weighting Mechanism
The non-linear re-weighting mechanism introduced by the softmax attention can concentrate the distribution of the attention weights and therefore stabilize the training process (Titsias, 2016; Gao & Pavel, 2017; Jang et al., 2016). We also empirically find that it can punish far-away connections and enforce locality in some cases. In fact, such locality bias, i.e., a large portion of contextual dependencies are from neighboring tokens, is commonly observed on downstream NLP tasks (Clark et al., 2019; Kovaleva et al., 2019), as shown in Figure 3 (1).
Based on the assumption above, what we need to fulfill the second property of softmax may be a decomposable re-weighting mechanism that can introduce recency bias to the attention matrix. Here, we propose a cos-based re-weighting mechanism as it perfectly fit our purpose: 1). the Ptolemy’s theorem ensures the weights can be decomposed into two summations; 2). as shown in Figure 3 (4), the will put more weights on the neighbouring tokens and therefore enforces locality. Also, by comparing the attention matrices in Figure 3 (2) and (3), the cosFormer enforces more locality than the one without the re-weighting mechanism.
Specifically, by combining with Eq 6, the model with cosine re-weighting is defined as:
By leveraging the Ptolemy’s theorem, we decompose this formulation as:
where is the output at the position of the sequence from the attention module. Detailed derivation are included in the Appendix. Without losing the generality, our method achieves a linear complexity as:
Relation to positional encoding.
cosFormer can be seen as a new way of introducing the relative positional bias to the efficient transformer. Compared with the Rotary Position Embedding (Su et al., 2021), they use a more complex position embedding strategy and did not enforce the non-negativity to the similarity scores as ours. Also, since they only change the position embedding on the numerator while keeping the denominator unchanged, the summation of their attention scores is not equal to 1. For Stochastic Positional Encoding (Liutkus et al., 2021), they use a sampling strategy to approximate the softmax, and introduce relative positional encoding to linear transformers.
Experiments
In this section, we experimentally validate the effectiveness of the proposed method in multiple settings. The purposes of the experiments are three-fold. First, we validate the capacity of cosFormer in language modeling through autoregressive (Sec. 3.1) and bidirectional (Sec. 3.2) setups using WikiText-103 (Merity et al., 2017). In this way, we validate the effectiveness of the proposed linear attention module in both causal and non-causal cases. Second, we investigate the generalization ability of cosFormer on downstream tasks by comparisons with other existing transformer variants. This is achieved by performing comparative finetuning experiments on five datasets, including GLUE (QQP, SST-2, MNLI) (Wang et al., 2018), IMDB (Maas et al., 2011) and AMAZON (Ni et al., 2019) (Sec. 3.3). We further compare cosFormer with other transformer variants on the long-range-arena benchmark (Tay et al., 2020b) to understand its ability in modeling long-range dependencies (Sec. 3.4) and show comparative analysis into model efficiency (Sec. 3.5). Third, we conduct ablation studies to understand each component in cosFormer (Sec. 3.6).
In autoregressive or left-to-right language modeling, we estimate the probability distribution of a token given its previous tokens. We use (Baevski & Auli, 2018) as our baseline model. Specifically, we adopt their large model which has 16 cascaded layers with a projected dimensions of 1024, and replace the self-attention module with our proposed linear attention module. We train our model on 8 Nvidia Tesla A100 GPUs with a sequence length of 512 for 150K updates on the WikiText-103 (Merity et al., 2017) and report perplexity on the validation and test splits in Table 2.
We observe that although the baseline model is a powerful standard transformer which requires quadratic computation complexity, cosFormer outperforms it with a clear margin in linear computation complexity. Besides, we achieve comparable perplexity to other methods on the validation set, and significantly outperform all competing methods on the test set by a clear gap, which further demonstrates the effectiveness of cosFormer .
2 Bidirectional Language Model
For bidirectional language modeling, we adopt RoBERTa (Liu et al., 2019) as the baseline model. Similarly, we replace the self-attention module in the RoBERTa by the proposed linear attention module, and keep other structures unchanged. We train this bidirectional task on 2 Nvidia Tesla A100 GPUs for 50K iterations with a input sequence length 512. As shown in Figure 4, cosFormer converges faster than vanilla transformer on both training and validation sets with a comparable or smaller loss values, despite it only consumes linear space and time computation complexity. In addition, the cosFormer variant with re-weighting mechanism has both notably better converge speed and final results over the counterpart without re-weighting, which further validates the effectiveness of our -based distance matrix and also demonstrates the effectiveness of recency bias on natural language data.
3 Downstream fine-tuning tasks
In this section, we fine-tune the pre-trained model on downstream tasks to demonstrate the generalization ability of cosFormer on downstream tasks. We use the pre-trained bidirectional model and fine-tune it on three downstream text classification tasks: GLUE (QQP, SST-2, MNLi) (Wang et al., 2018), IMDB (Maas et al., 2011) and AMAZON (Ni et al., 2019). For fair comparison, we first pre-train all the competing efficient transformer variants for the same 50K iterations on WikiText-103 (Merity et al., 2017) under the same setting, then we follow the same fine-tuning protocol as RoBERTa (Liu et al., 2019) to fine-tune these methods on the downstream tasks. From Table 3, we can see that cosFormer outperforms baseline (Liu et al., 2019) on three out of five datasets, and achieves either best or secondary place on all five downstream datasets compared to competing efficient transformers. It is worth noting that despite Longformer (Beltagy et al., 2020) achieves better results on MNLI than cosFormer , it requires a computation complexity of , where is window size. As shown in Figure 1, Longformer is slower and requires more memory overhead than cosFormer . Other competing methods(Peng et al., 2020; Choromanski et al., 2020; Kitaev et al., 2019) are all based on kernel functions and have substantial performance gaps compared with our model. This validates the effectiveness of the proposed cosFormer model compared with other efficient transformer variants.
4 Results on Long-range-arena Benchmark
To further evaluate the generalization ability of the proposed method, we train our model from scratch on Long-range-arena benchmark 2020b. Long-range-arena (Tay et al., 2020b) is a benchmark specifically designed for efficient transformers with long input sequences, thus serving as a suitable testbed to assess the quality of efficient transformer variants comparatively. To ensure fair comparison, we first implement our method on Jax (Bradbury et al., 2018), then carefully follow their preprocessing, data split, model structure and training protocol. We evaluate our method on a variety of tasks including Long sequence ListOps (Nangia & Bowman, 2018), Byte-level text classification (Maas et al., 2011), document retrieval using the ACL Anthology Network (Radev et al., 2013), image classification on sequence of pixels on CIFAR-10 (Krizhevsky & Hinton, 2009), and Pathfinder (Linsley et al., 2018). As shown in Table 4, cosFormer overall achieves competitive results across all the tasks while achieving best performance on ListOps and Document Retrieval. For the Pathfinder task, since the distance between the two points can be very far from each other, our introduced locality bias would have negative impact to this task and show a bit lags to other SOTA methods, despite that the performance gap between our method and the vanilla transformer is small It is worth mentioning that cosFormer achieves the best overall scores on Long-range-arena benchmark, being one of the only two models that surpass vanilla transformer architecture.
5 Efficiency Comparison
In this section, we compare the efficiency of cosFormer with other models, with a focus on long sequences as inputs. With the proposed linear attention module, we expect that cosFormer scales comparably with other linear variants while significantly surpassing the vanilla transformer architecture. For a fair and comprehensive comparison, we implement our method and competing methods on Jax (Bradbury et al., 2018). We use the byte-level text classification benchmark and report runtime speed during both training and inference under different sequence lengths (1k-4k). We conduct experiments on one Nvidia A6000 GPU and also report the corresponding inference-time memory foot prints as shown in Figure 1. As shown in Table 5 and Figure 1, most pattern based methods (Beltagy et al., 2020; Zaheer et al., 2020; Tay et al., 2020a; 2021) and vanilla transformer (Vaswani et al., 2017) are much slower and require greater memory than cosFormer prevents them from extending to longer sequence. Further, the kernel based methods like (Narang et al., 2021; Choromanski et al., 2020; Tay et al., 2020a) have comparable speed and memory overheads, but their performances are less satisfactory compared to cosFormer across above metrics. In summary, our model cosFormer achieves overall better efficiency than other linear variants while maintain superior modeling and generalization ability.
6 Ablation: cos𝑐𝑜𝑠cos-based Re-weighting
By introducing -based re-weighting, we provide a non-linear mechanism to concentrate the distribution of attention connections and stabilizes the training. In this way, we encourage the model to better take into account the locality inductive biases commonly observed on many natural language tasks. In particular, we investigate the effect of the -based re-weighting in two aspects. First, as shown in Figure 4, by adding -based re-weighting, we obtain both notably better converge speed and final results in autoregressive language modeling. Further, in Table 6, we present a comparison between cosFormer models with and without re-weighting mechanism. We use two composite metrics which comprehensively include 10 different datasets from bidirectional downstream fine-tuning tasks and long-range-arena (Tay et al., 2020b). cosFormer achieves overall better results over the counterpart without re-weighting, improving the average scores on bidirectional finetuning and long-range-arena by a clear margin. This verifies that the proposed re-weighting effectively incorporates the locality inductive biases for natural language tasks.
Related work
This section will introduce the existing works on improving the efficiency of Transformers, they can be broadly divided into two categories, Pattern based methods and Kernel based methods.
Pattern based methods sparsify the attention matrix with handcrafted or learnable patterns. As an early approach, Lee et al. (2019) leverages the inducing points from the sparse Gaussian process to reduce the quadratic complexities of a transformer. Child et al. (2019) reduces the complexity by applying combination of strided pattern and local pattern to the vanilla attention matrix. Longformer (Beltagy et al., 2020) designs fixed diagonal sliding windows combined with global window, and the sliding window pattern can also be extended with dilation to enlarge the receptive field. Zaheer et al. (2020) presents a more powerful and expressive sparse attention mechanism, which combines multiple types of attention patterns and gives a thorough study of sparse attention mechanism. Instead of fixed patterns, Kitaev et al. (2019) and Daras et al. (2020) group the attention computation process into buckets by local sensitive hashing, while Roy et al. (2020) uses mini-batch spherical -means. Nevertheless, Pattern based methods can only cope with sequences up to a certain length, and the computational complexity still grows rapidly when the input sequence becomes longer.
Kernel based method
When faced with longer input sequences, it is more efficient to directly reduce the complexity of the theoretical calculation method. Kernel based methods speed up self-attention by reducing the computation complexity of self-attention from quadratic to linear. Vyas et al. (2020) approximate the full attention with a fixed number of cluster attention groups by assuming neighbouring queries in Euclidean space should have similar attention distributions. Peng et al. (2020) chooses to use the production of Gaussian kernel functions to approximate Softmax, changing the order of scale dot product calculation, thus reducing the theoretical time to linear complexity and Choromanski et al. (2020) uses Haar measurement based kernel instead. Wang et al. (2020) imports the low-rank prior for attention matrix and approximate softmax with SVD decomposition manner. Xiong et al. (2021) utilizes the Nyström method with segment-means to generate a low-rank approximation of the Softmax matrix. Katharopoulos et al. (2020) formalizes the transformer layer as a recurrent neural network. In this paper, we demonstrate that the approximation to Softmax is unneccessary for Linearization of self-attention module. We instead propose a new method to replace Softmax with a linear operation with a re-weighting mechanism, which reduces both time complexity and space complexity to while maintaining the accuracy.
Conclusion
References
Appendix A Appendix
Following Equation 11, we give a detailed deviation of how to obtain output at position position:
A.2 Pseudo Code of cosFormer
Algorithm 1 describe the way to compute cosFormer attention
A.3 Algorithm to visualize attention matrix
Algorithm 2 describe the way to visualize attention matrix as Figure 3
A.4 Introduction of Dataset
We train both models on autoregressive language modeling and bidirectional modeling by Wikitext-103 dataset, it is split by tokens and its statistics as Table 7.Then we fine-tune the pre-trained bidirectional modeling on several text classification tasks.
QQP dataset contain thousands of sentence pair from community question-answering website Quora.Network need to determine pairs of question are semantically equivalent. SST-2 and IMDB are collections of movie reviews. The task is to determine whether a review is positive or not. AMAZON dataset contains millions of product reviews from Amazon.The requirement of this task is to infer the scoring of the product from the review text.MNLI is a crow-source collections of sentence pairs. The network must distinguish which of the three categories entailment, contradiction and neutral the given sentences belong to.
The long-range-aren benchmark contains 5 different datasets.ListOps contains some designed clever mathematical problem to clarify the parsing ability of neural models. IMDB is also used in this benchmark to examine the text classification ability of neural models. CIFAR-10 is a image collection of various of object, this task require models capture 2D spatial relations between flatten pixels.In pathfinder task, models need to determine the connection of two points in the picture, so as to examine the model’s ability to acquire 2D spatial relationships.AAN dataset is used to evaluate the ability for models to encode and store compressed representations for retrieving.
A.5 Qualitative Results of LRA
We provide our qualitative results of the ListOps and Document Retrieval tasks on Long-Range-Arena benchmark (Tay et al., 2020b) with a comparison to the vanilla transformer.
ListOps is a ten-way classification task which aims to prediction the results of a sequence with a hierarchical structure and operators MAX, MEAN, MEDIAN and SUM MOD that are enclosed by delimiters (brackets). The network needs to access all tokens and model the logical structure of the inputs in order to make a prediction.
Document Retrieval task is to decide whether the two input long documents are similar or not with a binary label. This task evaluates a model’s ability to encode and store compressed representations that are useful for matching and retrieval. Since the samples in LRA are too long, We substantially shorten some selected samples and display them as below: