Masked Siamese Networks for Label-Efficient Learning
Mahmoud Assran, Mathilde Caron, Ishan Misra, Piotr Bojanowski, Florian Bordes, Pascal Vincent, Armand Joulin, Michael Rabbat, Nicolas Ballas
Introduction
Self-Supervised Learning (SSL) has emerged as an effective strategy for unsupervised learning of image representations, eliminating the need to manually annotate vast quantities of data. By training large models on unlabeled data, SSL aims to learn representations that can be effectively applied to a downstream prediction task with few labels (Chen & He, 2020).
One of the core ideas of SSL is to remove a portion of the input and learn to predict the removed content (Pathak et al., 2016). Auto-regressive models and denoising auto-encoders instantiate this principle in vision by predicting the missing parts at the pixel or token level (Chen et al., 2020a; Vincent et al., 2010; He et al., 2021; Bao et al., 2021; Baevski et al., 2022). Masked auto-encoders in particular, which learn representations by reconstructing randomly masked patches from an input, have been successfully applied in vision (He et al., 2021; Xie et al., 2021; Wei et al., 2021; Bao et al., 2021). However, optimizing a reconstruction loss requires modelling low-level image details that are not necessary for classification tasks involving semantic abstraction. Thus, the resulting representations often need to be fine-tuned for semantic recognition tasks which can lead to overfitting in low-shot settings. Nevertheless, masked auto-encoders have enabled the training of large-scale models and demonstrated state-of-the-art performance when fine-tuning on large labeled datasets, with millions of labels (Bao et al., 2021; He et al., 2021; Xie et al., 2021; Baevski et al., 2022).
Joint-embedding architectures, on the other hand, avoid reconstruction. Approaches such as Siamese Networks (He et al., 2019; Caron et al., 2020; Chen & He, 2020; Grill et al., 2020; Caron et al., 2021; Zbontar et al., 2021; Bardes et al., 2021) learn a representation by training an encoder network to produce similar embeddings for two different views of the same image (Bromley et al., 1993; Dosovitskiy et al., 2014). Here the views are typically constructed by applying different image transforms — such as random scaling, cropping, and color jitter — to the input (Wu et al., 2018; Misra & van der Maaten, 2020). The inductive bias introduced by this invariance-based pre-training typically produces strong off-the-shelf representations of a high semantic level (Caron et al., 2021) but often disregards rich local structure that can be helpful to model.
In this work, we propose Masked Siamese Networks (MSNs), a self-supervised learning framework that leverages the idea of mask-denoising while avoiding pixel and token-level reconstruction. Figure 3 shows a schematic of the method. Given two views of an image, MSN randomly masks patches from one view while leaving the other view unchanged. The objective is to train a neural network encoder, parametrized with a vision transformer (ViT) (Dosovitskiy et al., 2020), to output similar embeddings for the two views. In this procedure, MSN does not predict the masked patches at the input level, but rather performs the denoising step implicitly at the representation level by ensuring that the representation of the masked input matches the representation of the unmasked one. Figure 2 qualitatively demonstrates the effectiveness of the MSN denoising process.
Empirically, we demonstrate that MSNs learn strong off-the-shelf representations that excel at low-shot prediction (cf. Figure 1). In particular, MSN achieves good classification performance using fewer labels than current mask-based auto-encoders (He et al., 2021; Xie et al., 2019). In the standard 1% ImageNet low-shot classification task, an MSN-trained ViT-B/4 (using a patch size of x pixels) achieves 75.7% top-1 accuracy, outperforming the previous 800M parameter state-of-the-art convolutional network (Chen et al., 2020c) while using nearly fewer parameters (cf. Figure 1(a)).
Since a good representation should not need many examples to learn about a concept (Goyal et al., 2019), we also consider a more challenging evaluation benchmark for label-efficient low-shot classification (Sohn et al., 2020; Lucas et al., 2021), using from 1 labeled image per class up to 5 images per class (cf. Table 2). MSN also achieves state-of-the-art performance in that regime. For instance, with only 5 labeled images per class, we can pre-train a ViT-L/7 with MSN on ImageNet-1K to achieve 72.1% top-1 accuracy surpassing the previous state-of-the-art method, DINO (Caron et al., 2021), by 8% top-1.
Similar to masked auto-encoders, MSNs also exhibit good computational scaling since only the unmasked patches are processed by the ViT encoder. For example, by randomly masking 70% of the patches, MSN uses half the computation and memory compared to an unmasked joint-embedding baseline. In practice, we pre-train a ViT-L/7 on as few as 18 AWS p4d-24xlarge machines. Without masking, the same job requires over 42 machines.
Finally, we also show that MSNs are competitive with prior works on other self-supervised benchmarks that use many labels for evaluation (e.g., fine-tuning, linear-evaluation, transfer learning).
Prerequisites
Consider a large collection of unlabeled images, , and a small dataset of annotated images, , with . Here, the images in may overlap with the images in the dataset . Our goal is to learn image representations by first pre-training on and then adapting the representation to the supervised task using .
The goal of siamese networks (Becker & Hinton, 1992; Bromley et al., 1993), as they are used in self-supervised learning, is to learn an encoder that produces similar image embeddings for two views of an image. Specifically, given an encoder and two views and of an image, the encoder independently processes each view and outputs representations and respectively, referred to as the anchor representation and the target representation. The objective of siamese networks is to learn an encoder that is not sensitive to differences between views, so the representations and should match. In practice, the encoder is usually parameterized as a deep neural network with learnable parameters .
The main challenge with siamese architectures is to prevent representation collapse in which the encoder produces a constant image embedding regardless of the input. Several approaches have been investigated in the literature. Contrastive losses explicitly push away embeddings of different images (Bromley et al., 1993; He et al., 2019; Chen & He, 2020). Information maximization approaches try to maximize the entropy of the average prediction (Caron et al., 2021; Assran et al., 2021) or spread out the embeddings uniformly on the surface of a sphere (Caron et al., 2020). Asymmetric approaches rely on an asymmetric architectural choice such as stop-gradient operations and a momentum encoder (Chen & He, 2020; Grill et al., 2020) to prevent collapse. Other approaches try to decorrelate the vector components of the embeddings to minimize redundancy across samples (Zbontar et al., 2021; Bardes et al., 2021).
We use a standard Vision Transformer (ViT) architecture (Dosovitskiy et al., 2020) as the encoder. Vision Transformers first extract a sequence of non-overlapping patches of resolution from an image. Next, they apply a linear layer to extract patch tokens, and subsequently add learnable positional embeddings to them. An extra learnable [CLS] token is added to the sequence. This token aims to aggregate information from the full sequence of patches (Dosovitskiy et al., 2020; Caron et al., 2021). The sequence of tokens is then fed to a stack of Transformer layers (Vaswani et al., 2017). A Transformer layer is composed of a self-attention (Vaswani et al., 2017) and a fully-connected layer with skip connections (He et al., 2016). Self-attention uses an attention mechanism (Bahdanau et al., 2014) applied to the entire sequence of elements to update the representation. The output representation associated to the [CLS] token is used as the output of the encoder.
Masked Siamese Networks
We now describe the proposed Masked Siamese Network (MSN) training procedure, which combines invariance-based pre-training with mask denoising; see Figure 3 for a schematic. MSNs first use random data augmentations to generate two views of an image, referred to as the anchor view and the target view. Subsequently, a random mask is applied to the anchor view, while the target view is left unchanged. Similar to clustering-based SSL approaches (Caron et al., 2020; 2021; Assran et al., 2021), learning occurs by computing a soft-distribution over a set of prototypes for both the anchor and target views. The objective is then to assign the representation of the masked anchor view to the same prototypes as the representation of the unmasked target view. We use a standard cross-entropy loss to optimize this criterion.
In contrast to previous work on masked image modelling, the mask-denoising process in MSN is discriminative, rather than generative (He et al., 2021; Xie et al., 2021; Wei et al., 2021; Bao et al., 2021; Zhou et al., 2021). MSN architectures do not directly predict pixel values (or tokens) for the masked patches. Instead, the loss is applied directly to the output corresponding to the [CLS] token of the encoder.
In each iteration of pre-training, we sample a mini-batch of images. For an index , let denote the image in the mini-batch. For each image , we first apply a random set of data augmentations to generate a target view, denoted , and anchor views, denoted .
Next, we “patchify” each view by converting it into a sequence of non-overlapping patches. After patchifying the anchor view , we also apply the additional step of masking by randomly dropping some of the patches. We denote by the sequence of masked anchor patches, and by the sequence of unmasked target patches. Because of masking, the anchor sequence can have a different length than the patchified target sequence , even if both image views originally have the same resolution.
We investigate two strategies for masking the anchor views, Random Masking and Focal Masking, which are depicted in Figure 4. When applying Random Masking, we randomly drop potentially non-contiguous patches across the sequence. Conversely, when applying Focal Masking, we randomly select a local continuous block of the anchor view and drop all the patches around it.
where is a temperature. Similarly, for each target representation , we generate a prediction by measuring the cosine similarity to the same prototypes matrix . When computing the target predictions, we also use a temperature parameter . Note, we always choose to encourage sharper target predictions, which implicitly guides the model to produce confident low entropy anchor predictions. As shown in Appendix B, target sharpening coupled with with other regularization like mean-entropy maximization (see below) is provably sufficient to eliminate collapsing solutions in the MSN framework. Empirically, we have observed that training without sharpening can result in collapsing solutions.
As previously mentioned, to train the encoder, we penalize when the anchor prediction is different from the target prediction . We enforce this criterion using a standard cross-entropy loss .
We also incorporate the mean entropy maximization (me-max) regularizer, also used in (Assran et al., 2021; Joulin & Bach, 2012), to encourage the model to utilize the full set of prototypes. Denote the average prediction across all the anchor views by
The me-max regularizer simply seeks to maximize the entropy of , denoted , or equivalently, minimize the negative entropy of . Thus, the overall objective to be minimized when training the encoder parameters and prototypes is
where controls the weight of the me-max regularization. Note that when training, we only compute gradients with respect to the anchor predictions , not the target predictions .
Related Work
Unsupervised pre-training for vision has seen rapid progress with the development of view-invariant representation learning and joint embedding architectures (Wu et al., 2018; He et al., 2019; Chen & He, 2020; Grill et al., 2020; Caron et al., 2021; Bardes et al., 2021). Most similar to our approach is DINO (Caron et al., 2021) which leverages a Siamese Network with a cross-entropy loss and a momentum encoder. DINO also uses multi-crop training, which is a form of focal masking, but it requires an unmasked anchor view during training. MSN can be seen as a generalization of DINO, leveraging both random and focal masking without requiring any unmasked anchor views. Since the cross-entropy loss in equation equation 1 is only differentiated with respect to the anchor predictions, not the target, MSN only backpropagates through the anchor network and only needs to store the activation associated with the masked view. MSN therefore reduces the computational and memory requirements. MSN also differs from DINO in its mechanism for preventing representation collapse (entropy maximization as opposed to centering and sharpening). Our empirical results show that MSN compares favourably to DINO across various degrees of supervision for the downstream task.
A prominent line of work in SSL is to remove a portion of the input and learn to reconstruct the removed content (Devlin et al., 2018). For example, in the field of image recognition, some works have proposed to predict augmented image channels (Zhang et al., 2017b), which can be regarded as a form of image colorization (Zhang et al., 2016; Larsson et al., 2016; 2017). Other approaches propose to remove and learn to regress entire image regions: the seminal Context Encoders of Pathak et al. (2016) train a network to generate missing image patches based on their surroundings. Recent works revisit this idea and investigate the pre-training of ViTs with masked auto-encoders (Chen et al., 2020a; He et al., 2021; Xie et al., 2021; Wei et al., 2021; Bao et al., 2021). These approaches corrupt images with mask-noise and predict missing input values at the pixel level (Dosovitskiy et al., 2020; He et al., 2021; Xie et al., 2019) or using a tokenizer (Bao et al., 2021; Wei et al., 2021). Our approach does not predict the missing value at the input level, but instead performs the denoising step implicitly by ensuring that the global representation of the noisy input matches that of the uncorrupted input.
Some recent approaches have started to explore the combination of joint-embedding architectures and denoising pre-training tasks (El-Nouby et al., 2021; Baevski et al., 2022; Zhou et al., 2021). Those approaches mask an image by replacing the masked patches with a learnable mask token, and output a single vector for each masked patch. The objective is then to directly match each computed patch vector to the equivalent patch token extracted from a target encoder. In addition to the patch-level loss, iBOT (Zhou et al., 2021) and SplitMask (El-Nouby et al., 2021) apply a joint-embedding loss to an output representing the global sequence (either the [CLS] token or a global average pool of the patch vectors). SplitMask shows that by using a patch-level loss, you can reduce the amount of unlabeled pre-training data. In contrast, we focus on reducing the amount of labeled data available for the downstream prediction task. Data2Vec (Baevski et al., 2022) demonstrates that this approach is suitable for multiple modalities such as vision, speech and text. Different from these approaches, we only match the view representations globally and do not consider a patch level loss. Consequently, we can completely ignore the masked patches, significantly reducing the computational and memory requirements. For example, when training our largest model, a ViT-L/7, we mask over 70% of the input patches, and reduce memory and computational overhead by half.
Results
We evaluate MSN representations learned on the ImageNet-1K dataset (Russakovsky et al., 2015). We first consider low-shot evaluation on ImageNet-1K using as few as 1–5 images per class. We also compare with the state-of-the-art in settings where more supervision is available and investigate transfer-learning performance. Finally, we conduct ablation experiments with MSN. By default, we pre-train with a batch-size of 1024 images, generating several anchor views from each image: 1 view with a random mask, and 10 views with focal masks. We find that the optimal masking ratio is model-dependent, with larger models benefiting from more aggressive patch dropping. We describe MSN implementation details in Appendix A.
The premise of SSL is to learn representations on unlabeled data that can be effectively applied to prediction tasks with few labels (Chen et al., 2020c). In this section we explore the performance of self-supervised approaches when very few labeled examples are available.
We first evaluate the classification performance of unsupervised models that have been pre-trained on ImageNet-1K, by using 1, 2, and 5 labeled images per class for supervised evaluation. We compare MSN to the joint-embedding approach, DINO (Caron et al., 2021), the auto-encoding approach, MAE (He et al., 2021), and the hybrid approach, iBOT (Zhou et al., 2021), which combines a joint-embedding architecture with a token-based patch-level loss. We download the official released models of each related approach for evaluation.
To adapt the joint-embeddings models to the supervised task, we freeze the weights of the pre-trained model and train a linear classifier on top using 1, 2 or 5 labeled samples (see Appendix A). For MAE, we rely on partial fine-tuning (He et al., 2021), except for the 1 image per class setting, and all results with the ViT-H/14 architecture, which use a linear classifier. Partial fine-tuning corresponds to fine-tuning the last block of the pre-trained model along with a linear head. MAE benefits from partial fine-tuning, but for sufficiently large models, such as the ViT-H/14, this leads to significant overfitting in the low-shot regime. We compare both protocols in more detail in Appendix C.
Table 1 reports the extreme low-shot evaluation results. MSN outperforms the other representation learning approaches across all levels of supervision. Moreover, the improvement offered by MSN increases as the amount of available labeled data is decreased. The performance of MSN also benefits from increased model size — settings with less labeled data appear to benefit more from increased model depth and smaller patch sizes.
We also observe that joint-embedding approaches appear to be more robust to the limited availability of downstream supervision than reconstruction-based auto-encoding approaches. To explain this observation, we refer to the Masked Auto-Encoders paper (He et al., 2021) which conjectures that using a pixel reconstruction loss results in encoder representations of a lower semantic level than other methods. Conversely, the inductive bias introduced by invariance-based pre-training appears to be helpful in the low-shot regime.
Table 2 reports a comparison on the 1% ImageNet-1K task, which is a standard benchmark for low-shot evaluation of self-supervised models (Chen et al., 2020b). For reference, the best reported result in the literature on 1% labeled data is 76.6%, achieved with a multi-stage semi-supervised pipeline, i.e., self-distilling from a fine-tuned ResNet-152 with 3 wider channels and selective kernels (Chen et al., 2020c). Here we focus on comparing to other models trained in a self-supervised setting. Our best MSN model using a ViT-B/4 achieves 75.7% top 1 accuracy, surpassing the previous 800M parameter state-of-the-art convolutional network (Chen et al., 2020c) while using significantly fewer parameters and no fine-tuning. When focusing the comparison on similar architectures (models with similar FLOP counts), MSN also consistently improves upon previous approaches.
2 Linear Evaluation and Fine-tuning
In this section we compare with the state-of-the-art on standard evaluation benchmarks where more supervised samples are available to adapt the representation. We use the full ImageNet-1K training images with 1.28M labels.
We evaluate self-supervised pretrained models by freezing their weights and training a linear classifier. Table 3 reports the linear evaluation results on ImageNet-1K. We observe that MSN performs competitively with the state-of-the-art. The best MSN model achieves 80.7% top-1 accuracy.
In this evaluation setting, we finetune all the weights of the self-supervised model using all the labels from the ImageNet-1K training set. We focus on the ViT-B/16 architecture. We adopt the same fine-tuning protocol as (Bao et al., 2021), and provide the details in Appendix A. Table 4 reports the comparison with fine-tuning evaluation using 100% labels on ImageNet-1K. MSN is competitive with joint-embedding approaches, such as DINO, and generative auto-encoding approaches, such as MAE.
3 Transfer Learning
We also report transfer learning experiments on the CIFAR10, CIFAR100 and iNaturalist datasets in Tables 5 and 6 when using a self-supervised ViT-B/16 pre-trained on ImageNet-1K. Across all tasks and various levels of supervision MSN either outperforms or achieves similar results to DINO pre-training. Recall that MSN pre-training is also less computationally expensive than DINO pre-training due to the anchor masking.
4 Ablations
We now conduct a series of experiments to gain insights into the important design decisions used in MSN such as the masking strategy and the data augmentation strategy. We measure the accuracy of the models by training a logistic regression classifier on the frozen trunk using 1% of ImageNet-1K labels (13 imgs/class).
In MSN we apply both random and focal masking to the anchor views. Focal masking corresponds to selecting a small crop from the anchor view. Random masking corresponds to randomly dropping potentially non-contiguous patches from the anchor view.
Table 7 reports the effect on low-shot evaluation when using a) No Masking, b) Focal Masking, c) Random Masking, or d) Random and Focal Masking. Applying a random mask to the anchor view is always better than applying no mask. By contrast, applying only a focal mask degrades the performance, which highlights the importance of maintaining a global view during pre-training. By combining both random and focal masking strategies,we obtain the strongest performance.
Here we explore the relationship between the optimal masking ratio and the model size. Table 8 reports the low-shot learning performance for various random masking ratios as we increase the model size.Note that the performance of the ViT-S/16 can be improved by removing the Sinkhorn normalization, as we do in Table 2, however for consistency of evaluation with other models, we keep it in for this this ablation.
When increasing the model size, we find that increasing the masking ratio (dropping more patches) is helpful for improving low-shot performance. We also find that the ViT-L/16 runs with weak masking are unstable, while the runs with more aggressive masking are quite stable. However, we do not have sufficient evidence to claim that increasing the masking ratio always improves the stability of large ViT pre-training.
We explore the importance of data-augmentation invariance for low-shot learning. We pretrain a ViT-B/16 with MSN, where the teacher and anchor networks either share the input image view or use different input views; in both cases, the anchor view is always masked. The views are constructed by applying random ColorJitter, Crop, Horizontal Flips, and GaussianBlur to the input image.
Table 9 reports top-1 accuracy when evaluating with 1% of ImageNet-1K labels. Sharing the view leads to a top-1 accuracy of ; MSN finds a shortcut solution relying on color statistics. Using different colors in the input views resolves this pathological behaviour and achieves a top-1 of . Further applying the geometric data-augmentations independently to the two views (as opposed to sharing views) further improves the performance to , showing the importance of learning view-invariant representations in the low-shot setting.
We look at the effect of the random masking ratio, i.e., the fraction of dropped patches from the global anchor view, on the computational requirements of large model pre-training. In each iteration we also generate 10 focal views (small crops) of each input image; the random masking ratio has no impact on these views.
Table 10 reports the memory consumption and throughput (imgs/s) of a ViT-L/7 model on a single AWS p4d-24xlarge machine using a batch-size of 2 images per GPU. As expected, using more aggressive masking of the global view progressively reduces device memory utilization and speeds up training. For example, by randomly masking 70% of the patches, we can use MSN to pre-train a full-precision ViT-Large with a patch-size of on as few as 18 AWS p4d-24xlarge machines. Without masking, the same job requires over 42 machines when using the default batch-size of 1024 images.
Conclusion
We propose Masked Siamese Networks (MSNs), a self-supervised learning framework that leverages the idea of mask-denoising while avoiding pixel and token-level reconstruction. We demonstrate empirically that MSNs learn strong off-the-shelf representations that excel at label-efficient learning, while simultaneously improving the scalability of joint-embedding architectures. By relying on view-invariant representation learning, MSN does require the specification of data transformations, and it may be that the optimal transformations and invariances are dataset and task dependant. In future work, we plan to explore more flexible mechanisms to learn those transformations and also explore the use of equivariant representations.
References
Appendix A Implementations Details
In this appendix section we provide the implementation details for MSN pre-training and evaluation.
We adopt similar hyper-parameter settings that have previously been reported in the self-supervised literature for training Vision Transformers (Caron et al., 2021; Chen & He, 2020). Specifically, for pre-training, we use the AdamW optimizer (Loshchilov & Hutter, 2017) with a batch-size of 1024. We linearly warm up the learning-rate from to during the first 15 epochs, and decay it following a cosine schedule thereafter. To construct the different image views, we apply the SimCLR data augmentations of Chen et al. (2020b) to each sampled image; namely random crop, horizontal flip, color distortion, and Gaussian blur. For each sampled image, we generate one large anchor view of size pixels, and apply a random mask with a pre-specified masking ratio (0.15 for the ViT-S/16, 0.3 for the ViT-B/16 and ViT-B/8, and 0.7 for the ViT-L/7 and the ViT-B/4). For each sampled image, we also generate 10 small focal anchor views of size pixels. We use a temperature of for the anchor network, and a temperature of for the target network. Following the DINO method of Caron et al. (2021), we update the target network via an exponential moving average of the anchor network with a momentum value of , and linearly increase this value to by the end of training. Similarly, following Caron et al. (2021), weight decay is set to and increased to throughout training via a cosine schedule. By default, we set the me-max regularization weight to and apply Sinkhorn normalization to the targets (Caron et al., 2020) to avoid having to tune the me-max regularization weight; however, in general, we observe stronger MSN performance when omitting Sinkhorn normalization (see Appendix C). We train with a 3-layer projection head with output dimension 256 and batch-normalization at the input and hidden layers, and use 1024 prototypes of dimension 256. We observe that using more prototypes has little effect on training, but using too few prototypes can hurt performance (see Appendix C). We discard the projection head during evaluation, and always use the representations computed from the output of the target encoder trunk for evaluation.
A.2 Low-Shot Evaluation
To avoid overfitting, we freeze the weights of the pre-trained model and train a linear classifier on top using 1, 2 or 5 labeled samples per class. Specifically, we take a single center crop of each labeled image, extract its representation using the pre-trained model, and then train a classifier on these representations using L2-regularized logistic regression. Following (Caron et al., 2021), we use the cyanure package (Mairal, 2019) to run logistic regression on the extracted representations. This objective is smooth and strongly-convex (i.e., has a unique minimizer) and can therefore be efficiently solved for using the cyanure python numerical solver on a single CPU core. All low-shot evaluations (including the 1% ImageNet-1K evaluation) are computed with this procedure, except for models pre-trained using MAE (He et al., 2021), which benefit from using partial fine-tuning (He et al., 2021).
Partial fine-tuning corresponds to fine-tuning the last block of the pre-trained model along with a linear head. MAE benefits from partial fine-tuning, but for sufficiently large models, such as the ViT-H/14, this leads to significant overfitting in the low-shot regime. Our results in Table 2 and Figure 1 report the best performance across evaluation methods for MAE. In particular, all the MAE results are obtained via partial fine-tuning, except for the 1 image per class setting, and all results with the ViT-H/14 architecture, which use a linear head. We compare both protocols in more detail in Appendix C.
A.3 Linear Evaluation
For linear evaluation, we use a similar procedure as He et al. (2021). Specifically, we use a large batch-size of 16,384 images and train a linear classifier for 100 epochs using a learning rate of , and decay it following a cosine schedule. We only apply basic data augmentations; namely, random resized crops to a resolution of pixels, and random horizontal flips. We also L2-normalize the representations before feeding them into the linear classifier, and optimize the classifier weights using SGD with Nesterov momentum. We do not apply any weight-decay and do not use any warmup.
A.4 Fine-Tuning Evaluation
We follow the common practice for fine-tuning SSL pre-trained ViT models. Specifically, we follow the setup of (Touvron et al., 2021; Bao et al., 2021; He et al., 2021). We fine-tune a pre-trained ViT model for 100 epochs on the full supervised ImageNet-1K training data set using the AdamW (Loshchilov & Hutter, 2017) optimizer. We use a batch size of 1024 with a learning rate of . The learning rate is linearly warmed-up during the first 5 epochs and decayed with a cosine schedule thereafter. A layer-wise decay of is also applied, along with the data augmentations defined by RandAugment(, ) (Cubuk et al., 2019). We additionally use label smoothing set to , mixup (Zhang et al., 2017a) set to , cutmix (Yun et al., 2019) set to , and drop path set to .
A.5 Transfer Learning
When performing linear evaluation for transfer learning, we freeze the weights of the ImageNet-1K pre-trained model and optimize a linear classifier on top. We resize each downstream image to pixels, and take a single center crop of size pixels. Next, we extract a representation of each image using the pre-trained model, and subsequently train a classifier on top using L2-regularized logistic regression.
A.5.2 Fine Tuning
When performing end-to-end fine-tuning for transfer learning, we follow the protocol of DeiT and DINO (Touvron et al., 2021; Caron et al., 2021). Models transferred to CIFAR10 and CIFAR100 are fine-tuned for 1000 epochs using a batch size of 768 and a learning rate of . Models transferred to iNat18 and iNat19 models are fine-tuned for 300 epochs using a batch size of 1024 and a learning of . All transfer fine-tuning experiments use the data augmentations defined by RandAugment(, ) (Cubuk et al., 2019). We also use label smoothing set to , mixup (Zhang et al., 2017a) set to , cutmix (Yun et al., 2019) set to , and drop path set to . The learning rate is linearly warmed-up during the 5 first epochs and decayed with a cosine schedule thereafter.
Appendix B Theoretical Guarantees
In this section we describe how MSN pre-training provably avoids representation collapse.
Recall that in each iteration of pre-training, we sample a mini-batch of images, and generate anchor views of each image. Here we show that MSN is guaranteed to avoid the trivial collapse of representations under the following assumption.
The target is sharpened, such that it is not equal to the uniform distribution.
Suppose Assumption 1 holds. If is such that the representations collapse, i.e., for all and , then for all .
For L2-normalized representations and prototypes, the prediction corresponding to the view of the image in the mini-batch is given by
Case 2: The predictions are not equal to the uniform distribution, i.e., . In that case, we have that the average prediction across all the anchor views is also not equal to the uniform distribution; i.e., , and hence . ∎
Proposition 1 provides a theoretical guarantee that MSN is immune to the trivial collapse of representations. In short, the underlying principle is that entropy maximization encourages the anchor predictions to utilize the full set of prototypes, thereby preventing collapse to a non-uniform distribution, while target sharpening encourages the anchor predictions to be confident, thereby preventing collapse to the uniform distribution.
Note that the sharpening mechanism defined in Section 3 (i.e., applying a temperature in the target network softmax) may not always satisfy Assumption 1, unless one introduces a simple tie-breaking rule. In practice, such a rule is not necessary as the targets never become uniform (since we apply sharpening from the start of the training), although, it is important to use a sufficiently small temperature value in this case.
Appendix C Additional Ablations
By default, we set the me-max regularization weight to and apply Sinkhorn normalization on the targets to avoid having to tune the me-max regularization weight. However, we find that tuning the me-max regularization weight and omitting Sinkhorn normalization can result in better performance; cf. Table 11.
C.2 Number of Prototypes
By default we train with 1024 prototypes of dimension 256. In this section we explore the effect of the number of prototypes on low-shot performance. We observe that using more prototypes has little effect on training, but using too few prototypes can hurt performance; cf. Table 12.
C.3 Masked Auto-Encoder Partial Fine-Tuning
Here we explore the low-shot performance of MAE when relying on alternative evaluation strategies. He et al. (2021) conjecture that using pixel reconstruction in their MAE objective results in encoder representations of a lower semantic level than other methods, which may explain their difficulty in training a linear classifier on the frozen features. In Table 13 we explore the effect of partial fine-tuning on the low-shot performance of pre-trained MAE models. Partial fine-tuning corresponds to fine-tuning the last block of the pre-trained model along with a linear head on the available labeled samples. As observed in (He et al., 2021), MAE benefits from partial fine-tuning. However, for sufficiently large models, such as the ViT-H/14, this leads to significant overfitting in the low-shot regime, where one must instead resort to linear evaluation. We report the best numbers for MAE across the two low-shot adaptation strategies in Figure 1.
Appendix D MSN Representation Robustness
Next we report the performance of MSN-pre-trained models on datasets that have been developed to evaluate the robustness of models trained on the standard ImageNet training set. We consider four datasets: ImageNet-A (Hendrycks et al. (2021b))https://github.com/hendrycks/natural-adv-examples, ImageNet-R (Hendrycks et al. (2021a))https://github.com/hendrycks/imagenet-r, ImageNet-Sketch (Wang et al. (2019))https://github.com/HaohanWang/ImageNet-Sketch, and ImageNet-C (Hendrycks & Dietterich (2019))https://github.com/hendrycks/robustness.
Table 14 shows results for a ViT-B/16 pre-trained using MSN and fine-tuned using the protocol described in Appendix A. For comparison, we also report the performance of a fine-tuned ViT-B/16 pre-trained using MAE (He et al., 2021), along with a supervised ResNet50 baseline, which is available in the PyTorch Torchvision packagehttps://github.com/pytorch/vision. For ImageNet-A, -R, and -Sketch, we report top-1 accuracy on each provided validation set. For ImageNet-C, we use the mean Corruption Error metric proposed in (Hendrycks & Dietterich, 2019), where values are normalized by AlexNet performance on the same validation set.
In each case we find that the performance of an MSN-pretrained ViT-B/16 is comparable or better than that of an MAE-pretrained ViT-B/16. Note also, that larger MAE-pretrained models achieve stronger performance on all four datasets (He et al., 2021).
Appendix E MSN Invariance to Masking
The goal of MSN pretraining is to denoise the input images at the representation level by ensuring that the representation of a masked input matches the representation of the unmasked one. Here, we shows that MSN pretraining learns representations that are robust to patch masking.
In Table 15, we evaluate the performance of MSN and DINO when masking parts of an image during evaluation. Models are evaluated on 1% of ImageNet-1K using logistic regression on top of frozen features. The logistic regression classifier is trained using masked images, and then evaluated on the standard ImageNet-1K validation set using unmasked images.
If the MSN representations are robust to missing image patches, then a linear classifier should be able to identify generalizable features when training on the representations of masked images. On the other hand, if the representations output by the learned encoder are not robust to missing image patches, then a linear classifier would have difficulty finding generalizable features when training on the representations of masked images.
We observe that masked pre-training results in representations that are more robust to patch removal, suggesting that MSN is performing an image denoising at the representation level. Furthermore, models pre-trained with more aggressive masking exhibit this quality to a higher degree. For example, the low-shot accuracy of ViT-L/7 pre-trained with aggressive masking is almost unaffected when we remove 70% of the patches at test time; 75.1% top-1 without dropping patches during evaluation versus 74.9% top-1 when dropping 70% of the patches during evaluation.
We also report the average cosine distance between masked and unmasked representations of the same image in Table 16. As expected, the cosine similarity between masked and unmasked representations of the same image is higher when pre-training with MSN, supporting the observation that masked-pretraining results in representations that are more robust to patch-removal.
Appendix F Qualitative Analysis
We qualitatively investigate the properties of the MSN pre-trained representations. We follow the RCDM framework (Bordes et al., 2021) and train a conditional generative diffusion model, which maps a learned image representation back to pixel space. Specifically, RCDM takes as input random noise and the representation vector of an image computed by an SSL model (either an MSN pre-trained model or a DINO pre-trained model in this analysis), and aims to reconstruct the image as close as possible to the original one through a diffusion process.
By using RCDM to sample an image based on its SSL representation, we can visualize how different pre-training strategies affect the degree of information contained in the representation. Qualities that vary across RCDM samples represent information that is not contained in the pre-trained representation. Qualities that are semantically common across samples represent information contained in the representation.
We apply RCDM on top of either a DINO or MSN pre-trained ViT-B/8 encoder to generate images of resolution pixels. RCDM is trained using unmasked images processed with the ViT-B/8 encoder. We then use masked images from the validation set at sampling time.
In Figure 5, we generate samples for RCDM when masking 50% of the conditioning images. The first column depicts images from the ImageNet validation set. The second column depicts the same image, but with 50% of the patches masked. The representation of the masked image is used as conditioning for the RCDM diffusion model. The subsequent columns in Figure 5 show various images sampled from the conditioned RCDM diffusion model. We observe that the RCDM samples conditioned on the MSN representations (cf. Figure 5(a)) preserve the semantic category of the masked images, and remain visually close to the original image, despite the missing patches. By contrast, the samples generated by the RCDM diffusion model conditioned on the DINO representations (cf. Figure 5(b)) are more blurry and do not preserve as well the semantic category of the masked images.
Figure 6 depicts similar visualizations, but with 80% of the patches masked. In this case, even with 80% of the patches missing, samples generated by RCDM conditioned on MSN representations preserve some of the structure in original images (cf. Figure 6(a)). On the other hand, conditioning on DINO representations leads to almost uniform background generation (cf. Figure 6(b)).
F.2 MSN ViT-L/7 Visualizations
We apply RCDM on top of the MSN pre-trained ViT-L/7 encoder to generate images with a resolution of pixels. RCDM is trained using images with 70% of patches masked. We then use masked images from the validation set (with various masking ratios) at sampling time, see Figures 7, 8, and 9.
Visualizations show that MSN discards instance-specific information such as background, pose, and lighting, while retaining semantic information about the images, even when a large fraction of the patches are masked.