Self-Supervised Learning from Images with a Joint-Embedding Predictive Architecture

Mahmoud Assran, Quentin Duval, Ishan Misra, Piotr Bojanowski, Pascal Vincent, Michael Rabbat, Yann LeCun, Nicolas Ballas

Introduction

In computer vision, there are two common families of approaches for self-supervised learning from images: invariance-based methods assran2022masked; he2019moco; caron2020unsupervised; chen2020exploring; grill2020bootstrap; caron2021emerging; zbontar2021barlow; bardes2021vicreg; asano2019self and generative methods devlin2018bert; baevski2022data2vec; pathak2016context; he2021masked.

Invariance-based pretraining methods optimize an encoder to produce similar embeddings for two or more views of the same image bromley1993signature; chen2020simple, with image views typically constructed using a set of hand-crafted data augmentations, such as random scaling, cropping, and color jittering chen2020simple, amongst others grill2020bootstrap. These pretraining methods can produce representations of a high semantic level caron2021emerging; assran2022masked, but they also introduce strong biases that may be detrimental for certain downstream tasks or even for pretraining tasks with different data distributions assran2022hidden. Often, it is unclear how to generalize these biases for tasks requiring different levels of abstraction. For example, image classification and instance segmentation do not require the same invariances bardes2022vicregl. Additionally, it is not straightforward to generalize these image-specific augmentations to other modalities such as audio.

Cognitive learning theories have suggested that a driving mechanism behind representation learning in biological systems is the adaptation of an internal model to predict sensory input responses rao1999predictive; friston2005theory. This idea is at the core of self-supervised generative methods, which remove or corrupt portions of the input and learn to predict the corrupted content vincent2010stacked; pathak2016context; he2021masked; bao2021beit; xie2021simmim; wei2021masked. In particular, mask-denoising approaches learn representations by reconstructing randomly masked patches from an input, either at the pixel or token level. Masked pretraining tasks require less prior knowledge than view-invariance approaches and easily generalize beyond the image modality baevski2022data2vec. However, the resulting representations are typically of a lower semantic level and underperform invariance-based pretraining in off-the-shelf evaluations (e.g., linear-probing) and in transfer settings with limited supervision for semantic classification tasks assran2022masked. Consequently, a more involved adaptation mechanism (e.g., end-to-end fine-tuning) is required to reap the full advantage of these methods.

In this work, we explore how to improve the semantic level of self-supervised representations without using extra prior knowledge encoded through image transformations. To that end, we introduce a joint-embedding predictive architecture lecun2022path for images (I-JEPA). An illustration of the method is provided in Figure 3. The idea behind I-JEPA is to predict missing information in an abstract representation space; e.g., given a single context block, predict the representations of various target blocks in the same image, where target representations are computed by a learned target-encoder network.

Compared to generative methods that predict in pixel/token space, I-JEPA makes use of abstract prediction targets for which unnecessary pixel-level details are potentially eliminated, thereby leading the model to learn more semantic features. Another core design choice to guide I-JEPA towards producing semantic representations is the proposed multi-block masking strategy. Specifically, we demonstrate the importance of predicting sufficiently large target blocks in the image, using an informative (spatially distributed) context block.

Through an extensive empirical evaluation, we demonstrate that:

I-JEPA learns strong off-the-shelf representations without the use of hand-crafted view augmentations (cf. Fig.1). I-JEPA outperforms pixel-reconstruction methods such as MAE he2021masked on ImageNet-1K linear probing, semi-supervised 1% ImageNet-1K, and semantic transfer tasks.

I-JEPA is competitive with view-invariant pretraining approaches on semantic tasks and achieves better performance on low-level visions tasks such as object counting and depth prediction (Sections 5 and 6). By using a simpler model with less rigid inductive bias, I-JEPA is applicable to a wider set of tasks.

I-JEPA is also scalable and efficient (Section 7). Pre-training a ViT-H/14 on ImageNet requires less than 1200 GPU hours, which is over 2.5×2.5\times faster than a ViT-S/16 pretrained with iBOTzhou2021ibotyes and over 10×10\times more efficient than a ViT-H/14 pretrained with MAE. Predicting in representation space significantly reduces the total computation needed for self-supervised pretraining.

Background

Self-supervised learning is an approach to representation learning in which a system learns to capture the relationships between its inputs. This objective can be readily described using the framework of Energy-Based Models (EBMs) lecun2006tutorial in which the self-supervised objective is to assign a high energy to incompatible inputs, and to assign a low energy to compatible inputs. Many existing generative and non-generative approaches to self-supervised learning can indeed be cast in this framework; see Figure 2.

Invariance-based pretraining can be cast in the framework of EBMs using a Joint-Embedding Architecture (JEA), which learns to output similar embeddings for compatible inputs, x,y{\bm{x}},{\bm{y}}, and dissimilar embeddings for incompatible inputs; see Figure 2(a). In the context of image-based pretraining, compatible x,y{\bm{x}},{\bm{y}} pairs are typically constructed by randomly applying hand-crafted data augmentations to the same input image chen2020simple.

The main challenge with JEAs is representation collapse, wherein the energy landscape is flat (i.e., the encoder produces a constant output regardless of the input). During the past few years, several approaches have been investigated to prevent representation collapse, such as contrastive losses that explicitly push apart embeddings of negative examples bromley1993signature; he2019moco; chen2020exploring, non-contrastive losses that minimize the informational redundancy across embeddings zbontar2021barlow; bardes2021vicreg, and clustering-based approaches that maximize the entropy of the average embedding caron2021emerging; assran2021semi; assran2022masked. There are also heuristic approaches that leverage an asymmetric architectural design between the xx-encoder and yy-encoder to avoid collapse chen2020exploring; grill2020bootstrap; baevski2022data2vec.

Generative Architectures.

Reconstruction-based methods for self-supervised learning can also be cast in the framework of EBMs using Generative Architectures; see Figure 2(b). Generative Architectures learn to directly reconstruct a signal y{\bm{y}} from a compatible signal x{\bm{x}}, using a decoder network that is conditioned on an additional (possibly latent) variable z{\bm{z}} to facilitate reconstruction. In the context of image-based pretraining, one common approach in computer vision is to produce compatible x,y{\bm{x}},{\bm{y}} pairs using masking he2016deep; bao2021beit where x{\bm{x}} is a copy of the image y{\bm{y}}, but with some of the patches masked. The conditioning variable z{\bm{z}} then corresponds to a set of (possibly learnable) mask and position tokens, that specifies to the decoder which image patches to reconstruct. Representation collapse is not a concern with these architectures as long as the informational capacity of z{\bm{z}} is low compared to the signal y{\bm{y}}.

Joint-Embedding Predictive Architectures.

As shown in Figure 2(c), Joint-Embedding Predictive Architectures lecun2022path are conceptually similar to Generative Architectures; however, a key difference is that the loss function is applied in embedding space, not input space. JEPAs learn to predict the embeddings of a signal y{\bm{y}} from a compatible signal x{\bm{x}}, using a predictor network that is conditioned on an additional (possibly latent) variable z{\bm{z}} to facilitate prediction. Our proposed I-JEPA provides an instantiation of this architecture in the context of images using masking; see Figure 3.

In contrast to Joint-Embedding Architectures, JEPAs do not seek representations invariant to a set of hand-crafted data augmentations, but instead seek representations that are predictive of each other when conditioned on additional information z{\bm{z}}. However, as with Joint-Embedding Architectures, representation collapse is also a concern with JEPAs; we leverage an asymmetric architecture between the x{\bm{x}}- and y{\bm{y}}-encoders to avoid representation collapse.

Method

We now describe the proposed Image-based Joint-Embedding Predictive Architecture (I-JEPA), illustrated in Figure 3. The overall objective is as follows: given a context block, predict the representations of various target blocks in the same image. We use a Vision Transformer touvron2021training; dosovitskiy2020image (ViT) architecture for the context-encoder, target-encoder, and predictor. A ViT is composed of a stack of transformer layers, each consisting of a self-attention vaswani2017attention operation followed by a fully-connected MLP. Our encoder/predictor architecture is reminiscent of the generative masked autoencoders (MAE) he2021masked method. However, one key difference is that the I-JEPA method is non-generative and the predictions are made in representation space.

We first describe how we produce the targets in the I-JEPA framework: in I-JEPA, the targets correspond to the representations of image blocks. Given an input image y{\bm{y}}, we convert it into a sequence of NN non-overlapping patches, and feed this through the target-encoder fθˉf_{\bar{\theta}} to obtain a corresponding patch-level representation sy={sy1,…,syN}{\bm{s}}_{y}=\{{\bm{s}}_{y_{1}},\dots,{\bm{s}}_{y_{N}}\} where syk{\bm{s}}_{y_{k}} is the representation associated with the kthk^{\text{th}} patch. To obtain the targets for our loss, we randomly sample MM (possibly overlapping) blocks from the target representations sy{\bm{s}}_{y}. We denote by BiB_{i} the mask corresponding of the ithi^{\text{th}} block and by sy(i)={syj}j∈Bi{\bm{s}}_{y}(i)=\{{\bm{s}}_{y_{j}}\}_{j\in B_{i}} its patch-level representation. Typically, we set MM equal to 4, and sample the blocks with a random aspect ratio in the range (0.75,1.5)(0.75,1.5) and random scale in the range (0.15,0.2)(0.15,0.2). Note that the target blocks are obtained by masking the output of the target-encoder, not the input. This distinction is crucial to ensure target representations of a high semantic level; see, e.g., baevski2022data2vec.

Context.

Recall, the goal behind I-JEPA is to predict the target block representations from a single context block. To obtain the context in I-JEPA, we first sample a single block x{\bm{x}} from the image with a random scale in the range (0.85,1.0)(0.85,1.0) and unit aspect ratio. We denote by BxB_{x} the mask associated with the context block x{\bm{x}}. Since the target blocks are sampled independently from the context block, there may be significant overlap. To ensure a non-trivial prediction task, we remove any overlapping regions from the context block. Figure 4 shows examples of various context and target blocks in practice. Next, the masked context block, x{\bm{x}}, is fed through the context encoder fθf_{\theta} to obtain a corresponding patch-level representation sx={sxj}j∈Bx{\bm{s}}_{x}=\{{\bm{s}}_{x_{j}}\}_{j\in B_{x}}.

Prediction.

Given the output of the context encoder, sx{\bm{s}}_{x}, we wish to predict the MM target block representations sy(1),…,sy(M){\bm{s}}_{y}(1),\ldots,{\bm{s}}_{y}(M). To that end, for a given target block sy(i){\bm{s}}_{y}(i) corresponding to a target mask BiB_{i}, the predictor gϕ(⋅,⋅)g_{\phi}(\cdot,\cdot) takes as input the output of the context encoder sx{\bm{s}}_{x} and a mask token for each patch we wish to predict, {mj}j∈Bi\{{\bm{m}}_{j}\}_{j\in B_{i}}, and outputs a patch-level prediction s^y(i)={s^yj}j∈Bi=gϕ(sx,{mj}j∈Bi)\hat{{\bm{s}}}_{y}(i)=\{\hat{{\bm{s}}}_{y_{j}}\}_{j\in B_{i}}=g_{\phi}({\bm{s}}_{x},\{{\bm{m}}_{j}\}_{j\in B_{i}}). The mask tokens are parameterized by a shared learnable vector with an added positional embedding. Since we wish to make predictions for MM target blocks, we apply our predictor MM times, each time conditioning on the mask tokens corresponding to the target-block locations we wish to predict, and obtain predictions s^y(1),…,s^y(M)\hat{{\bm{s}}}_{y}(1),\ldots,\hat{{\bm{s}}}_{y}(M).

Loss.

The loss is simply the average L2L_{2} distance between the predicted patch-level representations s^y(i)\hat{{\bm{s}}}_{y}(i) and the target patch-level representation sy(i){\bm{s}}_{y}(i); i.e.,

The parameters of the predictor, ϕ\phi, and context encoder, θ\theta, are learned through gradient-based optimization, while the parameters of the target encoder θˉ\bar{\theta} are updated via an exponential moving average of the context-encoder parameters. The use of an exponential moving average target-encoder has proven essential for training JEAs with Vision Transformers chen2021empirical; zhou2021ibotyes; caron2021emerging, we find the same to be true for I-JEPA.

Related Work

A long line of work has explored visual representation learning by predicting the values of missing or corrupted sensory inputs. Denoising autoencoders use random noise as input corruption vincent2010stacked. Context encoders regress an entire image region based on its surrounding pathak2016context. Other works cast image colorization as a denoising task zhang2016colorful; larsson2016learning; larsson2017colorization.

The idea of image denoising has recently been revisited in the context of masked image modelling he2021masked; xie2021simmim; bao2021beit, where a Vision Transformer dosovitskiy2020image is used to reconstruct missing input patches. The work on Masked Autoencoders (MAE) he2021masked proposed an efficient architecture that only requires the encoder to process visible image patches. By reconstructing missing patches in pixels space, MAE achieves strong performance when fine-tuned end-to-end on large labeled datasets and exhibits good scaling properties. BEiT bao2021beit predicts the value of missing patches in a tokenized space; specifically, tokenizing image patches using a frozen discreteVAE, which is trained on a dataset containing 250 million images ramesh2021zero. Yet, pixel-level pre-training has been shown to outperform BEiT for fine-tuning he2021masked. Another work, SimMIM xie2021simmim, explores reconstruction targets based on the classic Histogram of Gradients dalal2005histograms feature space, and demonstrates some advantage over pixel space reconstruction. Different from those works, our representation space is learned during training through a Joint-Embedding Predictive Architecture. Our goal is to learn semantic representations that do not require extensive fine-tuning on downstream tasks.

Closest to our work is data2vec baevski2022data2vec and Context Autoencoders chen2021empirical. The data2vec method learns to predict the representation of missing patches computed through an online target encoder; by avoiding handcrafted augmentations, the method can be applied to diverse modalities with promising results in vision, text and speech. Context Autoencoders use an encoder/decoder architecture optimized via the sum of a reconstruction loss and an alignment constraint, which enforces predictability of missing patches in representation space. Compared to these methods, I-JEPA exhibits significant improvements in computational efficiency and learns more semantic off-the-shelf representations. Concurrent to our work, data2vec-v2 baevski2022efficient explores efficient architectures for learning with various modalities.

We also compare I-JEPA with various methods based on joint-embedding architectures; e.g., DINO caron2021emerging, MSN assran2022masked and iBOT zhou2021ibotyes. Theses methods rely on hand-crafted data augmentations during pretraining to learn semantic image representations. The work on MSN assran2022masked, uses masking as an additional data-augmentation during pretraining, while iBOT combines a data2vec-style patch-level reconstruction loss with the DINO view-invariance loss. Common to these approaches is the need to process multiple user-generated views of each input image, thereby hindering scalability. By contrast, I-JEPA only requires processing a single view of each image. We find that a ViT-Huge/14 trained with I-JEPA requires less computational effort than a ViT-Small/16 trained with iBOT.

Image Classification

To demonstrate that I-JEPA learns high-level representations without relying on hand-crafted data-augmentations, we report results on various image classification tasks using the linear probing and partial fine-tuning protocols. In this section, we consider self-supervised models that have been pretrained on the ImageNet-1K dataset russakovsky2015imagenet. Pretraining and evaluation implementation details are described in the Appendix A. All I-JEPA models are trained at resolution 224×224224\times 224 pixels, unless stated otherwise.

Table 1 shows performance on the common ImageNet-1K linear-evaluation benchmark. After self-supervised pretraining, the model weights are frozen and a linear classifier is trained on top using the full ImageNet-1K training set. Compared to popular methods such as Masked Autoencoders (MAE) he2021masked, Context Autoencoders (CAE) chen2022context, and data2vec baevski2022data2vec, which also do not rely on extensive hand-crafted data-augmentations during pretraining, we see that I-JEPA significantly improves linear probing performance, while using less computational effort (see section 7). By leveraging the improved efficiency of I-JEPA, we can train larger models that outperform the best CAE model while using a fraction of the compute. I-JEPA also benefits from scale; in particular, a ViT-H/16 trained at resolution 448×448448\times 448 pixels matches the performance of view-invariant approaches such as iBOT zhou2021ibotyes, despite avoiding the use of hand-crafted data-augmentations.

Low-Shot ImageNet-1K.

Table 2 shows performance on the 1% ImageNet benchmark. Here the idea is to adapt the pretrained models for ImageNet classification using only 1% of the available ImageNet labels, corresponding to roughly 12 or 13 images per class. Models are adapted via fine-tuning or linear-probing, depending on whichever works best for each respective method. I-JEPA outperforms MAE while requiring less pretraining epochs when using a similar encoder architecture. I-JEPA, using a ViT-H/14 architecture, matches the performance of a ViT-L/16 pretrained with data2vec baevski2022data2vec, while using significantly less computational effort (see Section 7). By increasing the image input resolution, I-JEPA outperforms previous methods including joint-embedding methods that do leverage extra hand-crafted data-augmentations during pretraining, such as MSN assran2022masked, DINO caron2020unsupervised, and iBOT zhou2021ibotyes.

Transfer learning.

Table 3 shows performance on various downstream image classification tasks using a linear probe. I-JEPA significantly outperforms previous methods that do not use augmentations (MAE and data2vec), and decreases the gap with the best view-invariance-based methods, which leverage hand-crafted data augmentations during pretraining, even surpassing the popular DINO caron2021emerging on CIFAR100 and Place205 with a linear probe.

Local Prediction Tasks

As demonstrated in Section 5, I-JEPA learns semantic image representations that significantly improve the downstream image classification performance of previous methods, such as MAE and data2vec. Additionally, I-JEPA benefits from scale and can close the gap, and even surpass, view-invariance based methods that leverage extra hand-crafted data augmentations. In this section, we find that I-JEPA also learns local image features and surpasses view-invariance based methods on low-level and dense prediction tasks, such as object counting and depth prediction.

Table 4 shows performance on various low-level tasks using a linear probe. After pretraining, the encoder weights are frozen and a linear model is trained on top to perform object-counting and depth prediction on the Clevr dataset clevr. Compared to view-invariance methods such as DINO and iBOT, the I-JEPA method effectively captures low-level image features during pretraining and outperforms them in object counting (Clevr/Count) and (by a large margin) depth prediction (Clevr/Dist).

Scalability

I-JEPA is highly scalable compared to previous approaches. Figure 5 shows semi-supervised evaluation on 1% ImageNet-1K as a function of GPU hours. I-JEPA requires less compute than previous methods and achieves strong performance without relying on hand-crafted data-augmentations. Compared to reconstruction-based methods, such as MAE, which directly use pixels as targets, I-JEPA introduces extra overhead by computing targets in representation space (about 7% slower time per iteration). However, since I-JEPA converges in roughly 5×5\times fewer iterations, we still see significant compute savings in practice. Compared to view-invariance based methods, such as iBOT, which rely on hand-crafted data augmentations to create and process multiple views of each image, I-JEPA also runs significantly faster. In particular, a huge I-JEPA model (ViT-H/14) requires less compute than a small iBOT model (ViT-S/16).

Scaling data size.

We also find I-JEPA to benefit from pretraining with larger datasets. Table 5 shows transfer learning performance on semantic and low level tasks when increasing the size of the pretraining dataset (IN1K versus IN22K). Transfer learning performance on these conceptually different tasks improves when pretraining on a larger more diverse dataset.

Scaling model size.

Table 5 also shows that I-JEPA benefit from larger model size when pretraining on IN22K. Pretraining a ViT-G/16 significantly improves the downstream performances on image classification tasks such as Place205 and INat18 compared to a ViT-H/14 model, but does not improve performance on low-level downstream tasks — the ViT-G/16 uses larger input patches, which can be detrimental for the local prediction tasks.

Predictor Visualizations

The role of the predictor in I-JEPA is to take the output of the context encoder and, conditioned on positional mask tokens, to predict the representations of a target black at the location specified by the mask tokens. One natural question is whether the predictor conditioned on the positional mask tokens is learning to correctly capture positional uncertainty in the target. To qualitatively investigate this question, we visualize the outputs of the predictor. We use the following visualization approach to enable the research community to independently reproduce our findings. After pretraining, we freeze the context-encoder and predictor weights, and train a decoder following the RCDM framework bordes2022high to map the average-pool of the predictor outputs back to pixel space. Figure 6 shows decoder outputs for various random seeds. Qualities that are common across samples represent information that is contained in the average-pooled predictor representation. The I-JEPA predictor correctly captures positional uncertainty and produces high-level object parts with the correct pose (e.g., back of the bird and top of the car).

Ablations

Table 7 compares low-shot performance on 1% ImageNet-1K using a linear probe when the loss is computed in pixel-space versus representation space. We conjecture that a crucial component of I-JEPA is that the loss is computed entirely in representation space, thereby giving the target encoder the ability to produce abstract prediction targets, for which irrelevant pixel-level details are eliminated. From Table 7, it is clear that predicting in pixel-space leads to a significant degradation in the linear probing performance.

Masking strategy.

Table 6 compare our multi-block masking with other masking strategies such as rasterized masking, where the image is split into four large quadrants, and the goal is to use one quadrant as a context to predict the other three quadrants, and the traditional block and random masking typically used in reconstruction-based methods. In block masking, the target is a single image block and the context is the image complement. In random masking, the target is a set of random patches and the context is the image complement. Note that there is no overlap between the context and target blocks in all considered strategies. We find multi-block masking helpful for guiding I-JEPA to learning semantic representations. Additional ablations on multi-block masking can be found in Appendix C.

Conclusion

We proposed I-JEPA, a simple and efficient method for learning semantic image representations without relying on hand-crafted data augmentations. We show that by predicting in representation space, I-JEPA converges faster than pixel reconstruction methods and learns representations of a high semantic level. In contrast to view-invariance based methods, I-JEPA highlights a path for learning general representations with joint-embedding architectures, without relying on hand-crafted view augmentations.

References

Appendix A Implementation Details

For I-JEPA pretraining, we use Vision Transformer dosovitskiy2020image (ViT) architectures for the context-encoder, target-encoder, and the predictor. While the context-encoders and target-encoders correspond to standard ViT architectures, the predictor is designed as a light-weight (narrow) ViT architecture. Specifically, we fix the embedding dimension of the predictor to 384, while keeping the number of self-attention heads equal to that of the backbone context-encoder. For the smaller ViT-B/16 context-encoder, we set the depth of the predictor to 6. For ViT-L/16, ViT-H/16, and ViT-H/14 context-encoders, we set the depth of the predictor to 12. Finally, the ViT-G/16 uses a predictor of depth 16. I-JEPA is pretrained without a [cls] token. We use the target-encoder for evaluation and average pool its output to produce a global image representation.

Optimization.

We use AdamW loshchilov2017decoupled to optimize the context-encoder and predictor weights. Our default batch-size is 2048, and the learning rate is linearly increased from 10−410^{-4} to 10−310^{-3} during the first 15 epochs of pretraining, and decayed to 10−610^{-6} following a cosine schedule thereafter. Following caron2021emerging; assran2022masked, the weight-decay is linearly increased from 0.040.04 to 0.40.4 throughout pretraining. The target-encoder weights are identical to the context-encoder weights at initialization, and updated via an exponential moving average thereafter tarvainen2017mean; he2019moco; chen2020mocov2; grill2020bootstrap; caron2021emerging; assran2022masked. We use a momentum value of 0.9960.996, and linearly increase this value to 1.01.0 throughout pretraining, following caron2021emerging; assran2022masked.

Masking.

By default, we sample 44 possibly overlapping target blocks masks with random scale in the range (015,0.2)(015,0.2) and aspect ratio in the range (0.75,1.5)(0.75,1.5). We sample 11 context block mask with random scale in the range (0.85,1.0)(0.85,1.0) and unit aspect ratio. We subsequently eliminate any regions in the context block mask that overlap with any of the 44 target block masks. The context-block mask and target-block masks are sampled independently for each image in the mini-batch. To ensure efficient batch processing, we restrict the size of all context masks on a co-located GPU to be identical. Similarly, we restrict the size of all target masks on a co-located GPU to be identical. The mask-sampler is efficiently implemented in only a few lines of code in PyTorch paszke2019pytorch using a batch-collator function, which runs in the data loader processes. In short, in each iteration, the data loader returns a mini-batch of images and a set of context and target masks for each image, identifying the patch indices to keep for the context and target views.

A.2 Downstream Tasks

When evaluating methods such as iBOT zhou2021ibotyes, DINO caron2021emerging or MAE he2021masked, which leverage Vision Transformers dosovitskiy2020image with an additional [cls] token, we use the default configurations of VISSL goyal2021vissl to evaluate all the models on iNaturalist18 van2018inaturalist, CIFAR100 krizhevsky2009learning, Clevr/Count johnson2017clevr; vtab, Clevr/Dist johnson2017clevr; vtab, and Places205 zhou2014learning. We freeze the encoder and return the best number among the following representations: 1) the [cls] token representation of the last layer, 2) the concatenation of the last 44 layers of the [cls] token. For each representation, we try two different heads: 1) a linear head, or 2) a linear head preceded by a batch normalization, and return the best number. We use the default data augmentations of VISSL goyal2021vissl: random resize cropping and horizontal flipping, with the exception of Clevr/Count and Clevr/Dist, where we only use center crop and horizontal flipping, as random cropping interferes with the capability of counting objects and estimating distance, removing critical objects from the scene. For CIFAR100, we resize the images to 224×224224\times 224 pixels, so as to keep the number of patches equal to that used during pretraining.

Because our I-JEPA implementation uses Vision Transformer architectures without a [cls] token, we adapt the default VISSL evaluation recipe to utilize the average-pooled patch representation instead of the [cls] token. We therefore report the best linear evaluation number among the following representations: 1) the average-pooled patch representation of the last layer, 2) the concatenation of the last 44 layers of the average-pooled patch representations. We otherwise keep the linear-probing recipe identical.

ImageNet evaluations.

To evaluate the I-JEPA on ImageNet russakovsky2015imagenet, we adapt the VISSL recipe to use average pooled representations instead of the [cls] token. Following MAE he2021masked, we use the LARS lars optimizer with a batch-size of 1638416384, and train the linear probe for 50 epochs. We use a learning rate with a step-wise decay, dividing it by a factor of 1010 every 1515 epochs, and sweep three different reference learning rates [0.01,0.05,0.001][0.01,0.05,0.001], and two weight decay values [0.0005,0.0][0.0005,0.0].

Low-shot evaluation.

To evaluate our model on the ImageNet-1% low-shot task, we adapt the fine-tuning protocol of MAE he2021masked.We fine-tune our ViT-L/H models for 50 epochs on ImageNet-1% with the AdamW optimizer and a cosine learning rate scheduler. We use a batch size of 512, a learning rate layer decay of 0.750.75 and 0.10.1 label smoothing. We use the default randaugment data-augmentations as in MAE. In contrast to the fine-tuning done with MAE, we do not use mixup, cutmix, random erasing or drop path. For the I-JEPA, we use a learning rate /weight decay of 3e-5/5e-2 for the ViT-L/16, 3e-5/4e-1 for the ViT-H/14 and 3e-5/4e-1 for the ViT-H/16448. Similar fine-tuning strategy for low-shot learning has been explored by Semi-VIT in the context of semi-supervised learning cai2022semi.

Appendix B Broader Related Work

Self-supervised learning of visual representations with joint-embedding architectures is an active line of research wu2018unsupervised; he2019moco; chen2020exploring; grill2020bootstrap; chen2020mocov2; caron2021emerging; bardes2021vicreg; zhou2021ibotyes; bordes2022guillotine; mitrovic2020representation; assran2020supervision. These approaches train a pair of encoders to output similar embeddings for two or more views of the same image. To avoid pathological solution, many popular joint-embedding approaches use explicit regularization chen2020simple; caron2021emerging; bardes2021vicreg; assran2021semi or architectural constraints grill2020bootstrap; chen2020exploring. Collapse-prevention based on architectural constraints leverage specific network design choices to avoid collapse, for example, by stopping the gradient flow in one of the joint-embedding branches chen2020simple, using a momentum encoder in one of the joint-embedding branches grill2020bootstrap, or using an asymmetric prediction head grill2020bootstrap; chen2020simple; baevski2022data2vec. Recent work tian2021understanding attempts to theoretically understand (in certain simplified settings) how joint-embedding methods with architectural constraints avoid representation collapse without explicit regularization.

Typical regularization-based approaches to collapse prevention in joint-embedding architectures try to maximize the volume of space occupied by the representations. This is often motivated through the InfoMax ma2022principles principle. Indeed, a longstanding conviction in unsupervised representation learning is that the resulting representations should be both maximally informative about the inputs, while also satisfying certain simplicity constraints linsker1988self; goodfellow2016deep. The former objective is often referred to as the information-maximization principle (InfoMax), while the latter is sometimes referred to as the parsimony principle ma2022principles. Such approaches to representation learning have been proposed for decades (e.g., bridle1991unsupervised), where, historically, simplicity constraints were enforced by encouraging the learned representations to be sparse, low-dimensional, or disentangled, i.e., the individual dimensions of the representation vector should be statistically independent goodfellow2016deep. Modern approaches enforce the simplicity constraints coupled with InfoMax regularization through self-supervised loss terms hjelm2018learning; tschannen2019mutual; bachman2019learning; krause2010discriminative; hu2017learning; oord2018representation. One example is the widespread view-invariance penalty misra2020self, often coupled with with independence zbontar2021barlow; bardes2021vicreg or low-dimensionality constraints, e.g., by projecting representations on the unit hypersphere chen2020simple; he2019moco; grill2020bootstrap. However, despite its proliferation, there have also been many criticisms of the InfoMax principle, especially since it is does not discriminate between different types of information (e.g, noise and semantics) assran2022hidden. Indeed, the sets of features we wish the model to capture are not always those with the highest marginal entropy (maximal information content).

Orthogonal to the contributions of invariance-based pretraining, another line of work attempts to learn representations by artificially masking parts of the input and training a network to reconstruct the hidden content vincent2010stacked. Autoregressive models, and denoising autoencoders in particular, predict clean visual inputs from noisy views chen2020generative; vincent2010stacked; he2021masked; bao2021beit; baevski2022data2vec. Typically, the goal is to predict missing inputs at a pixel level dosovitskiy2020image; he2021masked; xie2019unsupervised, or at a patch token-level, using a tokenizer bao2021beit; wei2021masked. While these works demonstrate impressive scalability, they usually learn features at a low-level of semantic abstraction compared to joint-embedding approaches assran2022masked.

More recently, a set of approaches attempt to combine both joint-embedding architectures and reconstruction based approaches el2021large, wherein they combine an invariance pretraining loss with a patch-level reconstruction loss, as in the iBOT method zhou2021ibotyes. Since view-invariance based approaches are typically biased towards learning global image representations, thereby limiting their applicability to other computer vision tasks, the idea is that adding local loss terms can improve performance on other popular tasks in computer vision chen2022intra; gidaris2020learning; bardes2022vicregl. The framework of contrastive predictive coding oord2018representation is also closely related to this line of work on local loss terms. In the context of images henaff2020data, here the idea is to use a contrastive objective combined with a convolutional network to discriminate between overlapping image patch representations. Specifically, the goal is to encourage the representations of an image patch to be predictable of the image patches directly below it, while pushing away the representations of other patch views. In contrast to that work, the proposed I-JEPA method is non-contrastive and does not seek to discriminate between image patches. Rather, the goal is to predict the representations of various target blocks from a single context block. This is achieved with a Joint-Embedding Predictive Architecture, using a predictor network that is conditioned on positional embeddings corresponding to the location of the target block in the image. Qualitative experiments in Section 8 show that the predictor network in our architecture learns to correctly perform this local-to-local region feature mapping, and learns to correctly capture positional uncertainty in the image.

Appendix C Additional Ablations

This section follows the same experimental protocol as Section 9. We report the result of a linear probe with a frozen backbone, trained on the low-shot 1% ImageNet-1K benchmark.

We present an extended ablation of the multiblock masking strategy where we change the targets block scale (Table 8), the context scale (Table 9) and the number of target blocks (Table 10). We train a ViT-B/16 for 300 epochs using I-JEPA with various multi-block settings and compare performance on the 1% ImageNet-1K benchmark using a linear probe. In short, we find that it is important to predict several relatively large (semantic) target blocks, and to use a sufficiently informative (spatially distributed) context block.

Masking at the output of the target-encoder.

An important important design choice in I-JEPA is that the target blocks are obtained by masking the output of the target-encoder, not the input. Table 11 shows the effect of this design choice on the semantic level of the learned representations when pretraining a ViT-H/16 using I-JEPA for 300 epochs. In the case where masking is applied to the input, we forward-propagate through the target-encoder once for each target region. Masking the output of the target-encoder during pretraining results in more semantic prediction targets and improves linear probing performance.

Predictor depth.

We examine the impact of the predictor depth on the downstream low-shot performance in Table 12. We pretrain a ViT-L/16 for 500 epochs using either a 6-layer predictor network or a 12-layer predictor network. The model pretrained using a deeper predictor shows a significant improvement in downstream low-shot performance compared to the model pretrained with a shallower predictor.

Weight decay.

In Table 13, we evaluate the impact of weight-decay during pretraining. We explore two weight decay strategies: linearly increase the weight-decay from 0.040.04 to 0.40.4 or use a fix weight-decay of 0.050.05. Using a smaller weight decay during pretraining improves the downstream performance on ImageNet-1% when fine-tuning. However, this also leads to a degradation of performance in linear evaluation. In the main paper, we use the first weight decay strategy as it improves the performances in linear evaluation downstream tasks.

Predictor width.

We explore the impact of the predictor width in Table 14. We compare I-JEPA using a ViT-L encoder and a predictor with 386386 channels to a similar model using a predictor with 10241024 channels. Note that the ViT-L encoder has 10241024 channels. Using a bottleneck in the predictor width improves the downstream performance on ImageNet 1%.

Appendix D Finetuning on the full ImageNet

In this section, we report performance on I-JEPA when fine-tuning on the full ImageNet dataset. We focus on the ViT-H/16448 as this architecture achieves state-of-art performance with MAE he2021masked.

We use a fine-tuning protocol similar to MAE. Specifically, we fine-tune our model for 5050 epochs using AdamW and a cosine learning rate schedule. The base learning rate is set to 10−410^{-4} and the batch size to 528528. We train using mixup zhang2018mixup set to 0.80.8, cutmix yun2019cutmix set to 1.01.0, a drop path probability of 0.250.25 and a weight decay set to 0.040.04. We also use a layer decay of 0.750.75. Finally, we use the same rand-augment data-augmentations as MAE,

Table 15 reports the fine-tuning results. I-JEPA achieves 87.187.1 top-1 accuracy. Its performance is less than 1%1\% away from the best MAE model despite I-JEPA being trained for 5.35.3 times less epochs than MAE. This result demonstrates that I-JEPA is competitive when fine-tuning on the full ImageNet dataset.

Appendix E RCDM Visualizations

To visualize the representations of a pretrained neural network in pixel space, we use the RCDM framework bordes2022high. The RCDM framework trains a decoder network hωh_{\omega}, comprising a generative diffusion model, to reconstruct an image x{\bm{x}} from the representation vector of that image sx{\bm{s}}_{x} and a noisy version of that image x^≔x+ϵ\hat{{\bm{x}}}\coloneqq{\bm{x}}+\epsilon, where ϵ\epsilon is an additive noise vector. Concretely, the decoder objective is to minimize the loss function ∥hω(x^,sx)−ϵ∥\lVert h_{\omega}(\hat{{\bm{x}}},{\bm{s}}_{x})-\epsilon\rVert. We train each RCDM network for 300,000 iterations using the default hyperparameters bordes2022high. After training the decoder, one can subsequently feed the representation vector of an unseen test image sy{\bm{s}}_{y} into the decoder along with various random noise vectors to generate several pixel-level visualizations of the representation, thus providing insight into the features captured in the representations of the pretrained network. Qualities that are common across samples represent information that is contained in the representation. On the other hand, qualities that vary across samples represent information that is not contained in the representations

In Figure 6, the visualizations are obtained by feeding the average-pooled output of the predictor, conditioned on a specific target region, into the decoder network, along with various random noise vectors. In Figures 7 and 8, the visualizations are obtained by feeding the average-pooled output of the target-encoder into the decoder network, along with various random noise vectors.

In Figure 7, we visualize the average-pooled I-JEPA representations at the output of our ViT-H/14 target-encoder. The first column contains the original image, while subsequent columns contain synthetic samples obtained by feeding the average-pooled representation of the image into the decoder along with various random noise vectors. Figure 7 suggests that the I-JEPA target-encoder is able to correctly capture the high-level information regarding objects and their poses, while discarding low-level image details and background information.

Figure 8 shows similar visualizations, but when using an MSN assran2022masked pretrained ViT-L/7 target-encoder to compute the image representations. The MSN method trains a context- and target-encoder using a Joint-Embedding Architecture to enforce invariance of global image representations to various hand crafted data augmentations and missing patches. While the MSN pretrained network is able to capture high level semantic information about the image in the first column, it also exhibits higher variability in the generated samples, e.g., variability in the object pose, object scale, and number of instances. In short, the MSN pretrained discards much of the local structure in the image, which is in stark contrast to I-JEPA, which retains information about much of the local structure in the input image.