SdAE: Self-distillated Masked Autoencoder

Yabo Chen, Yuchen Liu, Dongsheng Jiang, Xiaopeng Zhang, Wenrui Dai, Hongkai Xiong, Qi Tian

Introduction

The masked language modeling task (MLM) has shown great success in self-supervised learning (SSL) for natural language processing. In computer vision, contrastive learning/instance discrimination is a promising direction recently, which regards each instance in the training dataset as a single category. Based on instance discrimination , some methods show the effectiveness in many computer vision tasks. With the development of vision transformer , inspired by natural language processing (NLP), the generative based self-supervised learning (SSL) methods using masked image modeling (MIM) task have grown in concern. MIM first randomly masks some proportion of image patches, and then recovers the masked patches based on the corrupted image.

BeiT transfers the MIM task into discrete token classification using a pre-trained discrete autoencoder dVAE (DALL-E ). PeCO modifies the generating procedure of codebook by enforcing the perceptual similarity during the VAE training. Similarly, these methods rely on a pre-trained feature descriptor to obtain the latent representation of masked tokens. These designs requiring an additional codebook are a kind of ‘‘``pre-pretraining"".

MAE proposes an asymmetric encoder-decoder architecture that can reconstruct the raw pixels from the latent representation. MaskFeat uses the hand-crafted feature descriptor Histograms of Oriented Gradients (HOG) to tokenize the image features. Although these methods do not need additional codebooks, they employ restoring low-level representations such as pixels for masked image modeling tasks. Nevertheless, restoring low-level representations such as pixels is redundant for high semantic level tasks. Moreover, directly reconstructing pixels may lead to an optimization gap between pre-training and downstream tasks, i.e., good reconstruction quality may not always lead to the high descriptive capability of the model.

Considering the above issue, we propose a simple yet effective self-distillated masked autoencoder structure called SdAE. In SdAE, we claim that MAE itself can produce good representations in an effective and efficient way, and can eliminate the representation gap when used as codebook appropriately. Without needing a codebook in advance nor modeling a low-level representation, SdAE uses a self-distillated teacher-student network to produce the latent representation as reconstruction targets. The student branch consists of the asymmetric encoder-decoder architecture that feeds unmasked images, and the teacher branch contains an encoder to produce latent representation and updates weights from the student using Exponential Moving Average (EMA).

When introducing the teacher branch into the masked autoencoder structure, the easiest way is to feed the full image into the teacher network directly. However, there is no computational loss on the unmasked tokens. Obviously, due to the spatial redundancy that exists in the image, it is not optimal to put the whole image into the teacher branch. In addition, simply putting all masked tokens is still resource-consuming and faced with performance degradation compared with using the whole image as input. So there grows another concern about how to produce better latent representation using the raw images.

We further discuss this problem from the perspective of information bottleneck and propose a multi-fold masking strategy to produce good views for the self-distillated masked autoencoders as well as reduce the computational complexity. The contributions of this paper are summarized as below:

We propose a novel self-distillated masked autoencoder structure that can construct a learnable high-level reconstruction target rather than extra pre-trained codebooks or low-level pixels, and find that MAE itself can produce a better codebook.

We discuss how to produce good views for the teacher branch and propose a multi-fold masking strategy to keep mutual information from the teacher branch relevant to the student one. This strategy can also save computation resources.

Methods

Firstly, we elaborate the basic framework of masked image modeling and derive the objective function of our masked feature reconstruction methods in Section 2.1. Figure 2 depicts our proposed self-distillated masked autoencoder that consists of a two-branch network, i.e., teacher and student branches. After that, we present theoretical discussions on how to produce good views for the teacher branch to build latent representation. In addition, we propose a multi-fold masking strategy, which is tailored for balancing the information between the teacher and student branches as in Section 2.2. Finally, in Section 2.3, the distillation strategy for the teacher-student framework is demonstrated.

Considering MIM is essentially a reconstruction task. It is related more to regression tasks than to classification tasks. Thus, we assume the noise in the deviation of reconstructed and ground truth values follow the standard Gaussian distribution as η∼N(0,1)\eta\sim N(0,1), and the objective is to minimize

Specifically, for MAE , fϕf_{\phi} is the identity function, and fθf_{\theta} is the masked autoencoder that reconstructs masked patches in the pixel spatial space. However, we argue that (1) The fϕf_{\phi} is a fixed identity function without adaptation during pre-training that undermines the effectiveness of the self-supervised training; (2) Memorying each pixel of the images by reconstruction in the low semantic level space is sub-optimal and less efficient for capturing representations for high-level tasks. There exists an optimization direction gap in that the quality of the reconstruction may not always increase the descriptive capability of the model.

Inspired by MAE , we propose a value normalization function upon teacher outputs. Specifically, we compute the mean and standard deviation of feature values within a patch and use them to normalize the teacher outputs as

where ϵ\epsilon is a small value to prevent the denominator from being 0. We find that using normalized features as the reconstruction target improves the representation quality.

Then we minimize the normalized teacher features with the output features of the student decoder based on the feature cosine similarity, and Eq. (2) is reformulated as:

2 Discussions on The Teacher Branch Feeding

As mentioned in MAE , languages are human-generated signals that are highly semantic and information-dense. Predicting a few missing words per sentence can induce sophisticated language understanding, but images are natural signals with heavy spatial redundancy. Directly feeding the whole image into the teacher network to produce features may be sub-optimal.

So the multi-mask strategy can save complexities by a factor of tt. Detailed theoretical complexity can be seen in Table 1.

In practice, directly feeding the whole image into the teacher network will increase almost half of the computation costs compared with MAE, as shown in Table 1. Using all masked patches as input like CAE will also increase 38.1%38.1\% pre-training time. Using multi-fold masking strategy, we can save the extra pre-training time costs by only 28.1%28.1\%.

The main difference between our multi-fold masking and previous work multi-crop is discussed as follows. The multi-crop strategy needs to train multiple random views with different sizes without concern of complexity. By contrast, multi-fold masking just rearranges the mask tokens and does not create new views. Multi-fold masking can also save training complexity by calculating self-attention on only a small group of tokens.

Information Bottleneck and Multi-fold masking In this section, we will follow the assumption from the views of information bottleneck theory in to discuss the benefits of multi-fold masking except for saving costs.

3 Distillation Strategy

Considering a teacher network, we hypothesize that it is desirable to build the reconstruction target representations in the (i) online and (ii) consistent way. Now that the student network fθf_{\theta} is trained by back-propagation to minimize the feature reconstruction loss. The teacher network is updated in a momentum update way using exponential moving average (EMA). Specifically, denoting the parameters of fϕf_{\phi} as ϕ\phi and those of fθf_{\theta} as θ\theta, we update ϕ\phi by:

Here η∈[0,1)\eta\in[0,1) is a momentum coefficient to control the frequency of updates from the student model. The codebook should not update too often. Otherwise, the model may fail to converge.

Experiments

This section evaluates our pre-trained feature representation on several unsupervised benchmarks. We first evaluate the classification performance on ImageNet-1k under fine-tuning and linear probing. Then we transfer the pre-trained features to several downstream tasks, i.e., semantic segmentation and object detection. Finally, we conducted an ablation study on the key components of SdAE.

We study the fine-tuning on the ILSVRC-2012 ImageNet dataset with 1k classes and 1.3M images. For a fair comparison, we directly follow most of the hyperparameters of MAE in our fine-tuning experiments. All experiments reported are only fine-tuning for 100 epochs (vs. 300 training from scratch). We compare our SdAE with Vision Transformers trained by random initialization and previous self-supervised learning methods. As shown in Table 6, compared with the models trained by random initialization which only achieves 81.8% top-1 accuracy with ViT-B, our SdAE achieves 84.1%, demonstrating the effectiveness of pre-training with unlabeled data.

Compared with previous self-supervised methods, our proposed SdAE surpasses them on ImageNet fine-tuning by a large margin. For ViT-B, our SdAE outperforms MAE by 1.2% top-1 accuracy with the same number of training epochs, demonstrating that MIM on high-level latent feature space is more effective than low-level pixel space. Besides, our SdAE outperforms BEiT by 1.1% top-1 accuracy. Moreover, compared to the recently proposed CAE, our SdAE achieves 0.8% top-1 accuracy gain, demonstrating the effectiveness of our self-distillated design and multi-fold masking strategy. In addition, with only 100 epochs pre-training, SdAE can achieve comparable performance with MAE using 1600 epochs pre-training and surpass 300 epochs pre-trained CAE. Our proposed SdAE also surpasses above methods on ImageNet linear probing with the same training epochs. As a MIM based method, SdAE can also surpasses MIM based methods with the same pre-training epochs. The phenomenon that contrastive based methods surpass the MIM based ones on linear probing is also discussed in MAE and CAE . In terms of linear probing, contrastive learning mainly cares about the 1000 classes and MIM methods may care about the classes beyond the 1000 classes. So fine-tuning measurement may better validate the effectiveness of MIM based methods.

2 Semantic Segmentation

We evaluate the learned representation of our SdAE on the ADE20K benchmark with 25K images and 150 semantic categories. The mean Intersection of Union (mIoU) averaged over all semantic categories is reported as the evaluation metric. Table 3 shows that our SdAE achieves the state-of-the-art performance with 48.6 mIoU for 300 pre-training epochs. Our SdAE outperforms BEiT, MAE, and CAE by 3.1, 2.8, and 1.9 mIoU, respectively. Besides, our SdAE with 300 pre-training epochs even outperforms MAE with 1600 pre-training epochs.

3 Object Detection

Following CAE , we fine-tune Mask R-CNN in an end-to-end manner on COCO . The ViT backbone is adapted for use with FPN . The box AP for object detection and the mask AP for instance segmentation is reported in Table 4. Our method (300 epochs, ViT-B) is consistently superior to all the other models. Our SdAE performs better than the recent published CAE (48.9 vs. 48.0 APb). Besides, it is worth mentioning that our SdAE (300 epochs) even outperforms MAE (1600 epochs) by 0.5 APb. As an effective framework for self-supervised learning, we achieve better performance with fewer training epochs.

Ablation Studies

In this section, we present ablation studies to better evaluate the contributions of each component and hyperparameter settings in our proposed SdAE. Unless specified, all results are compared with models pre-trained for 100 epochs for efficiency, and we report the top-1 accuracy after fine-tuning for 100 epochs.

In this subsection, we present ablation studies on each component. Table 5 shows that our proposed teacher normalization achieves 0.3% top-1 accuracy gain. Only inputting the masked tokens into the teacher network suffers from 0.5% performance degradation due to insufficient information exploration. While using our proposed multi-fold masking strategy, we can achieve 0.6% improvement compared with only masked token inputs. Besides, multi-fold masking even outperforms taking full image as inputs more efficiently.

2 The EMA Strategy

This experiment is conducted without the multi-fold masking strategy to evaluate the raw performance of the EMA strategy. Specifically, we have two settings for the EMA strategy: (1) update the parameters of the teacher branch with EMA for each training iteration. (2) update the parameters of the teacher branch with EMA for each training epoch. As shown in Figure 4 (a), conducting the EMA strategy to update the teacher branch each training iteration is extremely sensitive to the value of the momentum coefficient. Specifically, only changing the value by 0.001 results in the sharply degraded performance. In contrast, the EMA strategy to update the teacher branch per batch is more robust to the momentum coefficient. Besides, conducting the EMA strategy per epoch can better benefit from long training epochs.

3 The Multi-fold Masking Strategy

We further conduct an ablation study on the multi-fold masking strategy. Firstly, we study experiments on the teacher crop that mask some of the target tokens and feed the remaining tokens into the teacher. The whole image is divided into 196 image patches with the size of 16×1616\times 16, and we randomly sample a different number of total patches, e.g., 36, 49, 79, 122, 147 and 196 as whole image input into the teacher network. As shown in Figure 4 (b), for the case that 36 image patches (roughly 18% of the whole image) are input to the teacher network, our method can achieve 82.81% top-1 accuracy. 83.04% top-1 accuracy is achieved when we take the whole image as input. Comparably, only 0.23% performance gain is achieved when five times as many image patches are input. The problem of spatial redundancy is also common in the input views of the teacher network.

Correspondingly, we take the multi-fold masking strategy. Apart from the 49 patches that are input to the student network, the left 147 image patches are divided into 2 fold as 2×79\times 79 patches, 3 fold as 3×49\times 49 patches or 4 fold as 4×36\times 36 patches. As shown in Figure 4 (b), a multi-fold masking strategy can consistently improve the performance. For 3 fold with 49 patches, our method achieves the best performance. This experiment also proves our information bottleneck discussion, that with comparable mutual information between each fold and the student input, we can get the best performance.

4 Evaluation of Teacher and Student Models

As shown in Figure 4 (c), for every 20 epochs of our training, we finetune our pre-trained teacher and student models for 5 epochs on ImageNet-1K. The student model performs better than the teacher model in the initial epochs when the learning rate is relatively low due to warm up and then is very close to the teacher when the learning rate increases. We demonstrate that the teacher model and student model trained by our SdAE achieve similar performance. That proves the teacher branch can really learn high semantic level related representations. The teacher model can be seen as an ensemble over many student models. We follow most previous works to use the student model as the final model.

Related Work

Previous self-supervised methods utilize different priors of images to design clever pretext tasks, such as predicting the patch positions , inpainting , colorization , and rotation prediction . Recent progress in self-supervised learning focuses on contrastive learning and masked image modeling.

Masked Image Modeling. Motivated by BERT for masked language modeling(MLM), masked image modeling(MIM) is proposed to learn representations from images corrupted by masking. Recent works leverage the powerful Vision Transformer, which matches the masked image modeling task. iGPT operates on sequences of pixels and predicts unknown pixels. Vision transformer studies masked average value prediction for self-supervised learning. BEiT proposes to predict discrete tokens based on a pre-trained image tokenizer while iBOT proposes an online tokenizer. MAE proposes a masked autoencoder for reconstructing the image pixels. CAE provides a two-branch network for MIM, and the masked features are modeled as a regularization to align the mask and unmasked features.

Contrastive Learning. Contrastive learning is proposed based on the InfoMax principle, which aims at maximizing the mutual information across different augmentations of the same image . This augmentation invariance is achieved by enforcing the similarity over different views of the same image while avoiding model collapse. Model collapse can be avoided by introducing negative samples for noise-contrastive estimation . The models typically regard various data augmentations as different views of an image and then make the representations of positive pairs similar while pushing negative pairs away. Large memory banks or large batch size is leveraged to obtain more informative negative samples. BYOL and SimSiam employ an asymmetric network and eliminate the requirement of negative samples. Other methods use clustering to organize image examples.

Conclusion

In this paper, we propose a novel Self-distillated Masked Autoencoder for masked feature reconstruction, namely SdAE. We first formulate the framework of the masked image modeling task, based on which we analyze that existing methods are sub-optimal due to the need of pre-trained codebooks or just reconstructing the low-level pixels. We propose a self-distillated framework for reconstruction in the high-level feature space. Besides, we analyze how to build good views for the teacher branch to produce latent representation from the perspective of information bottleneck theory and propose a multi-fold masking strategy that can also relieve the spatial redundancy. Experimentally, based on a vanilla ViT-Base model, our SdAE achieves a new state-of-the-art of 84.1% top-1 accuracy with only 300 epochs pre-training.

Acknowledgments. This work was supported in part by the National Natural Science Foundation of China under Grant 61932022, Grant 61931023, Grant 61971285, Grant 62120106007, and in part by the Program of Shanghai Science and Technology Innovation Project under Grant 20511100100.

Appendix 0.A Appendix

Pretraining. The settings are almost the same as MAE . We use AdamW for optimization and train the CAE for 300 epochs with batch size 768. We set the learning rate as 8e-4, with cosine learning rate decay and a 60 epoch warmup, and set the weight decay as 0.05. We employ the drop path with the ratio of 0.25 only on the encoder. The momentum coefficient is set as 0.96 and with a cosine schedule to 0.99, and EMA is conducted per pre-training epoch. We set the mask ratio as 0.75 with only 49 tokens fed into the student branch. The masked tokens are divided into 3 folds where each fold contains 49 tokens, and all folds are fed into a shared weighted teacher branch.

Fine-tuning on ImageNet. We follow the fine-tuning setting almost the same as MAE to use layer-wise learning rate decay, weight decay, and AdamW. The batch size is 2048, and the weight decay is 0.05. For ViT-B, we train 100 epochs with base learning rate 5e-4, layer-wise decay rate 0.65, drop path rate 0.1, and warmup epoch 10. For ViT-L, we train 50 epochs with base learning rate 1e-3, layer-wise decay rate 0.75, drop path rate 0.2, and warmup epoch 5.

Object Detection and Instance Segmentation on COCO. We utilize the same setting as CAE that uses multi-scale training and resizes the image with the size of the short side between 480 and 800 and the long side no larger than 1333. The batch size is 32, the learning rate is 3e-4, and the layer-wise decay rate is 0.75. We train the network with the 1×\times schedule: 12 epochs with the learning rate decayed by 10×\times at epochs 9 and 11. We do not use multi-scale testing. The Mask R-CNN implementation follows MMDetection .

A.2 More Results for Larger Models and Longer Pre-training Epochs

SdAE can also perform well with only 300 epochs pre-training on a larger model scale such as ViT-L. We study the fine-tuning on the ILSVRC-2012 ImageNet dataset with 1k classes and 1.3M images. For a fair comparison, we directly follow most of the hyperparameters of MAE in our fine-tuning experiments. All reported experimental results are only fine-tuning for 50 epochs.

As shown in Table 6, compared with the models trained by random initialization (train from scratch), our pre-trained SdAE significantly improves the performance. Specifically, vision transformers trained from scratch only achieve 82.6% top-1 accuracy with ViT-L. While our SdAE achieves 85.7%, demonstrating the effectiveness of pre-training with unlabeled data.

Compared with previous self-supervised methods for vision transformers, our proposed SdAE surpasses them on ImageNet fine-tuning by a large margin. For ViT-L, our SdAE outperforms MoCo v3 by 1.6% top-1 accuracy with the same number of training epochs and our SdAE outperforms MAE by 1.4% top-1 accuracy with the less number of training epochs. Besides, our SdAE outperforms BEiT by 0.5% top-1 accuracy, while BEiT requires an additional pre-trained codebook and longer training epochs. In addition, our SdAE outperforms iBOT by 0.7% top-1 accuracy. Moreover, compared to the 1600 epoch pre-trained ViT-L of MAE and MaskFeat, which requires very large computational costs, our SdAE can achieve comparable performance.

As shown in Table 7, although the cost of SdAE de facto surpasses MAE per epoch, it can speed up convergence and achieve comparable performance in fewer epochs. It is also more efficient than SimMIM that adopts the mask token for the encoder. For longer pre-training epochs, SdAE is faced with a little bit performance degradation on ImageNet fine-tuning. We speculate that this is due to the fact that the Vit-base capacity is relatively close to the performance upper bound of the MIM tasks. However, SdAE shows continuous performance enhancement with longer training epochs on ADE20K semantic segmentation and COCO object detection, which also surpasses other methods by a considerable margin.

A.3 Comparison of iBOT, data2vec, and SplitMask

Except for the comparison of typical generative-based self-supervised learning methods in Figure 5 such as BeiT , PeCo , MAE and CAE , we also provide the comparison of several recently proposed works in Figure 5.

iBOT is more likely a contrastive learning/instance discrimination-based method. iBOT needs careful parameter setting of multi-crop augmentation, which uses 10 local crops with local scale being (0.05,0.32) and global scale being (0.32,1.0). In addition, iBOT heavily depends on the contrastive loss. MIM without the class token contrastive loss leads to undesirable results of 9.5%9.5\% kNN accuracy and 29.8%29.8\% linear accuracy on iBOT, indicating that iBOT benefits more from contrastive structure than MIM to extract the visual semantics.

Data2vec also uses an EMA parameterization of the two-branch teacher-student distillation structure. However, data2vec does not consider the spatial redundancy existing in the network structure. Not only visible unmasked patches but also learned mask embedding tokens are fed into the student branch. Furthermore, the whole input image will be fed into the teacher branch, ignoring the reconstruction loss computed only on masked tokens, which will also increase computational costs. In addition, data2vec conducts an EMA update per iteration. The MIM output and reconstruction target will be very similar if the momentum coefficient is not extremely small. So the network is sharply sensitive to the momentum coefficient. So, data2vec needs precisely tuning on this coefficient where in ViT-L they need to first set momentum coefficient as 0.9998 for the first 800 epochs and then reset the learning rate schedule and the teacher weights to the student and continue for another 800 epochs with momentum coefficient as 0.9999.

SplitMask considers the redundancy of split tokens. However, SplitMask still needs an additional tokenizer to produce discrete latent representations to conduct MIM. Moreover, SplitMask does not consider using a two-branch network to distill the representation between split tokens but adds a pooling module to calculate the contrastive loss between global representations.

A.4 Ablations on the Depth of Decoder

The decoder of the autoencoder, which maps the latent representation back to the reconstruction space, plays an essential role in the masked image modeling task. In the language MLM, the decoder predicts missing words that contain rich semantic information so that the decoder can be trivial (an MLP) in BERT . However, in MAE , the decoder reconstructs the image pixels, which reconstructs the latent representations into the low-level pixel space. Thus, MAE requires a relatively powerful decoder. In comparison, our student network maps the latent representations to the high-level semantic features so that the decoder can be lighter. That is another potential advantage of SdAE. As shown in Figure 6, the experiment shows that the depth of the decoder has little impact on the performance. Specifically, even with two layers of the decoder transformer, our SdAE achieves 82.76% top-1 accuracy and 45.2 mAP on COCO object detection, which only suffers 0.07% and 0.5 mAP performance degradation compared with eight layers of the decoder transformer.

A.5 Visualization

To analyze, we visualize the self-attention map with 300-epoch pre-trained ViT-B/16 of both MAE and SdAE. We choose the class token as the query and visualize attention maps from different heads of the last layer with different colors, following iBOT . As shown in Figure 7, we indicate that SdAE shows the capability to learn high-level semantic features to separate different parts of objects. Compared with MAE, SdAE is able to learn more meaningful high semantic information.

Specifically, in the figure, we observe SdAE can distinguish the bird from the tree or distinguish the eyes and ears of the Iberian wolf. Moreover, SdAE can also focus on the discriminative details of the object (e.g., the skeleton of a hot air balloon and sailboat rope) without using the contrastive loss. For more complex scenes like spiders on the surface of complex texture SdAE is still able to distinguish subjects. This is because SdAE does not need to reconstruct every pixel, so it did not pay attention to useless details.

With only a simple normalized feature MSE loss, we can achieve similar behaviors with intricately designed instance discrimination methods such as DINO .

References