Linear Complexity Randomized Self-attention Mechanism

Lin Zheng, Chong Wang, Lingpeng Kong

Introduction

Transformers (Vaswani et al., 2017) are powerful neural networks for sequence modeling. They have been successfully applied in various domains, such as natural language processing (Vaswani et al., 2017; Dehghani et al., 2019; Devlin et al., 2019; Raffel et al., 2020), computer vision (Carion et al., 2020; Dosovitskiy et al., 2021; Liu et al., 2021), bioinformatics (Rives et al., 2021; Jumper et al., 2021) and reinforcement learning (Chen et al., 2021c). The core building block of transformer models is the self-attention mechanism, which captures complex interactions among sequence elements (Vaswani et al., 2017).

However, the computational complexity of attention mechanism is quadratic in the number of tokens, making it prohibitive to process long sequences. In the past two years, there has been a community effort towards developing efficient attention architectures with improved computation complexity and memory usage (Tay et al., 2020b). Among them, one prominent is to view the attention mechanism through kernelization (Katharopoulos et al., 2020; Choromanski et al., 2021; Peng et al., 2021b, inter alia). In this work, we focus on random feature attentions (RFAs) (Peng et al., 2021b; Choromanski et al., 2021), which approximate softmax attention by linearizing the exponential kernel into a dot product of random feature maps. Despite achieving linear time and space complexity, this approximation is biased to the softmax attention as a whole.There are several variants of random feature maps that yield an unbiased estimate of the exponential kernel (Peng et al., 2021b; Choromanski et al., 2021). Nevertheless, RFAs still run a biased approximation to the whole softmax attention, since the softmax attention involves a ratio of these exponential kernels. Although the estimator is still consistent, the bias in question is elusive and might impair the approximation fidelity of random features.

In this work, we revisit RFA and show that it can be reinterpreted as a self-normalized importance sampler to softmax attention. This insight reveals that the source of the approximation bias in RFAs comes from the self-normalization in estimation (Owen, 2013). We further show softmax attention can be written as an expectation of linearized attention over an input-dependent mixture distribution. These findings suggest that we can in principle construct an unbiased estimator for the softmax attention as a whole, as opposed to merely exponential kernels in previous work. We call such unbiased estimation randomized attention or RA. To the best of our knowledge, this is the first unbiased approximation of the whole softmax attention via kernel linearization.

RA constructs positive random features via distributions exclusive to each query. Since RFAs only employ an input-agnostic standard Gaussian as the importance sampling proposal, RA enables a finer-grained treatment for query-specific information and greatly improves the approximation fidelity; however, it is as expensive as softmax attention computationally with quadratic complexity, because the key-value statistics are different for each query, unlike the ones in RFAs.

Based on the analysis, one question naturally arises: “Can we combine the expressiveness in RA and the efficiency in RFA to get the best of both worlds?” To achieve that, we generalize the importance sampling formulation of RFA by adopting multiple proposals, each of which depends on different subsets of queries. We further apply multiple importance sampling (Veach & Guibas, 1995) and put together these proposals to approximate softmax attention adaptively for different queries, retaining the query-specific property of RA. Meanwhile, since these proposals are shared among all queries, we inherit the efficient computation reuse in RFA and achieve linear complexity. We refer to this efficient attention mechanism as LineAr Randomized Attention (LARA). Extensive experiments and analyses demonstrate that RA, as well as its linear variant LARA, significantly reduce the approximation error of RFAs. They improve RFAs by a substantial margin across various tasks, including image classification, video action recognition, machine translation, and so on, while retaining computational efficiency.

Background

Intuitively, the softmax attention first computes the normalized similarity between the query and each key, which is then used to weight value vectors. In the case of self-attention in Transformers (Vaswani et al., 2017), we have N=MN=M; as a result, such mechanism suffers from quadratic time and memory complexity due to the explicit computation of the similarity scores between all pairs of queries and keys.

2 Random Feature Attention

To reduce the computational complexity of softmax attention, recent work (Choromanski et al., 2021; Peng et al., 2021b) proposes to linearize exponential kernels via random feature methods (Rahimi & Recht, 2008). According to Bochner’s theorem (Bochner, 2020), they re-write the exponential kernel exp⁡(x⊤y)\exp\left(\boldsymbol{\mathbf{x}}^{\top}\boldsymbol{\mathbf{y}}\right) as the following expectation,

Thanks to the linearized formulation, one can first pre-compute the corresponding key-value statistics ∑m=1Mξ(km,ωs)vm⊤\sum_{m=1}^{M}\xi(\boldsymbol{\mathbf{k}}_{m},\omega_{s})\boldsymbol{\mathbf{v}}_{m}^{\top} and ∑m=1Mξ(km,ωs)\sum_{m=1}^{M}\xi(\boldsymbol{\mathbf{k}}_{m},\omega_{s}) once, and then reuse them for each query. Consequently, it achieves linear complexity in both time and memory with respect to the sequence length.

3 Self-normalized Importance Sampling

The name self-normalized comes from the fact that the importance weights p(ω)/q(ω)p(\omega)/q(\omega) are normalized. Albeit at the cost of introducing a bias, this method cancels out the normalizing constant ZZ at both nominator and denominator. SNIS often works well in practice.

Randomized Attention

In this section, we present an alternative view of RFA, revealing new insights of how RFA approximates the softmax attention. In particular, we show that RFA can be recast as a self-normalized importance sampler and its target expectation is exactly softmax attention (§3.1). This reformulation allows us to construct an unbiased estimator for softmax attention. We refer this unbiased estimation as randomized attention (§3.2).

Solving this relation gives concise formulations for both f(ω)f(\omega) and p(ω)p(\omega) (see Appendix A for the proof):

where πm=exp⁡(qn⊤km)∑m′=1Mexp⁡(qn⊤km′)\pi_{m}=\frac{\exp\left(\boldsymbol{\mathbf{q}}_{n}^{\top}\boldsymbol{\mathbf{k}}_{m}\right)}{\sum_{m^{\prime}=1}^{M}\exp\left(\boldsymbol{\mathbf{q}}_{n}^{\top}\boldsymbol{\mathbf{k}}_{m^{\prime}}\right)} is the component weight. Besides, f(ω)f(\omega) is an attention-like aggregation function over value vectors, which computes the linearized similarity between queries and keys via randomized mappings,

This re-formulation offers alternative viewpoints to understand the approximation quality of RFA. It is straightforward to see that RFA is a biased (but consistent) estimator due to the self-normalization (Owen, 2013). In addition, RFA may exhibit large bias and variance since it only uses a standard Gaussian proposal, which is far away from the underlying input-dependent mixture pn(ω)p_{n}(\omega). These may explain its inferior performance and slow convergence observed in previous studies (Patrick et al., 2021; Tay et al., 2021b).

2 Randomized Attention

The analysis above further implies that the softmax attention itself can be formulated as an expectation.

Let pn(ω)p_{n}(\omega) and fn(ω)f_{n}(\omega) be defined by Equation 5 and Equation 6 respectively. Then for softmax attention we have

The detailed proof is in Appendix B. As a result, RFA can be viewed as using importance sampling to estimate softmax attention. Alternatively, one can directly sample from pn(ω)p_{n}(\omega) to construct an unbiased estimate of the softmax attention,

with ω1,…,ωS∼pn(ω)\omega_{1},\dots,\omega_{S}\sim p_{n}(\omega). To the best of our knowledge, this is the first kernel linearization estimator that approximates the whole softmax attention, instead of just exponential kernels, in an unbiased manner. We refer to this estimator as randomized attention (RA), since it computes attention-like aggregations but via randomized mappings.

Intuitively, RA constructs the randomized mapping by sampling from the contextual distribution pn(ω)p_{n}(\omega), which promotes ω\omega in the vicinity of the resultant of current queries and keys. Aware of locations of query-key pairs, ω\omega is likely to describe their similarity better than input-agnostic ones as in RFA. In addition, each query position nn in RA induces an exclusive distribution pnp_{n}, which makes the randomized mapping adaptive to each query. This allows the model to process query information at a finer-grained level and thus achieves higher approximation fidelity (see §5 for empirical validation). Nevertheless, the use of query-specific modeling requires to draw a different set of samples for different queries. As a result, the mapped key statistics ξ(km,ω)\xi(\boldsymbol{\mathbf{k}}_{m},\omega) will be different for different queries, which prevents reusing the computation and results in O(MN)\mathcal{O}(MN) complexity, rendering it less applicable in approximating softmax attention in practice.

This is in sharp contrast to RFA. RFA uses the same proposal N(ω;0,I)\mathcal{N}(\omega;0,\mathbf{I}) for all queries, and thus the modeling power is greatly reduced since the standard Gaussian would capture neither contextual information nor the inherent variations among queries. The advantage of the shared proposal is that it enables efficient computation reuse of key-value statistics (Equation 2), as the same randomized mapping is reused across queries. This property accounts for RFA’s linear complexity.

Linear Complexity Randomized Attention

In this section, we propose an improved estimator of softmax attention to combine both the expressiveness of RA and the efficiency of RFA. Motivated by the difference between RA and RFA, we generalize the importance sampling formulation of RFA by adopting multiple proposals. This strategy not only captures query information at a finer-grained level, but also allows the model to estimate softmax attention in a query-specific manner (§4.1). We further show that computation reuse in RFA can be achieved, which leads to linear complexity computation with the help of self-normalized importance sampling (§4.2).

This strategy not only enables a finer-grained treatment for query information, but also allows the model to estimate softmax attention in a query-specific way, which is the key advantage of RA. To be specific, since there are several proposals available for each query, and these proposals may provide complementary information to each other, we could combine them by invoking multiple importance sampling (MIS; Veach & Guibas, 1995). For each query, the MIS estimate takes the following form,Here we assume only one sample is drawn from each proposal distribution. A more general treatment would allow arbitrary numbers of samples to be drawn from each proposal.

where ωc∼qc(ω)\omega_{c}\sim q_{c}(\omega) for c=1,…,Cc=1,\dots,C and {αnc(⋅)}c=1C\{\alpha_{nc}(\cdot)\}_{c=1}^{C} are weighting functions. The MIS estimator is unbiased (Veach & Guibas, 1995) if ∑c=1Cαnc(ω)=1\sum_{c=1}^{C}\alpha_{nc}(\omega)=1 for any ω\omega (see the proof in Appendix F).Strictly speaking, for the MIS estimator to be unbiased, we additionally need the weighting functions to be zero for any ω\omega such that pn(ω)=0p_{n}(\omega)=0, although this holds trivially in our setting. Intuitively, MIS first computes individual importance sampling estimates with each proposal, which are averaged together according to the query-specific weighting functions.

Ideally, the nn-th set of weighting functions {αnc(⋅)}c=1C\{\alpha_{nc}(\cdot)\}_{c=1}^{C} should specialize in processing the nn-th query. To accomplish this goal, we expect weighting functions to be optimal (i.e., minimize the estimation variance) for the corresponding query. Optimal weighting functions takes the following form (detailed derivation can be found in Appendix D),

Here rnc(⋅)r_{nc}(\cdot) is roughly proportional to the closeness between qc(⋅)q_{c}(\cdot) and the query-specific optimal proposal. Intuitively, the optimal weighting function consists of two terms. The first term is query-agnostic and the second term is a query-specific correction. The correction term is defined by the difference between rnc(⋅)r_{nc}(\cdot) and its average weighted by qc(⋅)q_{c}(\cdot); consequently, if rnc(⋅)r_{nc}(\cdot) is large, the correction term will be positive, driving the weight of the cc-th proposal to be higher and vice versa.

In most cases, it is intractable to apply optimal weighting functions, since the closed form of rnc(⋅)r_{nc}(\cdot) is not available. We therefore approximate the optimal weighting functions by the following form,

where rnc′r^{\prime}_{nc} measures the degree of the proposal qcq_{c} favoring the nn-th query. For tractability, we implement rnc′r^{\prime}_{nc} as the normalized similarity between the nn-th query and the representation of the cc-th query subset. We also decouple the computation between proposal densities qc(ω)q_{c}(\omega) and rnc′r^{\prime}_{nc}, so that contributions from query-agnostic and query-specific terms can be independent of each other (see § G.3.2 for more details and ablations). Note that Equation 9 still ensures unbiasedness (or consistency) of MIS estimation due to ∑c=1Cαnc(ω)=1\sum_{c=1}^{C}\alpha_{nc}(\omega)=1.

2 Achieving Linear Time and Space Complexity

According to our MIS estimator (Equation 8), the key-value statistics under each proposal can be pre-computed once and then reused for all queries. This implies the computation reuse in RFA is achievable and so as the linear complexity.

The only problem left now is that the MIS estimator still requires explicitly evaluating the density pn(ω)p_{n}(\omega) for each query (Equation 5), which exhibits quadratic complexity. This is because pn(ω)p_{n}(\omega) is a Gaussian mixture with MM components, incurring O(NM)\mathcal{O}(NM) computations in total. We show that a self-normalized version of MIS allows us to further reduce the complexity to be linear. According to Proposition 3.1 (and Equation 15 in Appendix A), the mixture density pn(ω)p_{n}(\omega) can be equivalently expressed as

Our key observation is that now the numerator contains a linearized dot product of randomized mappings, which can be pre-computed and reused for all queries, while the denominator is similar to the normalizing constant in regular softmax attention and can only be computed in quadratic time. Fortunately, the denominator can be canceled out if we adopt the self-normalized estimator (see §2.3),

The resulting estimator is consistent and runs with linear complexity, similar to RFA. We name it linear randomized attention (LARA). See Algorithm 3 in § G.3 for an algorithmic sketch of LARA.

3 Discussion: RFA, RA, and LARA

LARA defines a flexible framework to bridge RFA and RA. To delineate the connection between RFA and LARA, we find LARA can be further rewritten as (see Appendix E for the derivation)

where ωc∼qc(ω)\omega_{c}\sim q_{c}(\omega) for c=1,…,Cc=1,\dots,C and αnc′(ωc)≔αnc(ωc)N(ωc;0,I)/qc(ωc)\alpha^{\prime}_{nc}(\omega_{c})\coloneqq\alpha_{nc}(\omega_{c})\mathcal{N}(\omega_{c};0,\mathbf{I})/q_{c}(\omega_{c}). Comparing to the formulation of RFA (Equation 2), we see that RFA is a special case of LARA if we set all proposals to N(ω;0,I)\mathcal{N}(\omega;0,\mathbf{I}) and all αnc(⋅)\alpha_{nc}(\cdot) to constant functions. On the other hand, LARA is equivalent to RA if we remove the use of self-normalization, set αnc(ω)=δnc\alpha_{nc}(\omega)=\delta_{nc} That is, weighting functions now become the Kronecker delta function, where αnc(ω)=1\alpha_{nc}(\omega)=1 if n=cn=c and otherwise. and maintain NN proposals, each of which takes the same form of pn(ω)p_{n}(\omega) (Equation 5). With general proposals and weighting functions, LARA approximates softmax attention in a query-specific manner as in RA while achieving linear complexity as in RFA, effectively combining the advantages of both estimators.

Experiments

In this section, we conduct extensive experiments across various domains to verify the effectiveness of linear randomized attention. Firstly, we start with an experiment to assess the approximation error of different random feature based methods (§5.1). We then perform a number of experiments on various data modalities, including image classification (§5.2), video action recognition (§5.3), machine translation (§5.4), and long sequence modeling on Long Range Arena benchmark (§ I.2). Additional details as well as ablation studies can be found in Appendices H and I. The implementation details of RA, Performer (RFA) and LARA are provided in Appendix G.

We conduct a preliminary experiment to assess the approximation fidelity of different random feature methods (details are deferred to § H.1). In particular, we consider vision transformers (ViT; Dosovitskiy et al., 2021; Touvron et al., 2021), keep Q,K\boldsymbol{\mathbf{Q}},\boldsymbol{\mathbf{K}} and V\boldsymbol{\mathbf{V}} the same across attention variants, and compute the Mean Squared Error (MSE) between the outputs of true softmax attention and its approximations. We use the ImageNet1k validation set (see more details in §5.2) as the input data and report MSE results averaged over all images. Figure 1 shows the results with respect to the number of random samples under different sequence lengths. We observe that RFA (Performer) soon plateaus at large approximation error and does not improve even with more samples, possibly due to low sample efficiency. On the other hand, LARA exhibits much lower MSE than Performer and the approximation error continually decreases as the number of samples increases. As for RA, it achieves the lowest MSE among these three methods. This clearly indicates that increasing the model’s resolution over query positions as in LARA and RA is more effective in improving approximation quality, compared to simply increasing the sample size from the same distribution (as in Performer).

2 Image Classification

For image classification, we conduct our experiment on the ImageNet1k benchmark (Deng et al., 2009), which consists of approximately 1,280K/50K images over 1,000 classes for training/validation splits respectively. We apply our attention mechanism to different vision transformer (ViT) architectures (Dosovitskiy et al., 2021), including DeiT (Touvron et al., 2021) and pyramid vision transformers v2 (PVTv2; Wang et al., 2021a, b). The former architecture adopts standard transformer layers with regular softmax attention and receives sequence with length 196 by default; while the latter processes much longer image sequences, which is therefore more suitable to evaluate the scalability of various efficient attention. More model and training details can be found in § H.2.

The comparison among different random feature based methods on DeiT model is demonstrated in Table 1. Consistent with previous studies (Zheng et al., 2021), Performer (RFA) incurs a significant performance drop due to its limited modeling capacity. Its unbiased counterpart, RA, performs much better than Performer and even slightly outperforms exact softmax attention under larger model sizes. This empirically validates the expressiveness of unbiasedness in approximating softmax attention. LARA achieves a good trade-off between Performer and RA. It enjoys linear complexity as Performer but performs substantially better. On the other hand, we note that a linear complexity variant enables the transformer model to scale to much longer sequences, which is often prohibitive for traditional softmax attention but delivers better predictive performance (El-Nouby et al., 2021). We thus train Performer and LARA with 8×88\times 8 image patches (resulting in sequence length 784) with all other settings unchanged. As shown in Table 1, increasing the sequence length (suffixed with “-8”) consistently boosts model performance. However, LARA benefits from longer sequences much more significantly than Performer and outperforms softmax attention by a large margin. This indicates the potential modeling power of our framework for long sequences. Also see § I.1 for additional experiments and ablations.

We then apply our method to the strong baseline PVTv2 and compare it against recent state-of-the-art model architectures. As presented by Table 2, we observe although replacing spatial reduction attention (SRA; details in §H.2) with Performer leads to inferior performance, LARA brings a consistent performance gain over vanilla SRA with much fewer model parameters. In addition, PVTv2 with LARA even performs highly competitive with state-of-the-art architectures across various model sizes, without introducing other inductive biases (such as locality). This implies the superior modeling capacity of LARA compared to SRA and Performer.

3 Video Action Recognition

In this section, we test our method on video action recognition with video transformers. We consider two standard datasets: (1) Kinetics-400 (K400; Kay et al., 2017), which contains 238,574 videos for training and 19,877 for evaluation at the time of writing and (2) Something-something-v2 (SSv2; Goyal et al., 2017), consisting of around 168K/25K videos of 174 classes for training/validation splits respectively. We base our model on the Motionformer architecture (Patrick et al., 2021) and follow their training and evaluation protocol; more details can be found in § H.3.

Table 3 reports the top-1 classification accuracy for both K400 and SSv2 datasets. We see that RA still achieves the best performance among attention approximations albeit falling behind the exact softmax attention. Since Motionformer is pretrained on images with softmax attention, this gap is likely introduced by employing a different attention mechanism during training the model further on video datasets. Besides, LARA outperforms Performer and Nyströmformer (Xiong et al., 2021) by a large margin on both K400 and SSv2 datasets. Although achieving strong performance, Orthoformer (Patrick et al., 2021) runs much slower (roughly 3×3\times or more) than other attention variants due to its sequential nature. As a result, LARA achieves better trade-offs than these baselines between predictive accuracy and efficiency.

4 Machine Translation

In this section, we conduct experiments on WMT14 EN–DE machine translation benchmark (Bojar et al., 2014) to evaluate the performance of our model under various sequence lengths. We follow Vaswani et al. (2017) and Ott et al. (2018) to preprocess this dataset, resulting in about 4.5M/3K/3K sentences pairs for training/validation/testing splits respectively. We adopt the standard transformer base architecture (Vaswani et al., 2017) and replace encoder self-attention with efficient attention variants. More detailed configurations are deferred to § H.4.

Table 4 presents the test BLEU scores under different attention mechanisms. Since this dataset consists mostly of short sentences, we set the number of samples to be relatively smaller. However, the training of Performer is quite unstable and a larger number of samples is required to mitigate this issue. Besides, we observe a similar trend that replacing the standard softmax attention with Performer leads to a significant performance drop, while increasing the number of samples does not improve the translation quality. RA, on the other hand, even outperforms softmax attention by over 0.3 BLEU score, clearly demonstrating the modeling capacity of unbiased approximations. LARA reaches performance close to softmax attention while runs with the same complexity as Performer; compared to other attention variants, LARA outperforms both Linformer (Wang et al., 2020) and ABC (Peng et al., 2021a) while obtaining similar BLEU scores to Nyströmformer (Xiong et al., 2021). This indicates RA and LARA are also capable of modeling natural language, which is typically hierarchically structured.

5 Analysis on Time and Memory Consumption

To evaluate the empirical efficiency of various attention methods, we conduct a simulation on a standard transformer architecture and report the running time and memory consumption under different sequence lengths. The detailed setup can be found in § H.5. As shown in Figure 2 (and Table 6 in § H.5 for exact statistics), we note that RA runs twice (or more) as slow as ordinary softmax attention with about 2.5×2.5\times memory consumption. This is as expected since RA needs to first compute full softmax probabilities to sample from pnp_{n}, and then compute fnf_{n}, both of which take a similar amount of computation to softmax attention. Nevertheless, its efficient variant LARA runs as fast as Performer with marginally increased memory usage. As for another baseline Nyströmformer (Xiong et al., 2021), which we found is a strong baseline and is used across experiments, it runs much slower than other variants at relatively short sequence lengths (e.g., less than 8192). Overall, the comparison result validates that LARA achieves a good balance between efficiency and expressiveness.

Related Work

Transformer models (Vaswani et al., 2017) are difficult to scale to long sequences due to the quadratic time and space complexity of self-attention mechanisms. Recently, a significantly large number of approaches have been proposed to improve the efficiency of attention mechanisms. A widely adopted paradigm is to utilize sparse attention, where each query is limited to only attend a subset of tokens. Such sparse attentive patterns can be pre-defined, such as sliding windows (Beltagy et al., 2020) or block-wise local chunks (Liu* et al., 2018; Parmar et al., 2018; Child et al., 2019; Ainslie et al., 2020; Zaheer et al., 2020; Liu et al., 2021); alternatively, the model can adaptively select tokens to take into account. This can be done via a trainable top-kk selecting operator (Pietruszka et al., 2020), learnable hash functions (Kitaev et al., 2020; Daras et al., 2020), clustering with K-Means (Vyas et al., 2020; Roy et al., 2021) or grouping tokens with a differentiable sorting module (Tay et al., 2020a). More recently, Combiner (Ren et al., 2021) is proposed to apply the sparse mechanism to factorize the softmax probability distribution so that the resulting approximation runs with sub-quadratic time but achieves full attention capacity.

Low-rank approximations to the softmax attention also received considerable interest. For instance, the Nyström method can be adopted to approximate the softmax attention map by a sub-sampled matrix (Xiong et al., 2021). Another approach is the kernel linearization, which aims to decompose the exponential kernel into a dot product of feature maps. Such feature maps can be randomized that yield unbiased estimates of exponential kernels (Choromanski et al., 2021; Peng et al., 2021b), or deterministic that enjoy better training convergence (Katharopoulos et al., 2020; Kasai et al., 2021b; Schlag et al., 2021). Alternatively, one can use a learnable matrix (including Linformer (Wang et al., 2020) and ABC (Peng et al., 2021a)) or other downsampling operations (Dai et al., 2020; Wang et al., 2021a, b) to project the key-value pairs into fixed-length sequences. Besides, a set of auxiliary points can also be incorporated to cache the information from the long sequence via an attention mechanism, which is adopted in LUNA (Ma et al., 2021), Set transformer (Lee et al., 2019) and Perceiver (Jaegle et al., 2021a, b). Our work falls into the category of kernel linearization methods, but in contrast to previous works, we propose an unbiased estimation for the whole softmax attention, which has not been explored and is orthogonal to previous works.

Recent studies also consider combining both the sparse and low-rank bias to achieve better approximation (Nguyen et al., 2021; Zhu et al., 2021; Chen et al., 2021a), or replace the softmax attention with other token-mixing mechanisms (Lee-Thorp et al., 2021; Lu et al., 2021; Chen et al., 2021d; Tay et al., 2021a). We refer readers to Tay et al. (2020b, 2021b); Lin et al. (2021) for a more detailed review on advances in the topic of efficient attention.

Conclusion

In this paper, we revisit the recently proposed random feature methods for approximating the softmax attention. By recasting RFA as self-normalized importance samplers, we identify an elusive bias in its approximation process. Built on this finding, we propose the unbiased estimation, called randomized attention (RA), which constructs positive random features via query-specific distributions. We then develop a novel linear complexity self-attention mechanism called linear randomized attention (LARA), which combines the expressiveness in RA and the efficiency in RFA. Extensive experiments demonstrate the effectiveness of RA and LARA, across various domains.

Acknowledgements

We thank Jianbo Yuan, Xiang Gao, Xiujun Li, Yanghua Peng, Ding Zhou, and Ruofan Ding for helpful discussions and feedback on early drafts of this paper. This research was supported in part by the joint research scheme of the National Natural Science Foundation of China (NSFC) and the Research Grants Council (RGC) under grant number N_HKU714/21.

References

Appendix A Proof for Proposition 3.1

Assume q(ω)=N(ω;0,I)q(\omega)=\mathcal{N}(\omega;0,\mathbf{I}). Recall that we define

Substituting the second equality into the first one yields the form of f(ω)f(\omega):

To more clearly illustrate the connection between RFA and RA, one can also manually verify this by rearranging terms in Equation 2:

where ZZ is the partition function. Recall that in Equation 1

which is effectively a mixture distribution and each component is selected with probability proportional to the similarity of queries and keys. As long as the randomized mapping is non-negative, p(ω∣m)p(\omega|m) would be a valid probability distribution since its density would be non-negative and integrate to 1, according to Equation 14.

In terms of the particular form of the distribution, we have the following lemma:

Then ω∼N(ω;qn+km,I)\omega\sim\mathcal{N}(\omega;\boldsymbol{\mathbf{q}}_{n}+\boldsymbol{\mathbf{k}}_{m},\mathbf{I}).

Note that q(ω)=N(ω;0,I)=1(2π)d/2exp⁡(−12ω⊤ω)q(\omega)=\mathcal{N}(\omega;0,\mathbf{I})=\frac{1}{(2\pi)^{d/2}}\exp\left(-\frac{1}{2}\omega^{\top}\omega\right). Based on the “complete the square” technique, we have

which is exactly the density function of a multivariate Gaussian with the mean qn+km\boldsymbol{\mathbf{q}}_{n}+\boldsymbol{\mathbf{k}}_{m} and covariance I\mathbf{I}. ∎

Following Lemma A.1, it is straightforward to obtain

where πm=exp⁡(qn⊤km)∑m′=1Mexp⁡(qn⊤km′)\pi_{m}=\frac{\exp\left(\boldsymbol{\mathbf{q}}_{n}^{\top}\boldsymbol{\mathbf{k}}_{m}\right)}{\sum_{m^{\prime}=1}^{M}\exp\left(\boldsymbol{\mathbf{q}}_{n}^{\top}\boldsymbol{\mathbf{k}}_{m^{\prime}}\right)}.

Due to the dependence on the randomized mapping ξ(⋅,⋅)\xi(\cdot,\cdot), different choices of feature maps would yield distinct density forms. Here we mainly study the positive randomized mapping in Performer (Choromanski et al., 2021) and leave other choices (such as trigonometric functions in Peng et al. (2021b)) as future work.

Appendix B Proof for Proposition 3.2

Since the vanilla random-feature-based attention estimation is consistent, softmax attention must be equal to expected randomized attention. However, such equality can also be verified as follows. Assume q(ω)=N(ω;0,I)q(\omega)=\mathcal{N}(\omega;0,\mathbf{I}) and p(ω)=q(ω)∑m=1Mξ(qn,ω)⊤ξ(km,ω)∑m′=1Mexp⁡(qn⊤km′)p(\omega)=q(\omega)\frac{\sum_{m=1}^{M}\xi(\boldsymbol{\mathbf{q}}_{n},\omega)^{\top}\xi(\boldsymbol{\mathbf{k}}_{m},\omega)}{\sum_{m^{\prime}=1}^{M}\exp(\boldsymbol{\mathbf{q}}_{n}^{\top}\boldsymbol{\mathbf{k}}_{m^{\prime}})} given by Proposition 1. Then we have

In addition, according to the definition of randomized mappings ξ(⋅,⋅)\xi(\cdot,\cdot),

Equipped with these helpers, we are ready to derive the equality as follows:

Appendix C Discussion on Different Randomized Mappings

The randomized mapping ξ(⋅,⋅)\xi(\cdot,\cdot) transforms the inputs to a ll-dimensional vector. There are various choices of ξ(⋅,⋅)\xi(\cdot,\cdot) for the resulting estimator to become unbiased in the context of attention mechanisms, such as

l=1l=1 and ξ(x,ω)=exp⁡(ω⊤x−∥x∥22)\xi(\boldsymbol{\mathbf{x}},\omega)=\exp{\left(\omega^{\top}\boldsymbol{\mathbf{x}}-\frac{\lVert\boldsymbol{\mathbf{x}}\rVert^{2}}{2}\right)} in Choromanski et al. (2021);

l=1l=1 and ξ(x,ω)=2exp⁡(∥x∥22)cos⁡(ω⊤x+b)\xi(\boldsymbol{\mathbf{x}},\omega)=\sqrt{2}\exp{\left(\frac{\lVert\boldsymbol{\mathbf{x}}\rVert^{2}}{2}\right)}\cos{\left(\omega^{\top}\boldsymbol{\mathbf{x}}+b\right)} with b∼Uniform⁡(0,2π)b\sim\operatorname{Uniform}(0,2\pi) in Rahimi & Recht (2008);

l=2l=2 and ξ(x,ω)=[exp⁡(∥x∥22)sin⁡(ω⊤x),exp⁡(∥x∥22)cos⁡(ω⊤x)]\xi(\boldsymbol{\mathbf{x}},\omega)=\left[\exp{\left(\frac{\lVert\boldsymbol{\mathbf{x}}\rVert^{2}}{2}\right)}\sin{\left(\omega^{\top}\boldsymbol{\mathbf{x}}\right)},\exp{\left(\frac{\lVert\boldsymbol{\mathbf{x}}\rVert^{2}}{2}\right)}\cos{\left(\omega^{\top}\boldsymbol{\mathbf{x}}\right)}\right] in Rahimi & Recht (2008); Peng et al. (2021b);

l=2l=2 and ξ(x,ω)=[12exp⁡(ω⊤x−∥x∥22),12exp⁡(−ω⊤x−∥x∥22)]\xi(\boldsymbol{\mathbf{x}},\omega)=\left[\frac{1}{\sqrt{2}}\exp{\left(\omega^{\top}\boldsymbol{\mathbf{x}}-\frac{\lVert\boldsymbol{\mathbf{x}}\rVert^{2}}{2}\right)},\frac{1}{\sqrt{2}}\exp{\left(-\omega^{\top}\boldsymbol{\mathbf{x}}-\frac{\lVert\boldsymbol{\mathbf{x}}\rVert^{2}}{2}\right)}\right] in Choromanski et al. (2021).

In the main paper, we focus on the positive randomized mappings (Choromanski et al., 2021); for other positive randomized mappings, it is also possible to derive a similar target expectation, such as the hyperbolic randomized mapping proposed in Choromanski et al. (2021):

Consider the hyperbolic randomized mapping

Consider the hyperbolic positive randomized mapping

According to proof of Proposition 3.1 in Appendix A, the density function p(ω)p(\omega) corresponding to the hyperbolic randomized mapping should also be a mixture (Equation 16) with the following form

where πm≔exp⁡(qn⊤km)∑m′=1Mexp⁡(qn⊤km′)\pi_{m}\coloneqq\frac{\exp(\boldsymbol{\mathbf{q}}_{n}^{\top}\boldsymbol{\mathbf{k}}_{m})}{\sum_{m^{\prime}=1}^{M}\exp(\boldsymbol{\mathbf{q}}_{n}^{\top}\boldsymbol{\mathbf{k}}_{m^{\prime}})} and p(ω∣m)p(\omega|m) denotes the density of the mm-th component distribution. By substituting the hyperbolic randomized mapping into the equation above, we have

It is straightforward to recognize that this can be viewed as the sum of two densities. We then invoke Lemma A.1 for each of them, which results in two Gaussians

Therefore, the true density function p(ω)p(\omega) can be expressed as follows

However, it is much more difficult to analyze the classical random Fourier mappings (Rahimi & Recht, 2008) since they may involve a negative density. As a result, the formulation of RFA with these randomized mappings may not define a valid self-normalized importance sampling estimate. We study positive randomized mappings through this paper and leave investigation into other cases as future work.

Appendix D Analysis on the Optimal Weighting Function in Multiple Importance Sampling

In this section, we analyze the optimal weighting function in MIS, which is self-normalized in our setting (§4.2).

Given the set of NN queries Q\boldsymbol{\mathbf{Q}} and the set of MM key-value pairs K\boldsymbol{\mathbf{K}} and V\boldsymbol{\mathbf{V}}, the regular softmax attention can be expressed as expected randomized attention according to Equation 7:

where the distribution is defined in Proposition 3.1 as

The attention mechanism outputs a DD-dimensional vector for each query. For brevity, we start with considering the dd-th dimension and denote fn,d(ω)f_{n,d}(\omega) as the dd-th dimension of the function output at query position nn. We then have

In our work, we estimate the expectation above by self-normalized multiple importance sampling (see §4.2). For the dd-th dimension of the output at query position nn, we have

where ωc∼pc(ω)\omega_{c}\sim p_{c}(\omega) for c=1,…,Cc=1,\dots,C. We also let AA and BB represent the nominator and denominator respectively. The expectations of AA and BB are

Unfortunately, the exact form of the variance of g^n,d\hat{g}_{n,d} is mostly intractable to compute. To this end, we follow previous practices (Owen, 2013) and approximate Var⁡[g^n,d]\operatorname{Var}\left[\hat{g}_{n,d}\right] via the delta method. In particular, we apply the first-order Taylor expansion approximation to the function g(A,B)≔A/Bg(A,B)\coloneqq A/B around point (μA,μB)\left(\mu_{A},\mu_{B}\right), yielding

where we denote g_{A}\coloneqq\frac{\partial g(A,B)}{\partial A}\Bigr{|}_{\begin{subarray}{c}A=\mu_{A}\\ B=\mu_{B}\end{subarray}} and g_{B}\coloneqq\frac{\partial g(A,B)}{\partial B}\Bigr{|}_{\begin{subarray}{c}A=\mu_{A}\\ B=\mu_{B}\end{subarray}} similarly. Note that both gAg_{A} and gBg_{B} are constants with respect to ω\omega. According to Equation 19, the approximate expectation is the following

The first three lines hold since ωc\omega_{c} is independent of ωc′\omega_{c^{\prime}} for any c≠c′c\neq c^{\prime}. Therefore, the approximate variance of our estimate at the dd-th dimension can be written as

Since we are using the same proposal distribution to estimate the output for all dimensions, we are interested in the sum of variance over every dimension (i.e., the trace of the covariance matrix):

According to our design choice, αnc(⋅)\alpha_{nc}(\cdot) is specific to each query position. Ideally, we hope these weighting functions can minimize the sum of variance at each position. Formally, we have

Optimizing weighting functions to minimize the variance of ordinary MIS estimator has been studied by a recent work (Kondapaneni et al., 2019). Our setting is different from it in that (1) we focus on self-normalized MIS and that (2) the function f(⋅)f(\cdot) is vector-valued instead of scalar-valued. These differences lead to a distinct objective (Equation 20). Here we adapt the analysis (Kondapaneni et al., 2019) to solve this problem. In particular, we rely on the calculus of variations and introduce the following Lagrangian

Solving ∂L(α,λ)∂αnc=0\frac{\partial\mathcal{L}(\alpha,\lambda)}{\partial\alpha_{nc}}=0 and ∂L(α,λ)∂λ=0\frac{\partial\mathcal{L}(\alpha,\lambda)}{\partial\lambda}=0 respectively yields

Here we denote rncd≔∫αnc(ω)pn(ω)(fn,d(ω)−μn,d)dωr_{ncd}\coloneqq\int\alpha_{nc}(\omega)p_{n}(\omega)\left(f_{n,d}(\omega)-\mu_{n,d}\right)d\omega. We then rearrange Equation 23 to obtain

Substituting Equation 25 into Equation 24 gives

Substituting Equation 26 back into Equation 25 yields

The characteristic of existence and uniqueness of the optimal weighting function is similar to Kondapaneni et al. (2019). Intuitively, the optimal weighting functions can be obtained by first calculating a query-dependent correction term, which sums to 0, and then adding such correction to the original balance heuristic weighting function. For large rncr_{nc}, the correction term will be positive, driving the weights for the cc-th proposal to be higher; and vice versa. Such formulation introduces the dependence between the current proposal index cc and the target query position nn, which allows the weighting functions αnc\alpha_{nc} (and thus the estimator) to specialize in the current query.

To obtain the exact form of rnc(ω)r_{nc}(\omega), we need to solve rncd=∫αnc(ω)pn(ω)(fn,d(ω)−μn,d)dωr_{ncd}=\int\alpha_{nc}(\omega)p_{n}(\omega)\left(f_{n,d}(\omega)-\mu_{n,d}\right)d\omega. However, deriving a closed form solution is mostly intractable given its complex structure, which not only involves an intractable integral but also mixes together the effect from different dimensions. To further analyze this problem, we start with a simplified case where D=1D=1. In this setting, we have the following:

for any c′=1,…,Cc^{\prime}=1,\dots,C. Although solving this linear system is intractable, it indicates that rncDr_{ncD} roughly describes how qc(ω)q_{c}(\omega) aligns with pn(ω)(fn,d(ω)−μn,d)p_{n}(\omega)(f_{n,d}(\omega)-\mu_{n,d}) under the expectation of different qc′q_{c}^{\prime}. Therefore, rnc(ω)r_{nc}(\omega) can be seen as an indicator for the correlation between the current proposal qc(ω)q_{c}(\omega) and pn(ω)(fn,d(ω)−μn,d)p_{n}(\omega)(f_{n,d}(\omega)-\mu_{n,d}) that is normalized by the strength of pn(ω)(fn,D(ω)−μn,D)p_{n}(\omega)\left(f_{n,D}(\omega)-\mu_{n,D}\right).

For larger DD, such concise equality involving rncdr_{ncd} is not available since the effect of different dimensions is mixed. We thus seek an heuristic approximation that not only reflects the same intuition but also becomes tractable in practice (see § G.3.2 for practical implementations).

Appendix E Derivation for the Formulation of LARA

In this section, we give the detailed derivation for the final expression of our estimator LARA:

The formulation (Equation 28) is obtained by substituting the equations above into the self-normalized estimator:

Note that we define αnc′(ωc)≔αnc(ωc)N(ωc;0,I)qc(ωc)\alpha_{nc}^{\prime}(\omega_{c})\coloneqq\alpha_{nc}(\omega_{c})\frac{\mathcal{N}(\omega_{c};0,\mathbf{I})}{q_{c}(\omega_{c})}.

Appendix F Proof for the Unbiasedness of Multiple Importance Sampling

As in §4.1, suppose our MIS estimator takes the following form

If ∑c=1Cαnc(ω)=1\sum_{c=1}^{C}\alpha_{nc}(\omega)=1, it can be shown that (Veach & Guibas, 1995)

Appendix G Details of RA, RFA and LARA

Some implementations of RFA (including Performer (Choromanski et al., 2021)) defines a sample-redrawing schedule, where the involved samples ω\omega are periodically redrawn according to a hand-crafted strategy. However, this requires a task-specific specification and we found tuning redrawing strategies only brings marginal performance gain over the simplest method that redraws samples at each training iteration (we use the same sample set during the entire evaluation phase). Therefore, we adopt this method to train Performer for all tasks. We also do not use orthogonal random samples as in Choromanski et al. (2021), as we found it does not improve empirical performance but increases the training time. Algorithm 2 provides a algorithm sketch for random feature attention and linear randomized attention, respectively. Note that every loop involved in all provided pseudo-codes (Algorithm 1, Algorithm 2 and Algorithm 3) can be trivially executed in parallel.

G.2 Specifics of Randomized Attention

In this section, we describe the details of RA approximation for softmax attention. Recall in Proposition 3.1 the RA sampling distribution is a Gaussian mixture

with πnm=exp⁡(qn⊤km)∑m′=1Mexp⁡(qn⊤km′)\pi_{nm}=\frac{\exp\left(\boldsymbol{\mathbf{q}}_{n}^{\top}\boldsymbol{\mathbf{k}}_{m}\right)}{\sum_{m^{\prime}=1}^{M}\exp\left(\boldsymbol{\mathbf{q}}_{n}^{\top}\boldsymbol{\mathbf{k}}_{m^{\prime}}\right)} and μnm=qn+km\boldsymbol{\mathbf{\mu}}_{nm}=\boldsymbol{\mathbf{q}}_{n}+\boldsymbol{\mathbf{k}}_{m}. To sample from this Gaussian mixture distribution, we first sample zn∼Categorical⁡(z;πn)z_{n}\sim\operatorname{Categorical}(z;\boldsymbol{\mathbf{\pi}}_{n}) with πn\boldsymbol{\mathbf{\pi}}_{n} being the probability masses at MM possible outcomes and then let an≔[an1,…,anM]\boldsymbol{\mathbf{a}}_{n}\coloneqq[a_{n1},\dots,a_{nM}] be an MM-dimensional one-hot vector with anzn=1a_{nz_{n}}=1. The discrete random variable an\boldsymbol{\mathbf{a}}_{n} defines which distribution component is selected. Since all components are Gaussian, we leverage reparameterization trick (Kingma & Welling, 2013; Rezende et al., 2014; Titsias & Lázaro-Gredilla, 2014) to draw independent ϵ∼N(ω;0,I)\epsilon\sim\mathcal{N}(\omega;0,\mathbf{I}) and add it to the selected mean, resulting in the final mixture sample. Formally, we express the sample ωn\omega_{n} from the Gaussian mixture as follows:

which is then used to compute fn(ωn)f_{n}(\omega_{n}) to obtain the RA estimation (see Algorithm 1 for a algorithm sketch). Assuming the number of samples is SS and the sequence length is NN, the overall time/space complexity for RA is O(SN2)\mathcal{O}(SN^{2}). Through experiments we take S=1S=1 sample in our randomized attention unless specified otherwise. We found this choice suffices to achieve good performance and increasing SS does not greatly improve the performance but introduces significant time/memory overheads.

Exact sampling from the mixture distribution requires us to first select a discrete component index a\boldsymbol{\mathbf{a}} from the mixture distribution and then sample from the corresponding component. Although such randomness might bring additional regularization effect, randomly selecting an index could lead to large variance and slow down training. To accelerate convergence, we also develop a biased sampling strategy from the Gaussian mixture. According to Equation 30, the sampled one-hot vector an\boldsymbol{\mathbf{a}}_{n} can be approximated by its expected value πn\boldsymbol{\mathbf{\pi}}_{n}:

This introduces a non-negligible sampling bias in estimating the softmax attention; however, it eliminates the need to randomly draw discrete indexing vectors an\boldsymbol{\mathbf{a}}_{n} and reduces the variance, especially in the case of long sequences. In fact, this biased sample can be equivalently viewed as drawn from a Gaussian:

Another advantage is that this formulation allows us to maintain fully deterministic during the evaluation mode, while not introducing large discrepancies from training time. Specifically, during evaluation we only pass the expectation Kπn+qn\boldsymbol{\mathbf{K}}\boldsymbol{\mathbf{\pi}}_{n}+\boldsymbol{\mathbf{q}}_{n} as the “sample”, which is a standard practice similar to the usage of Dropout (Srivastava et al., 2014). This is in contrast to unbiased RA sampling, which has to draw random indices even during evaluation (otherwise, replacing both random variables with their expected values would lead to larger discrepancies between training and testing, resulting in inferior performance). Same as the case of RFA, we also redraw random samples at every training iteration. Note that this can not be transferred to Performer since the expectation of ω\omega in RFA is 0, which leads to degeneration.

As a proof-of-concept experiment, we run randomized attention with biased sampling strategy on image classification with ImageNet1k dataset, video recognition with K400 and SSv2 datasets and machine translation with WMT dataset. From Table 5, we note that biased RA performs better than both its unbiased counterpart for visual tasks, which usually deal with longer sequences (196 for images and 1568 for videos); but it performs worse in machine translation, where either the source or target sentence only consist of dozens of tokens. On the other hand, RA outperforms softmax attention on both image and language tasks, indicating that the proposed estimation methods for softmax attention may enjoy better modeling capacity. This may shed light on some latent mechanism in such approximation that deviates from the standard softmax attention but does better in modeling the sequence representations. We leave detailed investigation in future work.

G.3 Specifics of Linear Randomized Attention

In this section, we provide more implementation details of linear randomized attention.

As mentioned in §4.1, each proposal qc(ω)q_{c}(\omega) is defined to depend on some subset of queries; and their union covers the whole set of queries. Since our goal is let these proposals behave similarly to the true RA distribution pn(ω)p_{n}(\omega), a straightforward choice is to specify qcq_{c} as the same formulation of pn(ω)p_{n}(\omega) (Equation 29):

Here we divide the input query sequence {qn}n=1N\{\boldsymbol{\mathbf{q}}_{n}\}_{n=1}^{N} into CC segments and compute the average (called landmarks, the number of which is equal to the number of samples) over queries {q~c}c=1C\{\widetilde{\boldsymbol{\mathbf{q}}}_{c}\}_{c=1}^{C} within the same segment. In particular, supposing NN is divisible by CC and T≔N/CT\coloneqq N/C is the segment length, each segment landmark can be expressed as

We then use each of these proposals to estimate the target expectation for the nn-th query and combine their results into the final estimation. However, this choice involves CMCM distributions in total (CC proposals are maintained, each of which is again a Gaussian mixture with MM components) and sampling from these distributions may introduce large noise. Motivated by the discussion of biased sampling in RA (Equation 31 in § G.2), we explore an alternative parameterization by defining each proposal as a Gaussian:

We find this choice performs better than the mixture formulation (Equation 32) empirically. Intuitively, this strategy aggregates the information from all keys based on the correlation between the query landmarks and each individual key. However, this introduces additional O(CM)\mathcal{O}(CM) computational costs.

In practice, we observe that for proposal landmark q~c\widetilde{\boldsymbol{\mathbf{q}}}_{c}, keys belonging to the same segment cc often contribute the most to the Gaussian mean. As a result, we develop another variant that also computes the key landmarks,

We observe this formulation works equally well; such parameterization is thus used throughout our experiments by default.

Comparing Equation 33 and Equation 34, we observe that for the former it only biases the Gaussian mean towards the direction of the current query landmark; while for the latter it only promotes information from key vectors that are in the same segment as q~c\widetilde{\boldsymbol{\mathbf{q}}}_{c} and ignores the global information of keys. Noticing these differences, we further propose a variant bridging these two formulations:

Intuitively, this performs an attention-like aggregation operation over key landmarks. The aggregation procedure not only computes the correlation between key vectors, which alleviates the bias of being closer to query landmarks, but also collects global information while still favoring local segments. In addition, it runs with O(C2)\mathcal{O}(C^{2}), which is much cheaper than O(CM)\mathcal{O}(CM). We find this yields better predictive performance in vision transformers, but improves marginally for other tasks. We hypothesize that this is because the attention-like operation smooths the Gaussian mean, which aligns with that ViT tends to produce smoothed patch representations. We leave in-depth investigation as future work. In summary, we adopt this parameterization only through experiments on image classification (§5.2).

See Algorithm 3 for a algorithm sketch of LARA.

G.3.2 On the Parameterization of Weighting Functions

Our MIS estimating strategy introduces a set of weighting functions α(⋅)\alpha(\cdot) for each proposal. A common choice of weighting functions in MIS (Owen, 2013) is the balance heuristic strategy

which is nearly optimal in that any other weighting schemes will not exhibit significantly smaller variance (Veach & Guibas, 1995). However, this strategy only considers the relative strengths of proposals and ignores contextual information from each query. As a result, a naïve application of MIS would disregard the inherent variation among different queries and fails to describe the specialized target distribution pn(ω)p_{n}(\omega).

Instead of balance heuristics, we adopt query-specific weighting functions that are inspired by query-optimal analysis. In our MIS scheme (Equation 9), the optimal weighting functions take the following form

Note that it sums to 1 over all cc’s and is a valid weighting function:

In particular, we observe the first term is the ordinary balance heuristic weighting function, while the second term is a query-specific correction that sums to 0.

As mentioned in Appendix D, the exact form of rnc(⋅)r_{nc}(\cdot) is mostly intractable to compute in practice. To this end, we introduce a heuristic yet tractable rnc′r^{\prime}_{nc} to roughly align with the intuition of original rnc(⋅)r_{nc}(\cdot):

Intuitively, we implement rnc′r^{\prime}_{nc} as the normalized similarity between the nn-th query and the cc-th segment-averaged query vector. In addition, we note that the query-specific information rnc′r^{\prime}_{nc} is influenced by the query-agnostic density qcq_{c}, which may be incorrectly suppressed or amplified if the drawn sample lies in a low-density region. Base on this, we further propose a simplified formulation:

where we decouple the computation between proposal densities qc(⋅)q_{c}(\cdot) and rnc′r^{\prime}_{nc}. In this way, query-dependent and query-agnostic information will be independent of each other.

We also notice that the query-specific information can be explicitly controlled by introducing a coefficient β\beta such that

This weighting function remains valid since the correction term still sums to 0. By setting β>1\beta>1, the mechanism tends to favor the query-specific information over the balance heuristic. We tried several choices of β\beta and found β=2\beta=2 slightly improves the performance. As reflected in our ablation study, we demonstrate the superior performance of query-specific weighting functions over vanilla balance heuristics (Veach & Guibas, 1995).

G.3.3 Training and Evaluation Details

LARA redraws samples from proposal sets at every training iteration; during evaluation, we simply pass corresponding expected values instead of drawing samples, in a similar way to dropout (Srivastava et al., 2014).

G.3.4 Complexity Analysis

Recall there are NN queries and MM key-value pairs. Like RFA, the involved computation of our LARA estimator includes (1) computing the proposal distribution, which may take O(C)\mathcal{O}(C) or O(C2)\mathcal{O}(C^{2}) time (§ G.3.1); (2) a pre-computing step over all key-value statistics, which takes O(CM)\mathcal{O}(CM) time and space; and (3) applying pre-computed statistics to all queries, taking O(CN)\mathcal{O}(CN) complexity. These steps result in overall O(CM+CN)\mathcal{O}(CM+CN) complexity given C≪min⁡(M,N)C\ll\min(M,N). Note that CC is analogous to the number of samples SS (often referred to as random feature dimension (Choromanski et al., 2021)) in RFA. Therefore, LARA does not incur a heavy computational overhead compared to RFA, as also reflected in §5.5.

Appendix H Additional Experimental Details

We conduct the preliminary experiment on vision transformers (ViT), which first split input images into small patches, serialize them as a 1D sequence and then processes the sequence through a transformer model. To be specific, we replace the standard softmax attention in vision transformers (ViT; Dosovitskiy et al., 2021; Touvron et al., 2021) with different approximation variants. The MSE is evaluated under three different sequence lengths NN by varying the image resolution and patch size: (a) resolution 224 x 224 with patch size 16 (N=196N=196), (b) resolution 384 x 384 with patch size 16 (N=576N=576) and (c) resolution 224 x 224 with patch size 8 (N=784N=784). To achieve a fair comparison, we use pretrained ViTs whose weights are trained under corresponding sequence lengths with standard softmax attention. The sequence length is selected according to whether the ViT weights pretrained by softmax attention are available. Since there are multiple attention blocks in ViT architecture, for each input image we average the attention MSE over all attention heads and transformer layers.

H.2 Image Classification

For image classification, we consider two vision transformer architectures: vanilla ViT (Dosovitskiy et al., 2021) and PVTv2 (Wang et al., 2021b). We refer to ViT as DeiT (Touvron et al., 2021) through this work, since DeiT follows the same model architecture as ViT but adopts greatly improved training protocols.

We do not use the distillation technique as in DeiT (Touvron et al., 2021). We following the same procedure to train DeiT on ImageNet1k dataset as in Touvron et al. (2021). In particular, we use AdamW optimizer (Loshchilov & Hutter, 2019) for 300 epochs, where we set the batch size to 1024 and the learning rate to 0.001 with cosine learning rate decay (Loshchilov & Hutter, 2016). The number of warm-up epochs is set to 10 for all models instead of 5, since we find it often stabilizes training and leads to better results. For data augmentation, we follow Touvron et al. (2021) and use random clipping, cropping, rand-augment (Cubuk et al., 2020) and random erasing (Zhong et al., 2020). We remove repeated augmentation (Hoffer et al., 2020) as it often slows down convergence, as also observed in previous studies (Berman et al., 2019; Xiao et al., 2021). For regularization, we employ stochastic depth (Huang et al., 2016), Mixup (Zhang et al., 2017), Cutmix (Yun et al., 2019), label smoothing and weight decay, all of which are set to default settings in DeiT (Touvron et al., 2021). Unless otherwise specified, the input image size is set to 224×224224\times 224 with patch size 1616, resulting in 14×14=19614\times 14=196 non-overlapping patch tokens. For LARA in DeiT models, we additionally transform the average query/key vector of each segment through a fully connected layer followed by a layer-norm operation. This corresponds to importance sampling with adaptive proposals (Owen, 2013), which improves the expressiveness of the proposal distributions. Note that the linear transformation is shared among all attention heads, which results in only marginal additional computational overheads.

Pyramid Vision Transformers v2 (PVTv2; Wang et al., 2021b) is a strong vision transformer baseline with pyramidal architectures that processes much longer token sequences. It first patchifies input images into a 56×5656\times 56 token sequence, which is then processed by 4 successive stages. Each stage consists of a stack of transformer layers and processes the input sequence from the previous stage by reducing both the height and width of patch tokens to the half and increasing the embedding dimension by a factor of 2. The detailed configuration for all model sizes follows Table 1 of Wang et al. (2021b). In such architecture, the sequence at early stages is too long to be handled by regular softmax attention. To address this issue, PVTv2 proposes an efficient variant Spatial Reduction Attention (SRA) and uses SRA for all attention blocks in the first three stages and ordinary softmax attention for the last stage due to reduced resolution. For each SRA module, it use a convolutional layer to reduce the length of input sequence to 49, which is then projected to key and value vectors correspondingly. The query set maintains the same resolution and performs attention over the shortened key-value sequence to obtain globally contextualized representations.

To evaluate our method on PVTv2, we replace all SRA modules with either Performer or LARA. For PVTv2 with Performer, we use 128 samples since it fails to converge with fewer samples. In terms of PVTv2 with LARA, we do not use convolutional blocks and simply use 2D average pooling (the same as segments) followed by a linear projection to obtain query and key landmarks, the number of which is set to 49 as in SRA. Since we do not use the convolutional block, Both Performer and LARA use much fewer model parameters than vanilla PVTv2.

In addition, vanilla PVTv2 model uses 1,2,5 and 8 attention heads for its 4 processing stages respectively in its original implementation; however, we found using 2×\times more heads consistently improves predictive performance for all methods (including baseline SRA, Performer and LARA) while introducing affordable overheads. Therefore, we use 2,4,10 and 16 heads for all PVTv2-based models across our experiments. We mostly follow the training protocol as Wang et al. (2021b) to train all PVTv2-based models, except that we increase the number of warm-up epochs from 5 to 10. We find a slightly longer warm-up schedule is helpful to improve the model performance.

H.3 Video Action Recognition

Our implementation is based on the PySlowFast (Fan et al., 2020) codebase and we follow the training protocol in Motionformer (Patrick et al., 2021). In particular, Motionformer adopts the vision transformer base (ViT/B) (Dosovitskiy et al., 2021) architecture which has 12 transformer encoder layers with 12 attention heads and 768-dimensional hidden representations. For K400 dataset, its parameter weights are pretrained on ImageNet21k dataset with regular softmax attention; while for SSv2 dataset, we use the trained weights on K400 with the corresponding attention variant. The model operates on videos with size 16×224×22416\times 224\times 224, which is then split into 8×14×148\times 14\times 14 tubes with separate space-time positional embedding. Motionformer introduces the trajectory attention, which first computes spatial attention to obtain probabilistic trajectories, which are then aggregated temporally. We use the trajectory attention module (Patrick et al., 2021) and replace the involved softmax attention with different attention approximation methods. Besides Performer (Choromanski et al., 2021), Nyströmformer (Xiong et al., 2021) and full trajectory attention, our baselines also include Orthoformer (Patrick et al., 2021), another strong baseline for video transformers that constructs a low-rank approximation of attention matrix via sequentially selecting orthogonal query landmarks. For all efficient attention variants, we set both the number of samples (in LARA and Performer) or the number of landmarks (in Nyströmformer and Orthoformer) to 128 for a fair comparison. For data augmentation, we also follow Patrick et al. (2021), adopting random scale jittering, random horizontal flips and color jittering for all datasets; and additionally rand-augment (Cubuk et al., 2020) for SSv2 dataset.

We use the AdamW (Loshchilov & Hutter, 2019) optimizer to train LARA for 40 epochs with weight decay 0.05, label smoothing rate 0.2 and total batch size 256. A slightly longer training schedule (compared to 35) is adopted since our method involves additional randomness and we found training for a longer time improves convergence. The initial learning rate is set to 0.0001 and gets decayed by a ratio of 10 at epochs 25 and 35 respectively. During training, the video clips are randomly sampled with cropped resolution 224×224224\times 224; while for testing, we sample 10 uniform temporal clips per video with 3 spatial crops per clip and average scores for these crops to obtain the final prediction.

H.4 Machine Translation

We use the Transformer-base architecture as specified in Vaswani et al. (2017) for our machine translation experiments. The model contains a transformer encoder and decoder, both of which consist of 6 layers with hidden size and number of heads being 512 and 8, respectively. The vocabulary is shared between source and target language, consisting of around 32K byte pair encoding (BPE; Sennrich et al., 2016) types. The hidden dimension of feed forward networks is set to 2048. The rate of dropout is set to 0.1. As mentioned in the main paper, we only replace encoder self-attention in transformer models with efficient attention variants. Since LARA does not support causal attention mode in its current version, this setting allows us to directly assess the ability of different attention mechanisms to learn contextualized representations. Recent studies also indicate that in neural machine translation the transformer encoder seems playing a more important role in extracting representations (Kasai et al., 2021a). Besides Performer, we also compare our method against other baselines including (1) Linformer (Wang et al., 2020), which is widely adopted in the context of NLP, (2) ABC (Peng et al., 2021a), a recently proposed unified framework of most low-rank attention approximations and (3) Nyströmformer (Xiong et al., 2021), which we find is a strong low-rank baseline across our experiments.

For training, we follow the same setup as in Vaswani et al. (2017). In particular, we use the Adam optimizer (Kingma & Ba, 2014) with learning rate 0.0007, label smoothing rate 0.1, inverse square root learning rate scheduler and 4000 warm-up steps. During decoding, we set beam size to 4, length penalty to 0.6, average last 10 checkpoints and apply a compound split post-processing to facilitate comparison.

H.5 Efficiency Analysis

For the simulation experiment conducted in §5.5, the same transformer architecture is used for all attention methods, which consists of 8 encoder layers with 192 embedding dimension and 3 attention heads. The use of smaller-size transformer model allows us to run longer lengths for softmax attention and randomized attention. The detailed running time (in ms) and memory consumption is listed in Table 6. For Nyströmformer (Xiong et al., 2021) and Linformer (Wang et al., 2020), the number of landmarks is set to 16; for Performer and LARA, the number of samples is set to 16 as well.

Appendix I Additional Experimental Results

We conduct additional experiments to evaluate the performance of our proposed method at various aspects. First, we vary the number of random samples to investigate the effect of sample size on the performance for ImageNet1k dataset. As presented in Table 7, although both Performer and LARA improves predictive accuracy as the number of samples increases, LARA benefits much more than Performer, and finally outperforms softmax attention with 196 samples, which is equal to the sequence length.

In addition, we also compare LARA against different efficient attention mechanisms, as shown in Table 8. We note that LARA outperforms most efficient attention mechanisms by a large margin, and slightly outperforms Nyströmformer (Xiong et al., 2021), which we found is a strong baseline across various domains.

I.2 Additional Experiments on Long Range Arena Benchmark

We also evaluate our model on the Long Range Arena (LRA) benchmark (Tay et al., 2021b), which is designed to test the ability to process long sequences and generalize over diverse tasks. In particular, LRA is a suite of tasks including Listops output prediction (Nangia & Bowman, 2018), byte-level text classification on IMDb (Maas et al., 2011), byte-level document retrieval on AAN (Radev et al., 2013), pixel-level image recognition on CIFAR-10 (Krizhevsky et al., 2009) and Pathfinder (Linsley et al., 2018). We follow the experimental setup in Xiong et al. (2021); Chen et al. (2021d) and adopt the same hyper-parameter setting across all attention variants to ensure a fair comparison. In particular, all tasks use a 2-layer Transformer model with 64 embedding dimension, 128 hidden dimension in feed forward neural networks and 2 attention heads. The transformer output is then aggregated by mean pooling (instead of class tokens) for task-specific prediction. The training details for each task are the same for all attention methods as in Xiong et al. (2021). For baselines, we compare our model against the standard softmax attention and Performer (Choromanski et al., 2021) as well as other efficient attention mechanisms, including Nyströmformer (Xiong et al., 2021), Linformer (Wang et al., 2020), Reformer (Kitaev et al., 2020) and BigBird (Zaheer et al., 2020).

As shown in Table 9, we see that RA performs better than softmax attention on 3 out of 5 tasks and obtains a higher averaged accuracy. Furthermore, its linear-complexity counterpart LARA also performs competitively with or slightly outperforms softmax attention except the image task. Both RA and LARA yield better performance than Performer and other baselines on all of 5 tasks, indicating the improved expressiveness of our proposed method. As the sequence length considered in this suite of tasks is typically longer, these results also validates the ability of RA and LARA to capture longer-term dependencies.

I.3 Ablation Study

In this section, we conduct an ablation study on vision transformers with ImageNet1k dataset to investigate the effect of component design in LARA. The main component design choices in LARA consist of the estimation framework (single proposal versus multiple proposals), the parameterization of proposal distributions (Gaussian mixtures versus Gaussian) and the weighting functions. The results are shown in Table 10. For the estimation framework, we compare our choice, which uses multiple proposal distributions, against a single proposal. This proposal is a Gaussian mixture with the similar formulation of true RA density (Equation 5) except that it only depends on the average of all queries. We see that an individual yet contextual proposal improves the performance of Performer, while generalizing the importance sampling in RFA to multiple proposals further boosts performance to be close to softmax attention. With multiple proposal distributions, even using a simple strategy (balance heuristic (Veach & Guibas, 1995)) to combine their estimates yields reasonable performance, which is improved further by adopting query-specific combinations. In addition, we validate the effectiveness of decoupling the effect of query-dependent and query-agnostic information inside the weighting function, which improves the coupled version by over 0.4 accuracy. In terms of the parameterization of each proposal distribution, we consider both the cases where each proposal is a Gaussian mixture and a Gaussian. As specified in § G.3, we train the transformer model with various parameterization choices (defined by Equation 32 for Gaussian mixtures and Equations 33, 34 and 35 for Gaussians). The results are consistent with the analysis in § G.3, where a simple parameterization suffices to yield good performance.