KVT: k-NN Attention for Boosting Vision Transformers

Pichao Wang, Xue Wang, Fan Wang, Ming Lin, Shuning Chang, Hao Li, Rong Jin

Introduction

Traditional CNNs provide state of the art performance in vision tasks, due to its ability in capturing locality and translation invariance, while transformer is the de-facto standard for natural language processing (NLP) tasks thanks to its advantages in modelling long-range dependencies. Recently, various vision transformers have been proposed by building pure or hybrid transformer models for visual tasks. Inspired by the transformer scaling success in NLP tasks, vision transformer converts an image into a sequence of image patches (tokens), with each patch encoded into a vector. Since self-attention in the transformer is position agnostic, different positional encoding methods have been developed, and in their roles have been replaced by convolutions. Afterwards, all tokens are fed into stacked transformer encoders for feature learning, with an extra CLSCLS token or global average pooling (GAP) for final feature representation. Compared with CNNs, transformer-based models explicitly exploit global dependencies and demonstrate comparable, sometimes even better, results than highly optimised CNNs .

Albeit achieving its initial success, vision transformers suffer from slow training. One of the key culprits is the fully-connected self-attention, which takes all the tokens to calculate the attention map. The dense attention not only neglects the locality of images patches, an important feature of CNNs, but also involves noisy tokens into the computation of self-attention, especially in the situations of cluttered background and occlusion. Both issues can slow down the training significantly . Recent works try to mitigate this problem by introducing convolutional operators into vision transformers. Despite encouraging results, these studies fail to resolve the problem fundamentally from the transformer structure itself, limiting their success. In this study, we address the challenge by directly attacking its root cause, i.e. the fully-connected self-attention.

To this end, we propose the kk-NN attention to replace the fully-connected attention. Specifically, we do not use all the tokens for attention matrix calculation, but only select the top-kk similar tokens from the sequence for each query token to compute the attention map. The proposed kk-NN attention not only naturally inherits the local bias of CNNs as the nearby tokens tend to be more similar than others, but also builds the long range dependency by choosing the most similar tokens from the entire image. Compared with convolution operator which is an aggregation operation built on Ising model and the feature of each node is aggregated from nearby pixels, in the kk-NN attention, the aggregation graph is no longer limited by the spatial location of nodes but is adaptively computed via attention maps, thus, the kk-NN attention can be regarded as a relieved version of local bias. The similar idea is proposed in where the kk-NN attention is mostly evaluated on NLP tasks. Despite the similarity in terms of the calculation of top-kk, our work focuses on the recent vision transformers, makes a deep theoretical understanding and presents a thoroughly analysis by defining several metrics. We verify, both theoretically and empirically, that kk-NN attention is effective in speeding up training and distilling noisy tokens of vision transformers. Eleven different available vision transformer architectures are adopted to verify the effectiveness of the proposed kk-NN attention.

Related Work

Self-attention has demonstrated promising results on NLP related tasks, and is making breakthroughs in speech and computer vision. For time series modeling, self-attention operates over sequences in a step-wise manner. Specifically, at every time-step, self-attention assigns an attention weight to each previous input element and uses these weights to compute the representation of the current time-step as a weighted sum of the past inputs. Besides the vanilla self-attention, many efficient transformers have been proposed. Among these efficient transformers, sparse attention and local attention are one of the main streams, which are highly related to our work. Sparse attention can be further categorized into data independent (fixed) sparse attention and content-based sparse attention . Local attention mainly considers attending only to a local window size. Our work is also content-based attention, but compared with previous works , our kk-NN attention has its merits for vision domain. For example, compared with routing transformer that clusters both queries and keys, our kk-NN attention equals only clustering keys by assigning each query as the cluster center, making the quantization more continuous which is a better fitting of image domain; compared with reformer which adopts complex hashing attention that cannot guarantee each bucket contain both queries and keys, our kk-NN attention can guarantee that each query has number kk keys for attention computing. In addition, our kk-NN attention is also a generalized local attention, but compared with local attention, our kk-NN attention not only enjoys the locality but also empowers the ability of global relation mining.

2 Transformer for Vision

Transformer is an effective sequence-to-sequence modeling network, and it has achieved state-of-the-art results in NLP tasks with the success of BERT . Due to its great success, it has also be exploited in computer vision community, and ‘Transformer in CNN’ becomes a popular paradigm . ViT leads the other trend to use ‘CNN in Transformer’ paradigm for vision tasks . Even though ViT has been proved compelling in vision recognition, it has several drawbacks when compared with CNNs: large training data, fixed position embedding, rigid patch division, coarse modeling of inner patch feature, single scale, unstable training process, slow speed training, easily fitting data and poor generalization, shallow & narrow architecture, and quadratic complexity. To deal with these problems, many variants have been proposed . For example, DeiT adopts several training techniques and uses distillation to extend ViT to a data-efficient version; CPVT proposes a conditional positional encoding that is adaptable to arbitrary input sizes; CvT , CoaT and Visformer safely remove the position embedding by introducing convolution operations; T2T ViT , CeiT , and CvT try to deal with the rigid patch division by introducing convolution operation for patch sequence generation; Focal Transformer makes each token attend its closest surrounding tokens at fine granularity and the tokens far away at coarse granularity; TNT proposes the pixel embedding to model the inner patch feature; PVT , Swin Transformer , MViT , ViL , CvT , PiT , LeViT , CoaT , and Twins adopt multi-scale technique for rich feature learning; DeepViT , CaiT , and PatchViT investigate the unstable training problem, and propose the re-attention, re-scale and anti-over-smoothing techniques respectively for stable training; to accelerate the convergence of training, ConViT , PiT , CeiT , LocalViT and Visformer introduce convolutional bias to speedup the training; conv-stem is adopted in LeViT , EarlyConv , CMT , VOLO and ScaledReLU to improve the robustness of training ViTs; LV-ViT adopts several techniques including MixToken and Token Labeling for better training and feature generation; T2T ViT , DeepViT and CaiT try to train deeper vision transformer models; T2T ViT , ViL and CoaT adopt efficient transformers to deal with the quadratic complexity; To further exploit the capacities of vision transformer, OmniNet , CrossViT and So-ViT propose the dense omnidirectional representations, coarse-fine-grained patch fusion and cross co-variance pooling of visual tokens, respectively. However, all of these works adopt the fully-connected self-attention which will bring the noise or irrelevant tokens for computing and slow down the training of networks. In this paper, we propose an efficient sparse attention, called kk-NN attention, for boosting vision transformers. The proposed kk-NN attention not only inherits the local bias of CNNs but also achieves the ability of global feature exploitation. It can also speed up the training and achieve better performance.

k𝑘k-NN Attention

The intuitive understanding of the attention is the weighted average over the old ones, where the weights are defined by the attention matrix A\bm{A}. In this paper, we consider the Q\bm{Q}, K\bm{K} and V\bm{V} are generated via the linear projection of the input token matrix X\bm{X}:

One shortcoming with fully-connected self-attention is that irrelevant tokens, even though assigned with smaller weights, are still taken into consideration when updating the representation V\bm{V}, making it less resilient to noises in V\bm{V}. This shortcoming motivates us to develop the kk-NN attention.

2 k𝑘k-NN Attention

Instead of computing the attention matrix for all the query-key pairs as in vanilla attention, we select the top-kk most similar keys and values for each query in the kk-NN attention. There are two versions of kk-NN attention, as described below.

Slow Version: For the ii-th query, we first compute the Euclidean distance against all the keys, and then obtain its kk-nearest neighbors Nik\mathcal{N}_{i}^{k} and Niv\mathcal{N}_{i}^{v} from keys and values, and lastly calculate the scaled dot product attention as:

where Tk(⋅)\mathcal{T}_{k}\left(\cdot\right) denotes the row-wise top-k selection operator:

3 Theoretical Analysis on k𝑘k-NN Attention

In this section, we will show theoretically that despite its simplicity, kk-NN attention is powerful in speeding up network training and in distilling noisy tokens. All the proof of the lemmas are provided in the supplementary.

Convergence Speed-up. Compared to CNNs, the fully-connected self-attention is able to capture long range dependency. However, the price to pay is that the dense self-attention model requires to mix each image patch with every other patch in the image, which has potential to mix irrelevant information together, e.g. the foreground patches may be mixed with background patches through the self-attention. This defect could significantly slow down the convergence as the goal of visual object recognition is to identify key visual patches relevant to a given class.

To see this, we consider the model with only learnable parameters WQ\bm{W}_{\bm{Q}}, WK\bm{W}_{\bm{K}} in attention layers and adopting Adam optimizer . According to Theorem 4.1 in , Adam’s convergence is proportional to O(α−1(G∞+1)+αG∞)\mathcal{O}\left(\alpha^{-1}(G_{\infty}+1)+\alpha G_{\infty}\right), where α\alpha is the learning rate and G∞G_{\infty} is an element-wise upper bound on the magnitude of the batch gradient Theorem 4.1 in describes the upper bound for regrets (the gap on loss function value between the current step parameters and optimal parameters). One can telescope it to the average regrets to consider the Adam’s convergence.. Let fif_{i} be the loss function corresponding to batch ii. Via chain rule of derivative, the gradient w.r.t the WQ\bm{W}_{Q} in a self-attention block can be represented as ∇WQfi=Fi(V^knn)⋅∂V^knn∂WQ\nabla_{\bm{W}_{\bm{Q}}}f_{i}=F_{i}(\hat{\bm{V}}^{knn})\cdot\frac{\partial\hat{\bm{V}}^{knn}}{\partial\bm{W}_{\bm{Q}}}, where Fi(V^knn)F_{i}(\hat{\bm{V}}^{knn}) is a matrix output function. Since the possible value of V^knn\hat{\bm{V}}^{knn} is a subset of its fully-connected counterpart, the upper bound of on the magnitude of Fi(V^knn)F_{i}(\hat{\bm{V}}^{knn}) is no larger than the full attention. We then introduce the weighed covariate matrix of patches to characterize the scale of ∂V^knn∂WQ\frac{\partial\hat{\bm{V}}^{knn}}{\partial\bm{W}_{\bm{Q}}} in the following lemma.

(Informal) Let V^lknn\hat{\bm{V}}^{knn}_{l} be the ll-th row of the V^knn\hat{\bm{V}}^{knn}. We have

where Varal(x)\textrm{Var}_{\bm{a}_{l}}(\bm{x}) is the covariate matrix on patches {x1,...,xn}\{\bm{x}_{1},...,\bm{x}_{n}\} with probability from ll-th row of the attention matrix. The same is true for V^\hat{\bm{V}} of the fully-connected self-attention.

Since kk-NN attention only uses patches with large similarity, its Varal(x)\textrm{Var}_{\bm{a}_{l}}(\bm{x}) will be smaller than that computed from the fully-connected attention. As indicated in Lemma 1, ∂V^knn∂WQ\frac{\partial\hat{\bm{V}}^{knn}}{\partial\bm{W}_{\bm{Q}}} is proportional to variance Varal(x)\textrm{Var}_{\bm{a}_{l}}(\bm{x}) and thus the scale of ∇WQfi\nabla_{\bm{W}_{\bm{Q}}}f_{i} becomes smaller in k-NN attention. Similarly, the scale of ∇WKfi\nabla_{\bm{W}_{\bm{K}}}f_{i} is also smaller in k-NN attention. Therefore, the element-wise upper bound on batch gradient G∞G_{\infty} in Adam analysis is also smaller for k-NN attention. For the same learning rate, the k-NN attention yields faster convergence. It is particularly significant at the beginning of training. This is because, due to the random initialization, we expect a relatively small difference in similarities between patches, which essentially makes self-attention behave like “global average". It will take multiple iterations for Adam to turn the "global average" into the real function of self-attention. In Table 2 and Figure 2, we numerically verify the training efficiency of kk-NN attention as opposed to the fully-connected attention.

Noisy patch distillation. As already mentioned before, the fully-connected self-attention model may mix irrelevant patches with relevant ones, particularly at the beginning of training when similarities between relevant patches are not significantly larger than those for irrelevant patches. kk-NN attention is more effective in identifying noisy patches by only considering the top kk most similar patches. To formally justify this point, we consider a simple scenario where all the patches are divided into two groups, the group of relevant patches and the group of noisy patches. All the patches are sampled independently from unknown distributions. We assume that all relevant patches are sampled from distributions with the same shared mean, which is different from the means of distributions for noisy patches. It is important to know that although distributions for the relevant patches share the mean, those relevant patches can look quite differently, due to the large variance in stochastic sampling. In the following Lemma, we will show that the kk-NN attention is more effective in distilling noises for the relevant patches than the fully-connected attention.

We consider the self-attention for query patch ll. Let’s assume the patch xi\bm{x}_{i} are bounded with mean μi\bm{\mu}_{i} for i=1,2,...,ni=1,2,...,n, and ρk\rho_{k} is the ratio of the noisy patches in all selected patches. Under mild conditions, the follow inequality holds with high probability:

In the above lemma, the quantity ∥V^lknn−μlWV∥∞\left\|\hat{\bm{V}}_{l}^{knn}-\bm{\mu}_{l}\bm{W}_{V}\right\|_{\infty} measures the distance between V^lknn\hat{\bm{V}}^{knn}_{l}, represention vector updated by the kk-NN attention, and its mean μlWV\mu_{l}\bm{W}_{\bm{V}}. We now consider two cases: the normal kk-NN attention with appropriately chosen kk, and fully-connected attention with k=nk=n. In the first case, with appropriately chosen kk, we should have most of the selected patches coming from the relevant group, implying a small ρk\rho_{k}. By combining with the fact that kk is decently large, we expect a small upper bound for the distance ∥V^lknn−μlWV∥∞\left\|\hat{\bm{V}}_{l}^{knn}-\bm{\mu}_{l}\bm{W}_{V}\right\|_{\infty}, indicating that kk-NN attention is powerful in distilling noise. For the case of fully-connected attention model, i.e. k=nk=n, it is clearly that ρn≈1\rho_{n}\approx 1, leading to a large distance between transformed representation V^l\hat{\bm{V}}_{l} and its mean, indicating that fully-connected attention model is not very effective in distilling noisy patches, particularly when noise is large.

Besides the instance with low signal-noise-ratio, the instance with a large volume of backgrounds can also be hard. In the next lemma, we show that under a proper choice of kk, with a high probability the kk-NN attention will be able to select all meaningful patches.

Let M∗\mathcal{M}^{*} be the index set contains all patches relevant to query ql\bm{q}_{l}. Under mild conditions, there exist c2∈(0,1)c_{2}\in(0,1) such that with high probability, we have

The above lemma shows that if we select the top O(nd−c2)\mathcal{O}(nd^{-c_{2}}) elements, with high probability, we will be able to eliminate almost all the irrelevant noisy patches, without losing any relevant patches. Numerically, we verify the proper kk gains better performance (e.g., Figure 1) and for the hard instance kk-NN gives more accurate attention regions. (e.g., Figure 4 and Figure 5).

Experiments for Vision Transformers

In this section, we replace the dense attention with kk-NN attention on the existing vision transformers for image classification to verify the effectiveness of the proposed method. The recent DeiT and its variants, including T2T ViT , TNT , PiT , Swin , CvT , So-ViT , Visformer , Twins , Dino and VOLO , are adopted for evaluation. These methods include both supervised methods and self-supervised method . Ablation studies are provided to further analyze the properties of kk-NN attention.

We perform image classification on the standard ILSVRC-2012 ImageNet dataset . In our experiments, we follow the experimental setting of original official released codes. For fair comparison, we only replace the vanilla attention with proposed k-NN attention. Unless otherwise specified, the fast version of kk-NN attention is adopted for evaluation. To speed up the slow version, we develop the CUDA version kk-NN attention. As for the value kk, different architectures are assigned with different values. For DeiT , So-ViT , Dino , CvT , TNT PiT and VOLO , as they directly split an input image into rigid tokens and there is no information exchange in the token generation stage, we suppose the irrelevant tokens are easy to filter, and tend to assign a smaller kk compared with these complicated token generation methods . Specifically, we assign kk to approximate n2\frac{n}{2} at each scale stage; for these complicated token generation methods , we assign a larger kk which is approximately 23n\frac{2}{3}{n} or 45n\frac{4}{5}{n} at each scale stage.

2 Results on ImageNet

Table 1 shows top-11 accuracy results on the ImageNet-1K validation set by replacing the dense attention with kk-NN attention using eleven different vision transformer architectures. From the Table we can see that the proposed kk-NN attention improves the performance from 0.2% to 0.8% for both global and local vision transformers. It is worth noting that on ImageNet-1k dataset, it is very hard to improve the accuracy after 85%, but our kk-NN attention can still consistently improve the performance even without model size increase.

3 The Impact of Number k𝑘k

The only parameter for kk-NN attention is kk, and its impact is analyzed in Figure 1. As shown in the figure, for DeiT-Tiny, kk = 100 is the best, where the total number of tokens nn = 196 (14 ×\times 14), meaning that kk approximates half of nn; for CvT-13, there are three scale stages with the number of tokens n1n_{1} = 3136, n2n_{2} = 784 and n3n_{3} = 196, and the best results are achieved when the kk in each stage is assigned to 1600/400/100, which also approximate half of nn in each stage; for Visformer-Tiny, there are two scale stages with the number of tokens n1n_{1} = 196 and n2n_{2} = 49, and the best results are achieved when kk in each stage is assigned to 150/45, as there are more than 21 conv layers for token generation and the information in each token are already mixed, making it hard to distinguish the irrelevant tokens, thus larger values of kk are desired; for PiT-Base, there are three scale stages with the number of tokens n1n_{1} = 961, n2n_{2} = 256 and n3n_{3} = 64, and the optimal values of kk also approximate the half of nn. Please note that, we do not perform exhaustive search for the optimal choice of kk, instead, a general rule as below is sufficient: k≈k\approx n2\frac{n}{2} at each scale stage for simple token generation methods and k≈23nk\approx\frac{2}{3}{n} or 45n\frac{4}{5}{n} for complicated token generation methods at each scale stage.

4 Convergence Speed of k𝑘k-NN Attention

In Table 2, we investigate the convergence speed of kk-NN attention. Three methods are included for comparison, i.e. DeiT-Small , CvT-13 and T2T-ViT-t-19 . From the Table we can see that the convergence speed of kk-NN attention is faster than full-connected attention, especially in the early stage of training. These observations reflect that removing the irrelevant tokens benefits the convergence of neural networks training.

5 Other properties of k𝑘k-NN attention

To analyze other properties of kk-NN attention, four quantitative metrics are defined as follows. Layer-wise cosine similarity between tokens: following this metric is defined as:

where tit_{i} represents the ii-th token in each layer and ∥⋅∥\lVert\cdot\rVert denotes the Euclidean norm. This metric implies the convergence speed of the network.

Layer-wise standard deviation of attention weights: Given a token tit_{i} and its softmax attention weight sfm(tit_{i}), the standard deviation of the softmax attention weight std(sfm(tit_{i})) is defined as the second metric. For multi-head attention, the standard deviations over all heads are averaged. This metric represents the degree of training stability.

Ratio between the norms of residual activations and main branch: The ratio between the norm of the residual activations and the norm of the activations of the main branch in each layer is defined as ∥fl(t)∥/∥t∥\lVert f_{l}(t)\rVert/\lVert t\rVert, where fl(t)f_{l}(t) can be the attention layer or the FFN layer. This metric denotes the information preservation ability of the network.

Nonlocality: following , the nonlocality is defined by summing, for each query patch ii, the distances ∥δij∥\left\|\delta_{ij}\right\| to all the key patches jj weighted by their attention score Aij\bm{A}_{ij}. The number obtained over the query patch is averaged to obtain the nonlocality metric of head hh, which can the be averaged over the attention heads to obtain the nonlocality of the whole layer ll:

where DlocD_{loc} is the number of patches between the center of attention and the query patch; the further the attention heads look from the query patch, the higher the nonlocality.

Comparisons of the four metrics on DeiT-tiny without distillation token are shown in Figure 2 and Figure 3. From Figure 2 (a) we can see that by using kk-NN attention, the averaged cosine similarity is larger than that of using dense self-attention, which reflects that the convergence speed is faster for kk-NN attention. Figure 2 (b) shows that the averaged standard deviation of kk-NN attention is smoother than that of fully-connected self-attention, and the smoothness will help make the training more stable. Figure 2 (c) and (d) show the ratio between the norms of residual activations and main branch are consistent with each other for kk-NN attention and dense attention, which indicates that there is nearly no information lost in kk-NN attention by removing the irrelevant tokens. Figure 3 shows that, with k-NN attention, lower layers tend to focus more on the local areas (with more lines being pushed toward the bottom area in Figure 3), while the higher layers still maintain their capability of extracting global information. Additionally, it is also observed that the non-locality of different layers is spreading more evenly, indicating that they can explore a larger variety of dependencies at different ranges.

6 Comparisons with temperature in softmax

kk-NN attention effectively zeros the bottom N−kN-k tokens out of the attention calculation. How does this compare with introducing a temperature parameter to softmax over the attention values? We compare our kk-NN attention with temperature tt in softmax as softmax(attn/tt). The performance over the tt is shown in Table 3. From the Table we can see that small tt makes the training crash due to large value of attention values; the performance increases a little bit to 72.5 (baseline 72.2) with tt assigned to appropriate values. The kk-NN attention is more robust compared with temperature in softmax, and achieves much better performance, 73.0 (kk-NN attention) vs 72.5 (best performance for temperature in softmax).

7 Visualization

Figure 4 visualizes the self-attention heads from the last layer on Dino-Small . We can see that different heads attend to different semantic regions of an image. Compared with dense attention, the kk-NN attention filters out most irrelevant information from background regions which are similar to the foreground, and successfully concentrates on the most informative foreground regions. Images from different classes are visualized in Figure 5 using Transformer Attribution method on DeiT-Tiny. It can be seen that the kk-NN attention is more concentrated and accurate, especially in the situations of cluttered background and occlusion.

8 Object Detection and Semantic Segmentation

To verify the effects of kk-NN attention on object detection and semantic segmentation tasks, the widely-used COCO and ADE20K are adopted for evaluation. We adopt Swin-Tiny and Twins-SVT-Base for comparisons due to the well released codes, and the results are shown in Table 4. From the Table we can see that by replacing the vanilla attention with our kk-NN attention, the performance increases with almost no overhead.

Conclusion

In this paper, we propose an effective kk-NN attention for boosting vision transformers. By selecting the most similar keys for each query to calculate the attention, it screens out the most ineffective tokens. The removal of irrelevant tokens speeds up the training. We theoretically prove its properties in speeding up training, distilling noises without losing information, and increasing the performance by choosing a proper kk. Several vision transformers are adopted to verify the effectiveness of the kk-NN attention.

References

Differences with the arXiv paper: Explicit Sparse Transformer: Concentrated Attention Through Explicit Selection (EST)

Similarities: Our method part is similar to EST in terms of the calculation of top-kk. Differences:

Our paper is focused not only on the methodology part, but also the deep understanding. There are many variants of Transformers in the NLP and vision community now, but few of them provide a deep and thorough analysis of their proposed methods. The proposed kk-NN attention indeed happens to be similar to EST, which was arxived 2 years ago and we were not aware of it when conducting our research. In addition to applying the idea to transformers and conducting extensive experiments as EST did, we provide theoretical justifications about the idea, which we think is equally or more important than the method itself, and helps with a more fundamental understanding.

The conclusion about how to select kk is different. In EST, it is found that a small kk is better (8 or 16), but we find a larger kk, namely, ≥12N\geq\frac{1}{2}N is better (NN is the sequence length).

The motivations of these two papers are different: EST targets to get sparse attention maps while ours aims to distill noisy patches.

Our paper focuses on vision transformers but EST focuses on NLP tasks, even though EST applied it to the image captioning task. Since late 2020, vision transformer backbones have become very popular, and kk-NN attention deserves a deeper analysis. Therefore, we apply the kk-NN attention on 11 different vision transformer backbones for empirical evaluations and find it simple and effective for vision transformer backbones.

More analysis about the properties of kk-NN attention in the context of vision transformer backbones are provided in our paper. Besides the kk selection and convergence speed as EST presented, we also define several metrics to facilitate the analysis, e.g. layer-wise cosine similarity between tokens, layer-wise standard deviation of attention weights, ratio between the norms of residual activation and main branch, and nonlocality. We also compare it with temperature in softmaxsoftmax and provide the visualizations.

In summary, our paper provides deeper understanding with comprehensive analysis of the kk-NN attention for vision transformers, which provides well-grounded knowledge advancement.

Source codes of fast version k𝑘k-NN attention in Pytorch

The source codes of fast version kk-NN attention in Pytorch are shown in Algorithm 1, and we can see that the core codes of fast version kk-NN attention is consisted of only four lines, and it can be easily imported to any architecture using fully-connected attention.

Comparisons between slow version and fast version

We develop two versions of kk-NN attention, one slow version and one fast version. The kk-NN attention is exactly defined by slow version, but its speed is extremely slow, as for each query it needs to select different kk keys and values, and this procedure is very slow. To speedup, we developed the CUDA version, but the speed is still slower than fast version. The fast version takes advantages of matrix multiplication and greatly speedup the computing. The speed comparisons on DeiT-Tiny using 8 V100 are illustrated in Table 5.

Evaluations on CIFAR10 or CIFAR100.

As vision transformers are data-hungry, directly training vision transformer backbones from scratch on small-size datasets such as CIFAR10 or CIFAR100 would yield much worse performances compared with ConvNets. Following the paradigm and codes of the NIPS2021 paper “Efficient Training of Visual Transformers with Small Datasets", we briefly conducted experiments on CIFAR10 and CIFAR100 using Swin-T and T2T-ViT-14 with kk-NN attention as shown in Table 6. Adding kk-NN attention brings much larger performance gain in the scratch training (ST) due to its faster convergence speed, while the gain in the setting of ImageNet-1k pretraining and CIFAR finetuning (FT) is not as large.

Proof

Notations. Throughout this appendix, we denote xix_{i} as ii-th element of vector x\bm{x}, Wij\bm{W}_{ij} as the element at ii-th row and jj-th column of matrix W\bm{W}, and Wj\bm{W}_{j} as the jj-th row of matrix W\bm{W}. Moreover, we denote xi\bm{x}_{i} as the ii-th patch (token) of the inputs with xi=Xi\bm{x}_{i}=\bm{X}_{i}.

Proof for Lemma 1 We first give the formal statement of Lemma 1.

The same is true for V^\hat{\bm{V}} of the fully-connected self-attention.

Let’s first consider the derivative of V^l\hat{\bm{V}}_{l} over WQ,ijW_{\bm{Q},ij}. Via some algebraic computation, we have

where we denote Tlknn(k)\mathcal{T}_{l}^{knn}(k) as follow for shorthand:

Let denote set \mathcal{S}\doteq\{i:\textrm{patchiisselectedinrowis selected in rowl}\} and then we consider the right-hand-side of (1).

Since al\bm{a}_{l} is the ll-th row of the attention matrix, we have alt≥0a_{lt}\geq 0 and ∑talt=1\sum_{t}a_{lt}=1. It is possible to treat terms (a)(a), (b)(b) and (c)(c) as the expectation of some quantities over tt replicates with probability alta_{lt}. Then (2) can be further simplified as

where the second equality uses the fact that xtWK,j\bm{x}_{t}\bm{W}_{\bm{K},j} is a scalar.

Due the symmetric on Q\bm{Q} and K\bm{K}, we can follow the similar procedure to show

Finally, by setting k=nk=n, one may verify that equations (4) and (5) also hold for fully-connected self-attention.

Proof for Lemma 2 Before given the formal statement of the Lemma 2, we first show the assumptions.

The token xi\bm{x}_{i} is the sub-gaussian random vector with mean μi\bm{\mu}_{i} and variance (σ2/d)I(\sigma^{2}/d)I for i=1,2,...,ni=1,2,...,n.

μ\bm{\mu} follows a discrete distribution with finite values μ∈V\bm{\mu}\in\mathcal{V}. Moreover, there exist 0<ν1,0<ν2<ν40<\nu_{1},0<\nu_{2}<\nu_{4} such that a) ∥μi∥=ν1\|\bm{\mu}_{i}\|=\nu_{1}, and b) μiWQWKTμi∈[ν2,ν4]\bm{\mu}_{i}\bm{W}_{\bm{Q}}\bm{W}_{\bm{K}}^{T}\bm{\mu}_{i}\in[\nu_{2},\nu_{4}] for all ii and ∣μiWQWK⊤μj⊤∣≤ν2|\bm{\mu}_{i}\bm{W}_{\bm{Q}}\bm{W}_{\bm{K}}^{\top}\bm{\mu}_{j}^{\top}|\leq\nu_{2} for all μi≠μj∈V\bm{\mu}_{i}\neq\bm{\mu}_{j}\in\mathcal{V}.

WV\bm{W}_{V} and WQWK⊤\bm{W}_{\bm{Q}}\bm{W}_{\bm{K}}^{\top} are element-wise bounded with ν5\nu_{5} and ν6\nu_{6} respectively, that is, ∣WV(ij)∣≤ν5|\bm{W}_{V}^{(ij)}|\leq\nu_{5} and ∣(WQWK⊤)(ij)∣≤ν6|(\bm{W}_{\bm{Q}}\bm{W}_{\bm{K}}^{\top})^{(ij)}|\leq\nu_{6}, for all i,ji,j from 1 to dd.

In Assumption 2 we ensure that for a given query patch, the difference between the clustering center and noises are large enough to be distinguished.

Let patch xi\bm{x}_{i} be σ2\sigma^{2}-subgaussian random variable with mean μi\bm{\mu}_{i} and there are k1k_{1} patches out of all kk patches follows the same clustering center of query ll. Per Assumption 2, when d≥3(ψ(δ,d)+ν2+ν4)\sqrt{d}\geq 3(\psi(\delta,d)+\nu_{2}+\nu_{4}), then with probability 1−5δ1-5\delta, we have

where ψ(δ,d)=2σν1ν62log⁡(1δ)+2σ2ν6log⁡(dδ)\psi(\delta,d)=2\sigma\nu_{1}\nu_{6}\sqrt{2\log\left(\frac{1}{\delta}\right)}+2\sigma^{2}\nu_{6}\log\left(\frac{d}{\delta}\right).

Without loss of generality, we assume the first kk patch are the top-kk selected patches. From Assumption 2.1, we can decompose xi=μi+hi\bm{x}_{i}=\bm{\mu}_{i}+\bm{h}_{i}, i=1,2,...,ki=1,2,...,k, where hi\bm{h}_{i} is the sub-gaussian random vector with zero mean. We then analyze the numerator part.

Below we will bound (a)(a), (b)(b) and (c)(c) separately.

Upper bound for (a)(a). Let denote index set S1={i:μ1=μi, i=1,2,...,k}\mathcal{S}_{1}=\{i:\bm{\mu}_{1}=\bm{\mu}_{i},\ i=1,2,...,k\}. We then have

where last inequality is from the Assumption 2.2.

Upper bound for (b)(b). Since each dimension in hl\bm{h}_{l} is the i.i.d random vector with zero mean variance σ2/d\sigma^{2}/d based on Assumption 2.1, we can use Hoeffding Inequality to derive the following result holds with probability 1−δ1-\delta.

We then build the upper bound for U1U_{1} and U2U_{2}. Since xi=μi+hi\bm{x}_{i}=\bm{\mu}_{i}+\bm{h}_{i} for i=1,2,...,ki=1,2,...,k, we have

Via Assumption 2.3 and Hoeffding Inequality, with probability 1−4δ1-4\delta, the follow results hold.

and we denote ψ(δ,d)=2σν1ν62log⁡(1δ)+2σ2ν6log⁡(dδ)\psi(\delta,d)=2\sigma\nu_{1}\nu_{6}\sqrt{2\log\left(\frac{1}{\delta}\right)}+2\sigma^{2}\nu_{6}\log\left(\frac{d}{\delta}\right) for shorthand and then we have

As a result, with a probability 1−5δ1-5\delta, we have:

Combine (15) with (11)-(13) and we have with probability 1−4δ1-4\delta:

From (7), (8), (14) and (16), with probability 1−5δ1-5\delta, we have

where τ1(k,k1)=(k−k1)exp⁡(ν2d)+k1exp⁡(ν5d)\tau_{1}(k,k_{1})=(k-k_{1})\exp\left(\frac{\nu_{2}}{\sqrt{d}}\right)+k_{1}\exp\left(\frac{\nu_{5}}{\sqrt{d}}\right) and τ2(k,k1)=k1exp⁡(2ν2d)+(k−k1)exp⁡(2ν4d)\tau_{2}(k,k_{1})=k_{1}\exp\left(\frac{2\nu_{2}}{\sqrt{d}}\right)+(k-k_{1})\exp\left(\frac{2\nu_{4}}{\sqrt{d}}\right).

Now we consider the upper bound the denominator part.

Via Assumption 2.2, (11)-(13) and the definition of ψ(δ,d)\psi(\delta,d), with probability 1−5δ1-5\delta the follow results hold.

Combining (18), (19) and Assumption 2.2, we have

When d≥3(ψ(δ,d)+ν2+ν4)\sqrt{d}\geq 3(\psi(\delta,d)+\nu_{2}+\nu_{4}), one may verify

where second inequality uses ν4≥μ2\nu_{4}\geq\mu_{2} and last inequality uses the definition of τ1(k,k1)\tau_{1}(k,k_{1}), τ2(k,k1)\tau_{2}(k,k_{1}) and d≥3(ψ(δ,d)+ν2+ν4)\sqrt{d}\geq 3(\psi(\delta,d)+\nu_{2}+\nu_{4}).

To see the connection between the Assumption 3.1 with attention scheme, we consider the follow problem.

If K\bm{K} is normalized with zero columns mean and we apply the exponential gradient method on the initial solution β0=1ne\bm{\beta}_{0}=\frac{1}{n}\bm{e} with step length 1/d1/\sqrt{d}, the one step updated solution β1\bm{\beta}_{1} is

The above equation (25) is just the attention scheme used standard transformer type model. Based on Assumption 3.1, we can treat β1\bm{\beta}_{1} as an approximation of underlying true parameters β∗\bm{\beta}^{*}.

On the other hand, it is commonly believed that only part of patches are correlated with the query patch (i.e., with non-zero similarity weights.) and it would be ideal if we could use a computational cheap method to eliminate the irrelevant patches. In this paper, we consider the top-k selection scheme. To see the rationality of the top-kk selection, we consider augmenting (24) with L2L_{2} regularization on β\bm{\beta}.

where we use the KKT optimal condition, and λ1\bm{\lambda}_{1}, λ2\lambda_{2} are Lagrange multipliers to make sure β∈Δ\bm{\beta}\in\Delta. If λ\lambda is large enough and β>0\bm{\beta}>0, we will have

The above result indicates that we may selection the important elements in β\bm{\beta} (e.t., with large magnitude) by considering rankness in vector Kq1⊤\bm{K}\bm{q}_{1}^{\top}.

We then discussion the correctness of the top-k selection with the following regularity assumptions on K\bm{K} and q1\bm{q}_{1}.

K\bm{K} is normalized with row zero mean. Let Σ=KK⊤\bm{\Sigma}=\bm{K}\bm{K}^{\top} and Z=Σ−1/2K⊤\bm{Z}=\bm{\Sigma^{-1/2}}\bm{K}^{\top}. We assume there exist some c,c4>1c,c_{4}>1 and C1>0C_{1}>0 such that the following inequality

var(q1)=O(1)\textrm{var}(\bm{q}_{1})=\mathcal{O}(1) and for some κ≥0\kappa\geq 0 and c5,c6>0c_{5},c_{6}>0,

There exist some τ∈[0,1)\tau\in[0,1) and c7>0c_{7}>0 such that

Let’s assume only be ss keywords are relevant to the query ll. Under Assumption 3.1 and 3.2, when 2κ+τ<12\kappa+\tau<1, with probability 1−O(sexp⁡(−Cd1−2κ/log⁡d))1-\mathcal{O}(s\exp(-Cd^{1-2\kappa}/\log d)), we have

where \mathcal{M}^{*}=\{i:\textrm{keywordiisrelevanttothequeryis relevant to the queryl.}\} , and τ\tau, κ\kappa, cc and CC are positive constants.

Our strategy is the similar to the proof of Theorem 1 in .

We then separately study ξ\bm{\xi} and η\bm{\eta}.

Analysis on ξ\bm{\xi}. We first bound ξ\bm{\xi} from above. Since {μi}\{\mu_{i}\} is the eigenvalues of n−1ZZ⊤n^{-1}\bm{Z}\bm{Z}^{\top}, we have

Let QQ belongs to the orthogonal group O(n)O(n) such that Σ1/2β=∥Σ1/2β∥Qe1\bm{\Sigma}^{1/2}\bm{\beta}=\|\bm{\Sigma}^{1/2}\bm{\beta}\|Q\bm{e}_{1}. Then, it follows from Lemma 1 in that

where we use the symbol =(d)\overset{(d)}{=} to denote being identical in distribution. By part 3 in Assumption 3.2, ∥Σ1/2β∥2=β⊤Σβ≤var(y)=O(1)\|\bm{\Sigma}^{1/2}\bm{\beta}\|^{2}=\bm{\beta}^{\top}\bm{\Sigma}\bm{\beta}\leq\textrm{var}(\bm{y})=\mathcal{O}(1), and thus via Lemma 4 in , we have for some C>0C>0,

We then consider the lower bound on ξi\xi_{i} for i∈M∗i\in\mathcal{M}_{*}. By (28), we have

Note that ∥Σ1/2ei∥=var(Xi)=1\|\bm{\Sigma}^{1/2}\bm{e}_{i}\|=\sqrt{\textrm{var}(\bm{X}_{i})}=1, ∥Σ1/2β∥=O(1)\|\bm{\Sigma}^{1/2}\bm{\beta}\|=\mathcal{O}(1). By part 2 of Assumption 3.2, there exists some c>0c>0 such that

Thus, there exists QQ in orthogonal group O(n)O(n) such that Σ1/2ei=Qe1\bm{\Sigma}^{1/2}\bm{e}_{i}=Q\bm{e}_{1} and

and thus by part 1 of assumption 3.2, Lemma 4 in , and union bound, we have for some c,C>0c,C>0,

where WW is N(0,1)\mathcal{N}(0,1)-distributed random variable. We then pick xd=c2Cd1−κ/log⁡dx_{d}=c\sqrt{2C}d^{1-\kappa}/\sqrt{\log d}. Then, by standard tail bound, we have

It implies that for ii with βi∗>0\beta_{i}^{*}>0, we have

Next, we examine term η=(η1,...,ηn)⊤=Kϵ⊤\bm{\eta}=(\eta_{1},...,\eta_{n})^{\top}=\bm{K}\bm{\epsilon}^{\top}. Clearly, we have

From Assumption 3.1, we know that {ϵi2/σ2}\{\epsilon_{i}^{2}/\sigma^{2}\} are i.i.d. χ12\chi_{1}^{2}-distributed random variables. Thus there exist c,C>0c,C>0 such that

Along with parts 1 and 3 of Assumption 3.2, we have

We then bound ∣ηi∣|\eta_{i}| from above. Given that η=Kϵ⊤∼N(0,σ2KK⊤)\bm{\eta}=K\bm{\epsilon}^{\top}\sim\mathcal{N}(0,\sigma^{2}KK^{\top}). Hence ηi∣K=K∼N(0,var(ηi∣K=K))\eta_{i}|_{K=\bm{K}}\sim\mathcal{N}(0,\textrm{var}(\eta_{i}|_{K=\bm{K}})) with var(ηi∣K=K)=σ2ei⊤Kei\textrm{var}(\eta_{i}|K=\bm{K})=\sigma^{2}\bm{e}_{i}^{\top}\bm{K}\bm{e}_{i}.

Let E\mathcal{E} be the event {var(ηi∣K)≤cd}\{\textrm{var}(\eta_{i}|\bm{K})\leq cd\} for some c>0. Then, using the same argument in the previous proof. we can show that

Condition on the event E\mathcal{E}, for all x>0x>0, we have

, where WW is N(0,1)\mathcal{N}(0,1) random variable. Via union bound, we have

By setting x=2cCd1−κ/log⁡dx=\sqrt{2cC}d^{1-\kappa}/\sqrt{\log d}, we have

Therefore, with probability 1−O(sexp⁡(−Cd1−2κ/log⁡d))1-\mathcal{O}(s\exp(-Cd^{1-2\kappa}/\log d)), the magnitudes of ωi\omega_{i} with βi∗>0\beta_{i}^{*}>0 are uniformly at least of order d1−κd^{1-\kappa} and for some c>0c>0, we have