Twins: Revisiting the Design of Spatial Attention in Vision Transformers

Xiangxiang Chu, Zhi Tian, Yuqing Wang, Bo Zhang, Haibing Ren, Xiaolin Wei, Huaxia Xia, Chunhua Shen

Introduction

Recently, Vision Transformers have received increasing research interest. Compared to the widely-used convolutional neural networks (CNNs) in visual perception, Vision Transformers enjoy great flexibility in modeling long-range dependencies in vision tasks, introduce less inductive bias, and can naturally process multi-modality input data including images, videos, texts, speech signals, and point clouds. Thus, they have been considered to be a strong alternative to CNNs. It is expected that vision transformers are likely to replace CNNs and serve as the most basic component in the next-generation visual perception systems.

One of the prominent problems when applying transformers to vision tasks is the heavy computational complexity incurred by the spatial self-attention operation in transformers, which grows quadratically in the number of pixels of the input image. A workaround is the locally-grouped self-attention (or self-attention in non-overlapped windows as in the recent Swin Transformer ), where the input is spatially grouped into non-overlapped windows and the standard self-attention is computed only within each sub-window. Although it can significantly reduce the complexity, it lacks the connections between different windows and thus results in a limited receptive field. As pointed out by many previous works , a sufficiently large receptive field is crucial to the performance, particularly for dense prediction tasks such as image segmentation and object detection. Swin proposes a shifted window operation to tackle the issue, where the boundaries of these local windows are gradually moved as the network proceeds. Despite being effective, the shifted windows may have uneven sizes. The uneven windows result in difficulties when the models are deployed with ONNX or TensorRT, which prefers the windows of equal sizes. Another solution is proposed in PVT . Unlike the standard self-attention operation, where each query computes the attention weights with all the input tokens, in PVT, each query only computes the attention with a sub-sampled version of the input tokens. Although its computational complexity in theory is still quadratic, it is already manageable in practice.

From a unified perspective, the core in the aforementioned vision transformers is how the spatial attention is designed. Thus, in this work, we revisit the design of the spatial attention in vision transformers. Our first finding is that the global sub-sampled attention in PVT is highly effective, and with the applicable positional encodings , its performance can be on par or even better than state-of-the-art vision transformers (e.g., Swin). This results in our first proposed architecture, termed Twins-PCPVT. On top of that, we further propose a carefully-designed yet simple spatial attention mechanism, making our architectures more efficient than PVT. Our attention mechanism is inspired by the widely-used separable depthwise convolutions and thus we name it spatially separable self-attention (SSSA). Our proposed SSSA is composed of two types of attention operations—(i) locally-grouped self-attention (LSA), and (ii) global sub-sampled attention (GSA), where LSA captures the fine-grained and short-distance information and GSA deals with the long-distance and global information. This leads to the second proposed vision transformer architecture, termed Twins-SVT. It is worth noting that both attention operations in the architecture are efficient and easy-to-implement with matrix multiplications in a few lines of code. Thus, all of our architectures here have great applicability and can be easily deployed.

We benchmark our proposed architectures on a number of visual tasks, ranging from image-level classification to pixel-level semantic/instance segmentation and object detection. Extensive experiments show that both of our proposed architectures perform favorably against other state-of-the-art vision transformers with similar or even reduced computational complexity.

Related Work

Convolutional neural networks. Characterized by local connectivity, weight sharing, shift-invariance and pooling, CNNs have been the de facto standard model for computer vision tasks. The top-performing models in image classification also serve as the strong backbones for downstream detection and segmentation tasks.

Vision Transformers. Transformer was firstly proposed by for machine translation tasks, and since then they have become the state-of-the-art models for NLP tasks, overtaking the sequence-to-sequence approach built on LSTM. Its core component is multi-head self-attention which models the relationship between input tokens and shows great flexibility.

In 2020, Transformer was introduced to computer vision for image and video processing . In the image classification task, ViT and DeiT divide the images into patch embedding sequences and feed them into the standard transformers. Although vision transformers have been proved compelling in image classification compared with CNNs, a challenge remains when it is applied to dense prediction tasks such as object detection and segmentation. These tasks often require feature pyramids for better processing objects of different scales, and take as inputs the high-resolution images, which significantly increase the computational complexity of the self-attention operations.

Recently, Pyramid Vision Transformer (PVT) is proposed and can output the feature pyramid as in CNNs. PVT has demonstrated good performance in a number of dense prediction tasks. The recent Swin Transformer introduces non-overlapping window partitions and restricts self-attention within each local window, resulting in linear computational complexity in the number of input tokens. To interchange information among different local areas, its window partitions are particularly designed to shift between two adjacent self-attention layers. The semantic segmentation framework OCNet shares some similarities with us and they also interleave the local and global attention. Here, we demonstrate this is a general design paradigm in vision transformer backbones rather than merely an incremental module in semantic segmentation.

Grouped and Separable Convolutions. Grouped convolutions are originally proposed in AlexNet for distributed computing. They were proved both efficient and effective in speeding up the networks. As an extreme case, depthwise convolutions use the number of groups that is equal to the input or output channels, which is followed by point-wise convolutions to aggregate the information across different channels. Here, the proposed spatially separable self-attention shares some similarities with them.

Positional Encodings. Most vision transformers use absolute/relative positional encodings, depending on downstream tasks, which are based on sinusoidal functions or learnable . In CPVT , the authors propose the conditional positional encodings, which are dynamically conditioned on the inputs and show better performance than the absolute and relative ones.

Our Method: Twins

We present two simple yet powerful spatial designs for vision transformers. The first method is built upon PVT and CPVT , which only uses the global attention. The architecture is thus termed Twins-PCPVT. The second one, termed Twins-SVT, is based on the proposed SSSA which interleaves local and global attention.

PVT introduces the pyramid multi-stage design to better tackle dense prediction tasks such as object detection and semantic segmentation. It inherits the absolute positional encoding designed in ViT and DeiT . All layers utilize the global attention mechanism and rely on spatial reduction to cut down the computation cost of processing the whole sequence. It is surprising to see that the recently-proposed Swin transformer , which is based on shifted local windows, can perform considerably better than PVT, even on dense prediction tasks where a sufficiently large receptive field is even more crucial to good performance.

In this work, we surprisingly found that the less favored performance of PVT is mainly due to the absolute positional encodings employed in PVT . As shown in CPVT , the absolute positional encoding encounter difficulties in processing the inputs with varying sizes (which are common in dense prediction tasks). Moreover, this positional encoding also breaks the translation invariance. On the contrary, Swin transformer makes use of the relative positional encodings, which bypasses the above issues. Here, we demonstrate that this is the main cause why Swin outperforms PVT, and we show that if the appropriate positional encodings are used, PVT can actually achieve on par or even better performance than the Swin transformer.

Here, we use the conditional position encoding (CPE) proposed in CPVT to replace the absolute PE in PVT. CPE is conditioned on the inputs and can naturally avoid the above issues of the absolute encodings. The position encoding generator (PEG) , which generates the CPE, is placed after the first encoder block of each stage. We use the simplest form of PEG, i.e., a 2D depth-wise convolution without batch normalization. For image-level classification, following CPVT, we remove the class token and use global average pooling (GAP) at the end of the stage . For other vision tasks, we follow the design of PVT. Twins-PCPVT inherits the advantages of both PVT and CPVT, which makes it easy to be implemented efficiently. Our extensive experimental results show that this simple design can match the performance of the recent state-of-the-art Swin transformer. We have also attempted to replace the relative PE with CPE in Swin, which however does not result in noticeable performance gains, as shown in our experiments. We conjecture that this maybe due to the use of shifted windows in Swin, which might not work well with CPE.

We report the detailed settings of Twins-PCPVT in Table 9 (in supplementary), which are similar to PVT . Therefore, Twins-PCPVT has similar FLOPs and number of parameters to .

2 Twins-SVT

Vision transformers suffer severely from the heavy computational complexity in dense prediction tasks due to high-resolution inputs. Given an input of H×WH\times W resolution, the complexity of self-attention with dimension dd is O(H2W2d)\mathcal{O}(H^{2}W^{2}d). Here, we propose the spatially separable self-attention (SSSA) to alleviate this challenge. SSSA is composed of locally-grouped self-attention (LSA) and global sub-sampled attention (GSA).

Motivated by the group design in depthwise convolutions for efficient inference, we first equally divide the 2D feature maps into sub-windows, making self-attention communications only happen within each sub-window. This design also resonates with the multi-head design in self-attention, where the communications only occur within the channels of the same head. To be specific, the feature maps are divided into m×nm\times n sub-windows. Without loss of generality, we assume H%m=0H\%m=0 and W%n=0W\%n=0. Each group contains HWmn\frac{HW}{mn} elements, and thus the computation cost of the self-attention in this window is O(H2W2m2n2d)\mathcal{O}(\frac{H^{2}W^{2}}{m^{2}n^{2}}d), and the total cost is O(H2W2mnd)\mathcal{O}(\frac{H^{2}W^{2}}{mn}d). If we let k1=Hmk_{1}=\frac{H}{m} and k2=Wnk_{2}=\frac{W}{n}, the cost can be computed as O(k1k2HWd)\mathcal{O}(k_{1}k_{2}HWd), which is significantly more efficient when k1≪Hk_{1}\ll H and k2≪Wk_{2}\ll W and grows linearly with HWHW if k1k_{1} and k2k_{2} are fixed.

Although the locally-grouped self-attention mechanism is computation friendly, the image is divided into non-overlapping sub-windows. Thus, we need a mechanism to communicate between different sub-windows, as in Swin. Otherwise, the information would be limited to be processed locally, which makes the receptive field small and significantly degrades the performance as shown in our experiments. This resembles the fact that we cannot replace all standard convolutions by depth-wise convolutions in CNNs.

Global sub-sampled attention (GSA).

A simple solution is to add extra standard global self-attention layers after each local attention block, which can enable cross-group information exchange. However, this approach would come with the computation complexity of O(H2W2d)\mathcal{O}(H^{2}W^{2}d).

Here, we use a single representative to summarize the important information for each of m×nm\times n sub-windows and the representative is used to communicate with other sub-windows (serving as the key in self-attention), which can dramatically reduce the cost to O(mnHWd)=O(H2W2dk1k2)\mathcal{O}(mnHWd)=\mathcal{O}(\frac{H^{2}W^{2}d}{k_{1}k_{2}}). This is essentially equivalent to using the sub-sampled feature maps as the key in attention operations, and thus we term it global sub-sampled attention (GSA). If we alternatively use the aforementioned LSA and GSA like separable convolutions (depth-wise + point-wise). The total computation cost is O(H2W2dk1k2+k1k2HWd)\mathcal{O}(\frac{H^{2}W^{2}d}{k_{1}k_{2}}+k_{1}k_{2}HWd). We have H2W2dk1k2+k1k2HWd≥2HWdHW\frac{H^{2}W^{2}d}{k_{1}k_{2}}+k_{1}k_{2}HWd\geq 2HWd\sqrt{HW}. The minimum is obtained when k1⋅k2=HWk_{1}\cdot k_{2}=\sqrt{HW}. We note that H=W=224H=W=224 is popular in classification. Without loss of generality, we use square sub-windows, i.e., k1=k2k_{1}=k_{2}. Therefore, k1=k2=15k_{1}=k_{2}=15 is close to the global minimum for H=W=224H=W=224. However, our network is designed to include several stages with variable resolutions. Stage 1 has feature maps of 56 ×\times 56, the minimum is obtained when k1=k2=56≈7k_{1}=k_{2}=\sqrt{56}\approx 7. Theoretically, we can calibrate optimal k1k_{1} and k2k_{2} for each of the stages. For simplicity, we use k1=k2=7k_{1}=k_{2}=7 everywhere. As for stages with lower resolutions, we control the summarizing window-size of GSA to avoid too small amount of generated keys. Specifically, we use the size of 4, 2 and 1 for the last three stages respectively.

As for the sub-sampling function, we investigate several options including average pooling, depth-wise strided convolutions, and regular strided convolutions. Empirical results show that regular strided convolutions perform best here. Formally, our spatially separable self-attention (SSSA) can be written as

where LSA means locally-grouped self-attention within a sub-window; GSA is the global sub-sampled attention by interacting with the representative keys (generated by the sub-sampling functions) from each sub-window z^ij∈Rk1×k2×C\hat{\bf{z}}_{ij}\in\mathcal{R}^{k_{1}\times k_{2}\times C}. Both LSA and GSA have multiple heads as in the standard self-attention.The PyTorch code of LSA is given in Algorithm 1 (in supplementary).

Again, we use the PEG of CPVT to encode position information and process variable-length inputs on the fly. It is inserted after the first block in each stage.

Model variants. The detailed configure of Twins-SVT is shown in Table 10 (in supplementary). We try our best to use the similar settings as in Swin to make sure that the good performance is due to the new design paradigm.

Comparison with PVT. PVT entirely utilizes global attentions as DeiT does while our method makes use of spatial separable-like design with LSA and GSA, which is more efficient.

Comparison with Swin. Swin utilizes the alternation of local window based attention where the window partitions in successive layers are shifted. This is used to introduce communication among different patches and to increase the receptive field. However, this procedure is relatively complicated and may not be optimized for speed on devices such as mobile devices. Swin Transformer depends on torch.roll() to perform cyclic shift and its reverse on features. This operation is memory unfriendly and rarely supported by popular inference frameworks such as NVIDIA TensorRT, Google Tensorflow-Lite, and Snapdragon Neural Processing Engine SDK (SNPE), etc. This hinders the deployment of Swin either on the server-side or on end devices in a production environment. In contrast, Twins models don’t require such an operation and only involve matrix multiplications that are already optimized well in modern deep learning frameworks. Therefore, it can further benefit from the optimization in a production environment. For example, we converted Twins-SVT-S from PyTorch to TensorRT , and its throughput is boosted by 1.7×\times. Moreover, our local-global design can better exploit the global context, which is known to play an important role in many vision tasks.

Finally, one may note that the network configures (e.g., such as depths, hidden dimensions, number of heads, and the expansion ratio of MLP) of our two variants are sightly different. This is intended because we want to make fair comparisons to the two recent well-known transformers PVT and Swin. PVT prefers a slimmer and deeper design while Swin is wider and shallower. This difference makes PVT have slower training than Swin. Twins-PCPVT is designed to compare with PVT and shows that a proper positional encoding design can greatly boost the performance and make it on par with recent state-of-the-art models like Swin. On the other hand, Twins-SVT demonstrates the potential of a new paradigm as to spatially separable self-attention is highly competitive to recent transformers.

Experiments

We first present the ImageNet classification results with our proposed models. We carefully control the experiment settings to make fair comparisons against recent works . All our models are trained for 300 epochs with a batch size of 1024 using the AdamW optimizer . The learning rate is initialized to be 0.001 and decayed to zero within 300 epochs following the cosine strategy. We use a linear warm-up in the first five epochs and the same regularization setting as in . Note that we do not utilize extra tricks in to make fair comparisons although it may further improve the performance of our method. We use increasing stochastic depth augmentation of 0.2, 0.3, 0.5 for small, base and large model respectively. Following Swin , we use gradient clipping with a max norm of 5.0 to stabilize the training process, which is especially important for the training of large models.

We report the classification results on ImageNet-1K in Table 1. Twins-PCPVT-S outperforms PVT-small by 1.4%1.4\% and obtains similar result as Swin-T with 18% fewer FLOPs. Twins-SVT-S is better than Swin-T with about 35%35\% fewer FLOPs. Other models demonstrate similar advantages.

It is interesting to see that, without bells and whistles, Twins-PCPVT performs on par with the recent state-of-the-art Swin, which is based on much more sophisticated designs as mentioned above. Moreover, Twins-SVT also achieves similar or better results, compared to Swin, indicating that the spatial separable-like design is an effective and promising paradigm.

One may challenge our improvements are due to the use of the better positional encoding PEG. Thus, we also replace the relative PE in Swin-T with PEG , but the Swin-T’s performance cannot be improved (being 81.2%).

2 Semantic Segmentation on ADE20K

We further evaluate the performance on segmentation tasks. We test on the ADE20K dataset , a challenging scene parsing task for semantic segmentation, which is popularly evaluated by recent Transformer-based methods. This dataset contains 20K images for training and 2K images for validation. Following the common practices, we use the training set to train our models and report the mIoU on the validation set. All models are pretrained on the ImageNet-1k dataset.

Twins-PCPVT vs. PVT. We compare our Twins-PCPVT with PVT because they have similar design and computational complexity. To make fair comparisons, we use the Semantic FPN framework and exactly the same training settings as in PVT. Specifically, we train 80K steps with a batch size of 16 using AdamW . The learning rate is initialized as 1×\times10-4 and scheduled by the ‘poly’ strategy with the power coefficient of 0.9. We apply the drop-path regularization of 0.2 for the backbone and weight decay 0.0005 for the whole network. Note that we use a stronger drop-path regularization of 0.4 for the large model to avoid over-fitting. For Swin, we use their official code and trained models. We report the results in Table 2. With comparable FLOPs, Twins-PCPVT-S outperforms PVT-Small with a large margin (+4.5% mIoU), which also surpasses ResNet-50 by 7.6% mIoU. It also outperforms Swin-T with a clear margin. Besides, Twins-PCPVT-B also achieves 3.3% higher mIoU than PVT-Medium, and Twins-PCPVT-L surpasses PVT-Large with 4.3% higher mIoU.

Twins-SVT vs. Swin. We also compare our Twins-SVT with the recent state-of-the-art model Swin . With the Semantic FPN framework and the above settings, Twins-SVT-S achieves better performance (+1.7%) than Swin-T. Twins-SVT-B obtains comparable performance with Swin-S and Twins-SVT-L outperforms Swin-B by 0.7% mIoU (left columns in Table 2). In addition, Swin evaluates its performance using the UperNet framework . We transfer our method to this framework and use exactly the same training settings as . To be specific, we use the AdamW optimizer to train all models for 160k iterations with a global batch size of 16. The initial learning rate is 6×\times10-5 and linearly decayed to zero. We also utilize warm-up during the first 1500 iterations. Moreover, we apply the drop-path regularization of 0.2 for the backbone and weight decay 0.01 for the whole network. We report the mIoU of both single scale and multi-scale testing (we use scales from 0.5 to 1.75 with step 0.25) in the right columns of Table 2. Both with multi-scale testing, Twins-SVT-S outperforms Swin-T by 1.3% mIoU. Moreover, Twins-SVT-L achieves new state of the art result 50.2% mIoU under comparable FLOPs and outperforms Swin-B by 0.5% mIoU. Twins-PCPVT also achieves comparable performance to Swin .

3 Object Detection and Segmentation on COCO

We evaluate the performance of our method using two representative frameworks: RetinaNet and Mask RCNN . Specifically, we use our transformer models to build the backbones of these detectors. All the models are trained under the same setting as in . Since PVT and Swin report their results using different frameworks, we try to make fair comparison and build consistent settings for future methods. Specifically, we report standard 1×\times-schedule (12 epochs) detection results on the COCO 2017 dataset in Tables 3 and 4. As for the evaluation based on RetinaNet, we train all the models using AdamW optimizer for 12 epochs with a batch size of 16. The initial learning rate is 1×\times10-4, started with 500-iteration warmup and decayed by 10×\times at the 8th and 11th epoch, respectively. We use stochastic drop path regularization of 0.2 and weight decay 0.0001. The implementation is based on MMDetection . For the Mask R-CNN framework, we use the initial learning rate of 2×\times10-4 as in . All other hyper-parameters follow the default settings in MMDetection. As for 3×\times experiments, we follow the common multi-scale training in , i.e., randomly resizing the input image so that its shorter side is between 480 and 800 while keeping longer one less than 1333. Moreover, for 3×\times training of Mask R-CNN, we use an initial learning rate of 0.0001 and weight decay of 0.05 for the whole network as .

For 1×\times schedule object detection with RetinaNet, Twins-PCPVT-S surpasses PVT-Small with 2.6% mAP and Twins-PCPVT-B exceeds PVT-Medium by 2.4% mAP on the COCO val2017 split. Twins-SVT-S outperforms Swin-T with 1.5% mAP while using 12% fewer FLOPs. Our method outperform the others with similar advantage in 3×\times experiments.

For 1×\times object segmentation with the Mask R-CNN framework, Twins-PCPVT-S brings similar improvements (+2.5% mAP) over PVT-Small. Compared with PVT-Medium, Twins-PCPVT-B obtains 2.6% higher mAP, which is also on par with that of Swin. Both Twins-SVT-S and Twins-SVT-B achieve better or slightly better performance compared to the counterparts of Swin. As for large models, our results are shown in Table 8 (in supplementary) and we also achieve better performance with comparable FLOPs.

4 Ablation Studies

We evaluate different combinations of LSA and GSA based on our small model and present the ablation results in Table 5. The models with only locally-grouped attention fail to obtain good performance (76.9%) because this setting has a limited and small receptive field. An extra global attention layer in the last stage can improve the classification performance by 3.6%. Local-Local-Global (abbr. LLG) also achieves good performance (81.5%), but we do not use this design in this work.

Sub-sampling functions.

We further study how the different sub-sampling functions affect the performance. Specifically, we compare the regular strided convolutions, separable convolutions and average pooling based on the ‘small’ model and present the results in Table 6. The first option performs best and therefore we choose it as our default implementation.

Positional Encodings.

We replace the relative positional encoding with CPVT for Swin-T and report the detection performance on COCO with RetinaNet and Mask R-CNN in Table 7. The CPVT-based Swin cannot achieve improved performance with both frameworks, which indicates that our performance improvements should be owing to the paradigm of Twins-SVT instead of the positional encodings.

Conclusion

In this paper, we have presented two powerful vision transformer backbones for both image-level classification and a few downstream dense prediction tasks. We dub them as twin transformers: Twins-PCPVT and Twins-SVT. The former variant explores the applicability of conditional positional encodings in pyramid vision transformer , confirming its potential for improving backbones in many vision tasks. In the latter variant we revisit current attention design to proffer a more efficient attention paradigm. We find that interleaving local and global attention can produce impressive results, yet it comes with higher throughputs. Both transformer models set a new state of the art in image classification, objection detection and semantic/instance segmentation.

References

Appendix A Experiment

Appendix B Algorithm

Appendix C Architecture Setting