RMT: Retentive Networks Meet Vision Transformers

Qihang Fan, Huaibo Huang, Mingrui Chen, Hongmin Liu, Ran He

Introduction

Vision Transformer (ViT) is an excellent visual architecture highly favored by researchers. However, as the core module of ViT, Self-Attention’s inherent structure lacking explicit spatial priors. Besides, the quadratic complexity of Self-Attention leads to significant computational costs when modeling global information. These issues limit the application of ViT.

Many works have previously attempted to alleviate these issues . For example, in Swin Transformer , the authors partition the tokens used for self-attention by applying windowing operations. This operation not only reduces the computational cost of self-attention but also introduces spatial priors to the model through the use of windows and relative position encoding. In addition to it, NAT changes the receptive field of Self-Attention to match the shape of convolution, reducing computational costs while also enabling the model to perceive spatial priors through the shape of its receptive field.

Different from previous methods, we draw inspiration from the recently successful Retentive Network (RetNet) in the field of NLP. RetNet utilizes a distance-dependent temporal decay matrix to provide explicit temporal prior for one-dimensional and unidirectional text data. ALiBi , prior to RetNet, also applied a similar approach and succeeded in NLP tasks. We extend this temporal decay matrix to the spatial domain, developing a two-dimensional bidirectional spatial decay matrix based on the Manhattan distance among tokens. In our space decay matrix, for a target token, the farther the surrounding tokens are, the greater the degree of decay in their attention scores. This property allows the target token to perceive global information while simultaneously assigning different levels of attention to tokens at varying distances. We introduce explicit spatial prior to the vision backbone using this spatial decay matrix. We name this Self-Attention mechanism, which is inspired by RetNet and incorporates the Manhattan distance as the explicit spatial prior, as Manhattan Self-Attention (MaSA).

Besides explicit spatial priors, another issue caused by global modeling with Self-Attention is the enormous computational burden. Previous sparse attention mechanisms and the way retention is decomposed in RetNet mostly disrupt the spatial decay matrix, making them unsuitable for MaSA. In order to sparsely model global information without compromising the spatial decay matrix, we propose a method to decompose Self-Attention along both axes of the image. This decomposition method decomposes Self-Attention and the spatial decay matrix without any loss of prior information. The decomposed MaSA models global information with linear complexity and has the same receptive field shape as the original MaSA. We compare MaSA with other Self-Attention mechanisms in Fig. 2. It can be seen that our MaSA introduces richer spatial priors to the model than its counterparts.

Based on MaSA, we construct a powerful vision backbone called RMT. We demonstrate the effectiveness of the proposed method through extensive experiments. As shown in Fig. 1, our RMT outperforms the state-of-the-art (SOTA) models on image classification tasks. Additionally, our model exhibits more prominent advantages compared to other models in tasks such as object detection, instance segmentation, and semantic segmentation. Our contributions can be summarized as follows:

We propose a spatial decay matrix based on Manhattan distance to augment Self-Attention, creating the Manhattan Self-Attention (MaSA) with an explicit spatial prior.

We propose a decomposition form for MaSA, enabling linear complexity for global information modeling without disrupting the spatial decay matrix.

Leveraging MaSA, we construct RMT, a powerful vision backbone for general purposes. RMT attains high top-1 accuracy on ImageNet-1k in image classification without extra training data, and excels in tasks like object detection, instance segmentation, and semantic segmentation.

Related Work

Transformer architecture was firstly proposed in to address the training limitation of recurrent model and then achieve massive success in many NLP tasks. By splitting the image into small, non-overlapped patches sequence, Vision Transformer (ViTs) also have attracted great attention and become widely used on vision tasks . Unlike in the past, where RNNs and CNNs have respectively dominated the NLP and CV fields, the transformer architecture has shined through in various modalities and fields . In the computer vision community, many studies are attempting to introduce spatial priors into ViT to reduce the data requirements for training . At the same time, various sparse attention mechanisms have been proposed to reduce the computational cost of Self-Attention .

Prior Knowledge in Transformer.

Numerous attempts have been made to incorporate prior knowledge into the Transformer model to enhance its performance. The original Transformers use trigonometric position encoding to provide positional information for each token. In vision tasks, proposes the use of relative positional encoding as a replacement for the original absolute positional encoding. points out that zero padding in convolutional layers could also provide positional awareness for the ViT, and this position encoding method is highly efficient. In many studies, Convolution in FFN has been employed for vision models to further enrich the positional information in the ViT. For NLP tasks, in the recent Retentive Network , the temporal decay matrix has been introduced to provide the model with prior knowledge based on distance changes. Before RetNet, ALiBi also uses a similar temporal decay matrix.

Methodology

Retentive Network (RetNet) is a powerful architecture for language models. This work proposes the retention mechanism for sequence modeling. Retention brings the temporal decay to the language model, which Transformers do not have. Retention firstly considers a sequence modeling problem in a recurrent manner. It can be written as Eq. 1:

For a parallel training process, Eq. 1 is expressed as:

2 Manhattan Self-Attention

Starting from the retention in RetNet, we evolve it into Manhattan Self-Attention (MaSA). Within MaSA, we transform the unidirectional and one-dimensional temporal decay observed in retention into bidirectional and two-dimensional spatial decay. This spatial decay introduces an explicit spatial prior linked to Manhattan distance into the vision backbone. Additionally, we devise a straightforward approach to concurrently decompose the Self-Attention and spatial decay matrix along the two axes of the image.

In RetNet, retention is unidirectional due to the causal nature of text data, allowing each token to attend only to preceding tokens and not those following it. This characteristic is ill-suited for tasks lacking causal properties, such as image recognition. Hence, we initially broaden the retention to a bidirectional form, expressed as Eq. 3:

From One-dimensional to Two-dimensional Decay:

While retention now supports bi-directional modeling, this capability remains confined to a one-dimensional level and is inadequate for two-dimensional images. To address this limitation, we extend the one-dimensional retention to encompass two dimensions.

In the context of images, each token is uniquely positioned with a two-dimensional coordinate within the plane, denoted as (xn,yn)(x_{n},y_{n}) for the nn-th token. To adapt to this, we adjust each element in the matrix DD to represent the Manhattan distance between the respective token pairs based on their 2D coordinates. The matrix DD is redefined as follows:

Decomposed Manhattan Self-Attention.

In the early stages of the vision backbone, an abundance of tokens leads to high computational costs for Self-Attention when attempting to model global information. Our MaSA encounters this challenge as well. Utilizing existing sparse attention mechanisms , or the original RetNet’s recurrent/chunk-wise recurrent form directly, disrupts the spatial decay matrix based on Manhattan distance, resulting in the loss of explicit spatial prior. To address this, we introduce a simple decomposition method that not only decomposes Self-Attention but also decomposes the spatial decay matrix. The decomposed MaSA is represented in Eq. 6. Specifically, we calculate attention scores separately for the horizontal and vertical directions in the image. Subsequently, we apply the one-dimensional bidirectional decay matrix to these attention weights. The one-dimensional decay matrix signifies the horizontal and vertical distances between tokens (DnmH=γ∣yn−ym∣D^{H}_{nm}=\gamma^{|y_{n}-y_{m}|}, DnmW=γ∣xn−xm∣D^{W}_{nm}=\gamma^{|x_{n}-x_{m}|}):

Based on the decomposition of MaSA, the shape of the receptive field of each token is shown in Fig. 4, which is identical to the shape of the complete MaSA’s receptive field. Fig. 4 indicates that our decomposition method fully preserves the explicit spatial prior.

To further enhance the local expression capability of MaSA, following , we introduce a Local Context Enhancement module using DWConv:

3 Overall Architecture

We construct the RMT based on MaSA, and its architecture is illustrated in Fig. 3. Similar to previous general vision backbones , RMT is divided into four stages. The first three stages utilize the decomposed MaSA, while the last uses the original MaSA. Like many previous backbones , we incorporate CPE into our model.

Experiments

We conducted extensive experiments on multiple vision tasks, such as image classification on ImageNet-1K , object detection and instance segmentation on COCO 2017 , and semantic segmentation on ADE20K . We also make ablation studies to validate the importance of each component in RMT. More details can be found in Appendix.

We train our models on ImageNet-1K from scratch. We follow the same training strategy in , with the only supervision being classification loss for a fair comparison. The maximum rates of increasing stochastic depth are set to 0.1/0.15/0.4/0.5 for RMT-T/S/B/L , respectively. We use the AdamW optimizer with a cosine decay learning rate scheduler to train the models. We set the initial learning rate, weight decay, and batch size to 0.001, 0.05, and 1024, respectively. We adopt the strong data augmentation and regularization used in . Our settings are RandAugment (randm9-mstd0.5-inc1), Mixup (prob=0.8), CutMix (prob=1.0), Random Erasing (prob=0.25). In addition to the conventional training methods, similar to LV-ViT and VOLO , we train a model that utilizes token labeling to provide supplementary supervision.

Results.

We compare RMT against many state-of-the-art models in Tab. 1. Results in the table demonstrate that RMT consistently outperforms previous models across all settings. Specifically, RMT-S achieves 84.1% Top1-accuracy with only 4.5 GFLOPs. RMT-B also surpasses iFormer by 0.4% with similar FLOPs. Furthermore, our RMT-L model surpasses MaxViT-B in top1-accuracy by 0.6% while using fewer FLOPs. Our RMT-T has also outperformed many lightweight models. As for the model trained using token labeling, our RMT-S outperforms the current state-of-the-art BiFormer-S by 0.5%.

2 Object Detection and Instance Segmentation

Results.

3 Semantic Segmentation

We adopt the Semantic FPN and UperNet based on MMSegmentation , apply RMTs which are pretrained on ImageNet-1K as backbone. We use the same setting of PVT to train the Semantic FPN, and we train the model for 80k iterations. All models are trained with the input resolution of 512×512512\times 512. When testing the model, we resize the shorter side of the image to 512 pixels. As for UperNet, we follow the default settings in Swin . We take AdamW with a weight decay of 0.01 as the optimizer to train the models for 160K iterations. The learning rate is set to 6×10−56\times 10^{-5} with 1500 iterations warmup.

Results.

The results of semantic segmentation can be found in Tab. 5. All the FLOPs are measured with the resolution of 512×2048512\times 2048, except the group of RMT-T, which are measured with the resolution of 512×512512\times 512. All our models achieve the best performance in all comparisons. Specifically, our RMT-S exceeds Shunted-S for +1.2 mIoU with Semantic FPN. Moreover, our RMT-B outperforms the recent InternImage-S for +1.8 mIoU. All the above results demonstrate our model’s superiority in dense prediction.

4 Ablation Study

In order to make a strict comparison with previous methods, we align RMT’s hyperparameters (such as whether to use hierarchical structure, the number of channels in the four stages of the hierarchical model, whether to use positional encoding and convolution stem, etc.) of the overall architecture with DeiT and Swin , and only replace the Self-Attention/Window Self-Attention with our MaSA. The comparison results are shown in Tab. 6, where RMT significantly outperforms DeiT-S, Swin-T, and Swin-S.

MaSA.

We verify the impact of Manhattan Self-Attention on the model, as shown in the Tab. 6. MaSA improves the model’s performance in image classification and downstream tasks by a large margin. Specifically, the classification accuracy of MaSA is 0.8% higher than that of vanilla attention.

Softmax.

In RetNet, Softmax is replaced with a non-linear gating function to accommodate its various computational forms . We replace the Softmax in MaSA with this gating function. However, the model utilizing the gating function cannot undergo stable training. It is worth noting that this does not mean the gating function is inferior to Softmax. The gating function may just not be compatible with our decomposed form or spatial decay.

LCE.

Local Context Enhancement also plays a role in the excellent performance of our model. LCE improves the classification accuracy of RMT by 0.3% and enhances the model’s performance in downstream tasks.

CPE.

Just like previous methods, CPE provides our model with flexible position encoding and more positional information, contributing to the improvement in the model’s performance in image classification and downstream tasks.

Convolutional Stem.

The initial convolutional stem of the model provides better local information, thereby further enhancing the model’s performance on various tasks.

Decomposed MaSA.

In RMT-S, we substitute the decomposed MaSA (MaSA-d) in the third stage with the original MaSA to validate the effectiveness of our decomposition method, as illustrated in Tab. 7. In terms of image classification, MaSA-d and MaSA achieve comparable accuracy. However, for semantic segmentation, employing MaSA-d significantly reduces computational burden while yielding similar result.

MaSA v.s. Retention.

As shown in Tab. 8, we replace MaSA with the original retention in the architecture of RMT-S. We partition the tokens into chunks using the method employed in Swin-Transformer for chunk-wise retention. Due to the limitation of retention in modeling one-dimensional causal data, the performance of the vision backbone based on it falls behind RMT. Moreover, the chunk-wise and recurrent forms of retention disrupt the parallelism of the vision backbone, resulting in lower inference speed.

Inference Speed.

We compare the RMT’s inference speed with the recent best performing vision backbones in Tab. 9. Our RMT demonstrates the optimal trade-off between speed and accuracy.

Conclusion

In this work, we propose RMT, a vision backbone with explicit spatial prior. RMT extends the temporal decay used for causal modeling in NLP to the spatial level and introduces a spatial decay matrix based on the Manhattan distance. The matrix incorporates explicit spatial prior into the Self-Attention. Additionally, RMT utilizes a Self-Attention decomposition form that can sparsely model global information without disrupting the spatial decay matrix. The combination of spatial decay matrix and attention decomposition form enables RMT to possess explicit spatial prior and linear complexity. Extensive experiments in image classification, object detection, instance segmentation, and semantic segmentation validate the superiority of RMT.

Appendix A Architecture Details

Our architectures are illustrated in the Tab. 10. For convolution stem, we apply five 3×33\times 3 convolutions to embed the image into 56×5656\times 56 tokens. GELU and batch normalization are used after each convolution except the last one, which is only followed by batch normalization. 3×33\times 3 convolutions with stride 2 are used between stages to reduce the feature map’s resolution. 3×33\times 3 depth-wise convolutions are adopted in CPE. Moreover, 5×55\times 5 depth-wise convolutions are adopted in LCE. RMT-DeiT-S, RMT-Swin-T, and RMT-Swin-S are models that we used in our ablation experiments. Their structures closely align with the structure of DeiT and Swin-Transformer without using techniques like convolution stem, CPE, and others.

Appendix B Experimental Settings

We adopt the same training strategy with DeiT with the only supervision is the classification loss. In particular, our models are trained from scratch for 300 epochs. We use the AdamW optimizer with a cosine decay learning rate scheduler and 5 epochs of linear warm-up. The initial learning rate, weight decay, and batch size are set to 0.001, 0.05, and 1024, respectively. Our augmentation settings are RandAugment (randm9-mstd0.5-inc1), Mixup (prob=0.8), CutMix (probe=1.0), Random Erasing (prob=0.25) and Exponential Moving Average (EMA) . The maximum rates of increasing stochastic depth are set to 0.1/0.15/0.4/0.5 for RMT-T/S/B/L, respectively. For a more comprehensive comparison, we train two versions of the model. The first version uses only classification loss as the supervision, while the second version, in addition to the classification loss, incorporates token labeling introduced by for additional supervision. Models using token labeling are marked with“*”.

COCO Object Detection and Instance Segmentation.

We apply RetinaNet , Mask-RCNN and Cascaded Mask-CNN as the detection frameworks to conduct experiments. We implement them based on the MMDetection . All models are trained under two common settings:“1×1\times” (12 epochs for training) and“3×3\times+MS” (36 epochs with multi-scale augmentation for training). For the “1×1\times” setting, images are resized to the shorter side of 800 pixels. For the “3×3\times+MS”, we use the multi-scale training strategy and randomly resize the shorter side between 480 to 800 pixels. We apply AdamW optimizer with the initial learning rate of 1e-4. For RetinaNet, we use the weight decay of 1e-4 for RetinaNet while we set it to 5e-2 for Mask-RCNN and Cascaded Mask-RCNN. For all settings, we use the batch size of 16, which follows the previous works

ADE20K Semantic Segmentation.

Based on MMSegmentation , we implement UperNet and SemanticFPN to validate our models. For UperNet, we follow the previous setting of Swin-Transformer and train the model for 160k iterations with the input size of 512×512512\times 512. For SemanticFPN, we also use the input resolution of 512×512512\times 512 but train the models for 80k iterations.

Appendix C Efficiency Comparison

We compare the inference speed of RMT with other backbones, as shown in Tab. 11. Our models achieve the best trade-off between speed and accuracy among many competitors.

Appendix D Details of Explicit Decay

We use different γ\gamma for each head of the multi-head ReSA to control the receptive field of each head, enabling the ReSA to perceive multi-scale information. We keep all the γ\gamma of ReSA’s heads within a certain range. Assuming the given receptive field control interval of a specific ReSA module is [a,b][a,b], where both aa and bb are positive real numbers. And the total number of the ReSA module’s heads is NN. The γ\gamma for its iith head can be written as Eq. 8:

For different stages of different backbones, we use different values of aa and bb, with the details shown in Tab. 12.

References