Separable Self-attention for Mobile Vision Transformers

Sachin Mehta, Mohammad Rastegari

Introduction

Vision transformers (ViTs) have become ubiquitous for a wide variety of visual recognition tasks , including mobile vision tasks . At the heart of the ViT-based models, including mobile vision transformers, is the transformer block . The main efficiency bottleneck in ViT-based models, especially for inference on resource-constrained devices, is the multi-headed self-attention (MHA). MHA allows the tokens (or patches) to interact with each other, and is a key for learning global representations. However, the complexity of self-attention in transformer block is O(k2)O(k^{2}), i.e., it is quadratic with respect to the number of tokens (or patches) kk. Besides this, computationally expensive operations (e.g., batch-wise matrix multiplication; see Fig. 1) are required to compute attention matrix in MHA. This, in particular, is concerning for deploying ViT-based models on resource-constrained devices, as these devices have reduced computational capabilities, restrictive memory constraints, and a limited power budget. Therefore, this paper seeks to answer this question: can self-attention in transformer block be optimized for resource-constrained devices?

Several methods [e.g., 7, 8, 9, 10] have been proposed for optimizing the self-attention operation in transformers (not necessarily for ViTs). Among these, a widely studied approach in sequence modeling tasks is to introduce sparsity in self-attention layers, wherein each token attends to a subset of tokens in an input sequence . Though these approaches reduces the time complexity from O(k2)O(k^{2}) to O(kk)O(k\sqrt{k}) or O(klog⁡k)O(k\log{k}), the cost is a performance drop. Another popular approach for approximating self-attention is via low-rank approximation. Linformer decomposes the self-attention operation into multiple smaller self-attention operations via linear projections, and reduces the complexity of self-attention from O(k2)O(k^{2}) to O(k)O(k). However, Linformer still uses costly operations (e.g., batch-wise matrix multiplication; Fig. 1) for learning global representations in MHA, which may hinder the deployment of these models on resource-constrained devices.

This paper introduces a novel method, separable self-attention, with O(k)O(k) complexity for addressing the bottlenecks in MHA in transformers. For efficient inference, the proposed self-attention method also replaces the computationally expensive operations (e.g., batch-wise matrix multiplication) in MHA with element-wise operations (e.g., summation and multiplication). Experimental results on standard vision datasets and tasks demonstrates the effectiveness of the proposed method (Fig. 2).

Related work

Improving the efficiency of MHA in transformers is an active area of research. The first line of research introduces locality to address the computational bottleneck in MHA [e.g., 7, 9, 11, 12]. Instead of attending to all kk tokens, these methods use predefined patterns to limit the receptive field of self-attention from all kk tokens to a subset of tokens, reducing the time complexity from O(k2)O(k^{2}) to O(kk)O(k\sqrt{k}) or O(klog⁡k)O(k\log{k}). However, such methods suffer from large performance degradation with moderate training/inference speed-up over the standard MHA in transformers. To improve the efficiency of MHA, the second line of research uses similarity measures to group tokens . For instance, Reformer uses locality-sensitive hashing to group the tokens and reduces the theoretical self-attention cost from O(k2)O(k^{2}) to O(klog⁡k)O(k\log{k}). However, the efficiency gains over standard MHA are noticeable only for large sequences (k>2048k>2048) . Because k<1024k<1024 in ViTs, these approaches are not suitable for ViTs. The third line of research improves the efficiency of MHA via low-rank approximation . The main idea is to approximate the self-attention matrix with a low-rank matrix, reducing the computational cost from O(k2)O(k^{2}) to O(k)O(k). Even though these methods speed-up the self-attention operation significantly, they still use expensive operations for computing attention, which may hinder the deployment of these models on resource-constrained devices (Fig. 1).

In summary, existing methods for improving MHA are limited in their reduction of inference time and memory consumption, especially for resource-constrained devices. This work introduces a separable self-attention method that is fast and memory-efficient (see Fig. 1), which is desirable for resource-constrained devices.

Improving transformer-based models

There has been significant work on improving the efficiency of transformers . The majority of these approaches reduce the number of tokens in the transformer block using different methods, including down-sampling and pyramidal structure . Because the proposed separable self-attention module is a drop-in replacement to MHA, it can be easily integrated with any transformer-based model to further improve its efficiency.

Other methods

Transformer-based models performance can be improved using different methods, including mixed-precision training , efficient optimizers , and knowledge distillation . These methods are orthogonal to our work, and by default, we use mixed-precision during training.

MobileViTv2

MobileViT is a hybrid network that combines the strengths of CNNs and ViTs. MobileViT views transformers as convolutions, which allows it to leverage the merits of both convolutions (e.g., inductive biases) and transformers (e.g., long-range dependencies) to build a light-weight network for mobile devices. Though MobileViT networks have significantly fewer parameters and deliver better performance as compared to light-weight CNNs (e.g., MobileNets ), they have high latency. The main efficiency bottleneck in MobileViT is the multi-headed self-attention (MHA; Fig. 3(a)).

MHA uses scaled dot-product attention to capture the contextual relationships between kk tokens (or patches). However, MHA is expensive as it has O(k2)O(k^{2}) time complexity. This quadratic cost is a bottleneck for transformers with a large number of tokens kk (Fig. 1). Moreover, MHA uses computationally- and memory-intensive operations (e.g., batch-wise matrix multiplication and softmax for computing attention matrix; Fig. 1); which could be a bottleneck on resource-constrained devices. To address the limitations of MHA for efficient inference on resource-constrained devices, this paper introduces separable self-attention with linear complexity (Fig. 3(c)).

The main idea of our separable self-attention approach, shown in Fig. 4(b), is to compute context scores with respect to a latent token LL. These scores are then used to re-weight the input tokens and produce a context vector, which encodes the global information. Because the self-attention is computed with respect to a latent token, the proposed method can reduce the complexity of self-attention in the transformer by a factor kk. A simple yet effective characteristic of the proposed method is that it uses element-wise operations (e.g., summation and multiplication) for its implementation, making it a good choice for resource-constrained devices. We call the proposed attention method separable self-attention because it allows us to encode global information by replacing the quadratic MHA with two separate linear computations. The improved model, MobileViTv2, is obtained by replacing MHA with separable self-attention in MobileViT.

In the rest of this section, we first briefly describe MHA (Section 3.1), and then elaborate on the details of separable self-attention (Section 3.2) and MobileViTv2 architecture (Section 3.3).

2 Separable self-attention

The context vector cv\mathbf{c_{v}} is analogous to the attention matrix a\mathbf{a} in Eq. 1 in a sense that it also encodes the information from all tokens in the input x\mathbf{x}, but is cheap to compute.

where ∗* and ∑\sum are broadcastable element-wise multiplication and summation operations, respectively.

Fig. 1 compares the proposed method with Transformer and Linformer. Because time complexity of self-attention methods do not account for the cost of operations that are used to implement these methods, some of the operations may become bottleneck on resource-constrained devices. For holistic understanding, module-level latency on a single CPU core with varying kk is also measured in addition to theoretical metrics. The proposed separable self-attention is fast and efficient as compared to MHA in Transformer and Linformer.

Besides these module-level results, when we replaced the MHA in the transformer with the proposed self-separable attention in the MobileViT architecture, we observe 3×3\times improvement in inference speed with similar performance on the ImageNet-1k dataset (Table 1). These results show the efficacy of the proposed separable self-attention at the architecture-level. Note that self-attention in Transformer and Linformer yields similar results for MobileViT. This is because the number of tokens kk in MobileViT is fewer (k≤1024k\leq 1024) as compared to language models, where Linformer is significantly faster than the transformer.

Relationship with additive addition

The proposed approach resembles the attention mechanism of Bahdanau et al. , which also encodes the global information by taking a weighted-sum of LSTM outputs at each time step. Unlike , where input tokens interact via recurrence, the input tokens in the proposed method interact only with a latent token.

3 MobileViTv2 architecture

To demonstrate the effectiveness of the proposed separable self-attention on resource-constrained devices, we integrate separable self-attention with a recent ViT-based model, MobileViT . MobileViT is a light-weight, mobile-friendly hybrid network that delivers significantly better performance than other competitive CNN-based, transformer-based, or hybrid models, including MobileNets . To avoid ambiguity, we refer to MobileViT as MobileViTv1 in the rest of the paper.

Specifically, we replace MHA in the transformer block in the MobileViTv1 with the proposed separable self-attention method. We call the resultant architecture MobileViTv2. We also do not use the skip-connection and fusion block in the MobileViT block (Fig. 1b in ) as it improves the performance marginally (Fig. 12 in ). Furthermore, to create MobileViTv2 models at different complexities, we uniformly scale the width of MobileViTv2 network using a width multiplier α∈{0.5,2.0}\alpha\in\{0.5,2.0\}. This is in contrast to MobileViTv1 which trains three specific architectures (XXS, XS, and S) for mobile devices. More details about MobileViTv2’s architecture are given in Appendix A.

Experimental results

We train MobileViTv2 for 300 epochs with an effective batch size of 1024 images (128 images per GPU ×\times 8 GPUs) using AdamW on the ImageNet-1k dataset with 1.28 million and 50 thousand training and validation images respectively. We linearly increase the learning rate from 10−610^{-6} to 0.0020.002 for the first 20k iterations. After that, the learning rate is decayed using a cosine annealing policy . To reduce stochastic noise during training, we use exponential moving average (EMA) as we find it helps larger models. We implement our models using CVNets , and use their provided scripts for data processing, training, and evaluation.

Pre-training on ImageNet-21k-P and finetuning on ImageNet-1k

We train on the ImageNet-21k (winter’21 release) that contains about 13 million images across 19k classes. Specifically, we follow to pre-process (e.g., remove classes with fewer samples) the dataset and split it into about 11 million and 522 thousand training and validation images spanning over 10,450 classes, respectively. Following , we refer to this pre-processed dataset as ImageNet-21k-P. Note that the ImageNet-21k-P validation set does not overlap with the validation and test sets of ImageNet-1k.

We follow for pre-training MobileViTv2 on ImageNet-21k-P. For faster convergence, we initialize MobileViTv2 models with ImageNet-1k weights and finetune it on ImageNet-21k-P for 80 epochs with an effective batch size of 4096 images (128 images per GPU x 32 GPUs). We do not use any linear warm-up. Other settings follow ImageNet-1k training.

We finetune ImageNet-21k-P pre-trained models on ImageNet-1k for 50 epochs using SGD with momentum (0.9) and cosine annealing policy with an effective batch size of 256 images (128 images per GPU ×\times 2 GPUs).

Finetuning at higher resolution

MobileViTv2 is a hybrid architecture that combines convolution and separable self-attention to learn visual representations. Unlike many ViT-based models (e.g., DeiT), MobileViTv2 does not require adjustment to patch embeddings or positional biases for different input resolutions and is simple to finetune. We finetune MobileViTv2 models at higher resolution (i.e., 384<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><mo>×</mo></mrow><annotationencoding="application/x−tex">×</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.6667em;vertical−align:−0.0833em;"></span><spanclass="mord">×</span></span></span></span></span>384384<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mo>×</mo></mrow><annotation encoding="application/x-tex">\times</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6667em;vertical-align:-0.0833em;"></span><span class="mord">×</span></span></span></span></span>384) for 10 epochs with a fixed learning rate of 10−310^{-3} using SGD.

Comparison with existing methods

Table 2 and Fig. 2 compares MobileViTv2’s performance with recent methodsFor additional results including ablations, see Appendix B, Appendix C, and Appendix E.. We make following observations:

When MHA in MobileViTv1 is replaced with separable self-attention, the resultant model, MobileViTv2, is faster and better (Fig. 2); validating the effectiveness of the proposed separable self-attention method for mobile ViTs.

Compared to transformer-based (including hybrid) models, MobileViTv2 models are fast on mobile devices. For example, MobileViTv2 is about 8×8\times faster on a mobile device and delivers 2.5% better performance on the ImageNet-1k dataset than MobileFormer , even though MobileFormer is FLOP efficient (R9 vs. R10). However, on GPU, both MobileFormer and MobileViTv2 run at a similar speed. The discrepancy in FLOPs and speed of MobileFormer across devices is primarily because of its architectural design. MobileFormer has conditional operations between mobile and former blocks. Such conditional operations, especially on resource-constrained devices, have a low degree of parallelism and create memory bottlenecks, resulting in a high latency network. Ma et al. also makes a similar observation for CNN-based architectures.

MobileViTv2 bridges the latency gap between CNN- and ViT-based models on mobile devices while maintaining performance with similar or fewer parameters. For example, on a mobile device, ConvNexT (CNN-based model) is 2×2\times and 3.6×3.6\times faster than MobileViTv2 (hybrid model) and DeiT (transformer-based model) for similar performance respectively (see R11, R13, and R14). The low latency of fully CNN-based models on mobile devices can be attributed to several device-level optimizations that have been done for CNN-based models over the past few years (e.g., dedicated hardware implementations for convolutions and folding batch normalization with convolutions). ViT-based models still lack such optimizations and therefore, the resultant inference graphs are sub-optimal. Though MobileViTv2 bridges the latency gap between CNNs and ViTs, we believe the latency of ViT-based models will improve in the future with similar optimizations.

The delta in speed (on GPU) between ConvNext and MobileViTv2 (R15-R18) at higher model complexities reduces from 1.6×1.6\times to 1.3×1.3\times when input resolution is increased from 224×224224\times 224 (or 256×256256\times 256) to 384×384384\times 384, suggesting ViT-based (including hybrid) models exhibit better scaling properties as compared to CNNs. This is because of a higher degree of parallelism that ViT-based models offer at a large scale . Our results on down-stream tasks in Section 4.2 and previous work on scaling ViTs further supports this observation.

2 Evaluation on down-stream tasks

We integrate MobileViTv2 with two standard segmentation architectures, PSPNet and DeepLabv3 , and study it on two standard semantic segmentation datasets, ADE20k and PASCAL VOC 2012 . For training details including hyper-parameters, see supplementary material.

Table 3 and Fig. 2(c) compares the segmentation performance in terms of validation mean intersection over union (mIOU) of MobileViTv2 with different segmentation methods. MobileViTv2 delivers competitive performance at different complexities while having significantly fewer parameters and FLOPs. Interestingly, the inference speed of MobileViTv2 models is comparable to CNN-based models, including light-weight MobileNetv2 and heavy-weight ResNet-50 model. This is consistent with our observation in Section 4.1 (R17 vs. R18; Table 2) where we also observe that ViT-based models scale better than CNN’s at higher input resolutions and model complexities.

Object detection

We integrate MobileViTv2 with SSDLite (SSD head with separable convolutions) for mobile object detection, and study its performance on MS-COCO dataset . We follow for training detection models. Table 4 and Fig. 2(b) compares SSDLite’s detection performance in terms of validation mean average precision (mAP) using different ImageNet-1k backbones. MobileViTv2 delivers competitive performance to models with different capacities, further validating the effectiveness of the proposed self-separable attention method.

Visualizations of self-separable attention scores

Fig. 5 visualizes what the context scores learn at different output stridesOutput stride is the ratio of the spatial dimension of the input to the feature map. of MobileViTv2 network. We found that separable self-attention layers pay attention to low-, mid-, and high-level features, and allow MobileViTv2 to learn representations from semantically relevant image regions.

Conclusions

Transformer-based vision models are slow on mobile devices as compared to CNN-based models because multi-headed self-attention is expensive on resource-constrained devices. In this paper, we introduce a separable self-attention method that has linear complexity and can be implemented using hardware-friendly element-wise operations. Experimental results on standard datasets and tasks demonstrate the effectiveness of the proposed method over multi-headed self-attention.

Acknowledgements

We are grateful to Ali Farhadi, Peter Zatloukal, Oncel Tuzel, Rick Chang, Fartash Faghri, Farzad Abdolhosseini, Lailin Chen, and Max Horton for their helpful comments. We are also thankful to Apple’s infrastructure and open-source teams for their help with training infrastructure and open-source release of the code and pre-trained models.

References

Appendix A Detailed architecture of MobileViTv2

MobileViTv2’s architecture follows MobileViTv1 and is given in Table 5. MobileViTv2 block, shown in Fig. 6, makes two changes to the MobileViTv1 block: (1) it replaces the multi-headed self-attention with the proposed separable self-attention to learn global representations and (2) it does not use fusion block and skip-connection (see Fig. 1b in ) as they improve the performance marginally (see Fig. 12 in ). The expansion factor in MobileNetv2 blocks and feed-forward layers is two. Similar to , we use Swish as a non-linear activation function. Unlike MobileViTv1 that creates three specific architectures (XXS, XS, and S) for mobile devices, we uniformly scale the width of MobileViTv2 network using a width multiplier α∈0.5,2.0\alpha\in{0.5,2.0} to create models at different complexities.

Appendix B MobileViTv2’s classification performance

Table 6 shows the results of MobileViTv2 on the ImageNet-1k dataset. Finetuning MobileViTv2 models at higher resolution (384×384384\times 384) shows improvement across the board. For example, the performance of MobileViTv2-0.50 with 1.4 million parameters improves by about 2% when finetuned at higher resolution (R1 vs. R2). Similarly, pre-training on the ImageNet-21k-P dataset helps improve the performance of MobileViTv2 models. For example, ImageNet-21k-P pretraining improves the performance of MobileViTv2-2.0 improves by 1.2% (R17 vs. R18). Notably, MobileViTv2 models pretrained on the ImageNet-21k-P are able to achieve the similar performance with fewer FLOPs to models finetuned on ImageNet-1k with a higher resolution (e.g., R10 vs. R11; R14 vs. R15; R18 vs. R19 in Table 6).

ImageNet-21k-P

Table 7 shows the results on the ImageNet-21k-P validation dataset. The performance of MobileViTv2 improves with increase in model size.

Appendix C Comparisons with light-weight networks on the ImageNet-1k dataset

Comparison with light-weight CNNs. Fig. 7(a) shows that MobileViTv2 outperforms light-weight CNNs across different network sizes (MobileNetv1 , MobileNetv2 , ShuffleNetv2 , ESPNetv2 , and MobileNetv3 ).

Fig. 7(b) shows that MobileViTv2 achieves better performance than previous light-weight ViT-based models acorss different network sizes (DeIT , T2T , CrossViT , LocalViT , ConViT , and Mobile-former ).

Appendix D Visualizations of separable self-attention scores

The context score maps for different input images at different output strides of MobileViTv2 model are shown in Fig. 8. These visualizations show that the proposed separable self-attention method is able to (1) aggregate information from entire image under different settings, including complex backgrounds, illumination & view-point changes, and different objects, and (2) learn high-, mid-, and low-level representations.

Appendix E MobileViTv2’s ablation studies on the ImageNet-1k dataset

In this section, we study the effect on different methods on the performance of MobileViTv2 models, including augmentation methods.

We study two different augmentation methods: (1) standard augmentation that uses Inception-style augmentation , i.e., random resized cropping and horizontal flipping and (2) advanced augmentation that uses RandAugment , CutMix , MixUp , and RandomErase along with standard augmentation methods. The effect of these augmentations on the performance of MobileViTv2 is shown in Figure 9. Smaller models (<4.5<4.5 million parameters) benefit from standard augmentation while larger models (≥4.5\geq 4.5 million parameters) benefit from advanced augmentation. For simplicity, we use advanced augmentation for all variants of MobileViTv2 in this paper.

Loss functions

CutMix and Mixup augmentations mixes the samples in a batch. As a result, each sample has multiple labels. Therefore, in presence of these augmentations, ImageNet classification can be thought as a multi-label classification task. Similar to , we trained MobileViTv2 by minimizing binary cross-entropy loss. Unlike , we did not observe any improvements in the performance when cross-entropy loss with label smoothing is replaced with binary cross-entropy loss. Therefore, we use cross-entropy with label smoothing for training MobileViTv2 models.

Effect of multiple latent tokens

Similar to multi-head attention in transformers, the proposed separable self-attention can have multiple latent tokens. When we changed the number of latent tokens from 11 to 88, the performance improvements on the ImageNet-1k dataset were negligible (within ±0.1\pm 0.1 top-1 accuracy). Therefore, we use only one latent token in our experiments.

We note that changing the number of heads from 44 to 11 in multi-headed self-attention in the transformer block of the MobileViTv1-S architecture dropped the top-1 accuracy by 0.7%. This observation is similar to Vaswani et al. , who also found that multiple heads in multi-headed self-attention improve transformers performance on the task of neural machine translation.

Improving FLOP-efficiency via pixel- and patch-sampling

The MobileViTv1 model unfolds an input feature map into NN patches, each patch with M=hwM=hw pixels and applies a transformer block for each pixel in a patch independently, where hh and ww are patch’s height and width respectively. Because pixels in a patch are spatially correlated, one can sub-sample mm pixels from MM pixels and learn non-local representations by applying self-attention layers on mm pixels only. Such sub-sampling methods should help in reducing model FLOPs.

We tried following sampling methods at pixel- as well as patch-level:

Random sampling, wherein mm pixels (or nn patches) from MM pixels (or NN patches) are randomly selected during training and uniformly during validation.

Top-mm (or top-nn) sampling, wherein top-mm pixels (or top-nn patches) are selected based on their magnitude computed using L2 norm.

Uniform sampling, wherein mm pixels (or nn patches) are sampled uniformly from MM pixels (or NN patches).

We found that these methods can reduce the FLOPs by 1.2×1.2\times to 1.6×1.6\times with little or no drop in top-1 accuracy on the ImageNet-1k dataset for both MobileViTv1 (with multi-headed self-attention) and MobileViTv2 (with the proposed separable self-attention) models. However, these improvements in FLOPs did not translate to latency improvements on a mobile device. In fact, models with these sampling methods were significantly slower than the models without these methods. The high-latency of models with these sampling methods on mobile devices can be attributed to their high memory access cost, as these methods change the memory order of tensor. Because of their high-latency on mobile devices, we did not use these methods in the MobileViTv2 model.

Appendix F MobileViTv2 training configurations

Configurations for training and finetuning MobileViTv2-2.0 on the ImageNet-1k and ImageNet-21k-P datasets are given in Table 8 and Table 9 respectively while configurations for finetuning MobileViTv2 on downstream tasks are given in Table 10.