Understanding Self-supervised Learning with Dual Deep Networks
Yuandong Tian, Lantao Yu, Xinlei Chen, Surya Ganguli
Introduction
While self-supervised learning (SSL) has achieved great empirical success across multiple domains, including computer vision (He et al., 2020; Goyal et al., 2019; Chen et al., 2020a; Grill et al., 2020; Misra & Maaten, 2020; Caron et al., 2020), natural language processing (Devlin et al., 2018), and speech recognition (Wu et al., 2020; Baevski & Mohamed, 2020; Baevski et al., 2019), its theoretical understanding remains elusive, especially when multi-layer nonlinear deep networks are involved (Bahri et al., 2020). Unlike supervised learning (SL) that deals with labeled data, SSL learns meaningful structures from randomly initialized networks without human-provided labels.
In this paper, we propose a systematic theoretical analysis of SSL with deep ReLU networks. Our analysis imposes no parametric assumptions on the input data distribution and is applicable to state-of-the-art SSL methods that typically involve two parallel (or dual) deep ReLU networks during training (e.g., SimCLR (Chen et al., 2020a), BYOL (Grill et al., 2020), etc). We do so by developing an analogy between SSL and a theoretical framework for analyzing supervised learning, namely the student-teacher setting (Tian, 2020; Allen-Zhu & Li, 2020; Lampinen & Ganguli, 2018; Saad & Solla, 1996), which also employs a pair of dual networks. Our results indicate that SimCLR weight updates at every layer are amplified by a fundamental positive semi definite (PSD) covariance operator that only captures feature variability across data points that survive averages over data augmentation procedures designed in practice to scramble semantically unimportant features (e.g. random image crops, blurring or color distortions (Falcon & Cho, 2020; Kolesnikov et al., 2019; Misra & Maaten, 2020; Purushwalkam & Gupta, 2020)). This covariance operator provides a principled framework to study how SimCLR amplifies initial random selectivity to obtain distinctive features that vary across samples after surviving averages over data-augmentations.
While the covariance operator is a mathematical object that is valid for any data distribution and augmentations, we further study its properties under specific data distributions and augmentations. We first start with a simple one-layer case where two 1D objects undergo 1D translation, then study a fairly general case when the data are generated by a hierarchical latent tree model (HLTM), which can be regarded as an abstract conceptual model for object compositionality in computer vision. In this case, training deep ReLU networks on the data generated by the HLTM leads to the emergence of learned representations of the latent variables in its intermediate layers, even if these intermediate nodes have never been directly supervised by the unobserved and inaccessible latent variables. This shows that in theory, useful hidden features can automatically emerge by contrastive self-supervised learning.
To the best of our knowledge, we are the first to provide a systematic theoretical analysis of modern SSL methods with deep ReLU networks that elucidates how both data and data augmentation, drive the learning of internal representations across multiple layers.
Related Works
In addition to SimCLR and BYOL, many concurrent SSL frameworks exist to learn good representations for computer vision tasks. MoCo (He et al., 2020; Chen et al., 2020b) keeps a large bank of past representations in a queue as the slow-progressing target to train from. DeepCluster (Caron et al., 2018) and SwAV (Caron et al., 2020) learn the representations by iteratively or implicitly clustering on the current representations and improving representations using the cluster label. (Alwassel et al., 2019) applies similar ideas to multi-modality tasks. Contrastive Predictive Coding (Oord et al., 2018) learns the representation by predicting the future of a sequence in the latent space with autoregressive models and InfoNCE loss. Contrastive MultiView Coding (Tian et al., 2019) uses multiple sensory channels (called “views”) of the same scene as the positive pairs and independently sampled views as the negative pairs to train the model. Recently, (Li et al., 2020) moves beyond instance-wise pairs and proposes to use prototypes to construct training pairs that are more semantically meaningful.
In contrast, the literature on the (theoretical) analysis of SSL is sparse. (Wang & Isola, 2020) shows directly optimizing the alignment/uniformity of the positive/negative pairs leads to comparable performance against contrastive loss. (Arora et al., 2019b) proposes an interesting analysis of how contrastive learning aids downstream classification tasks, given assumptions about data generation. (Lee et al., 2020) analyzes how learning pretext tasks could reduce the sample complexity of the downstream task and (Tosh et al., 2020) analyzes contrastive loss with multi-view data in the semi-supervised setting, with different generative models. However, they either work on linear models or treat deep models as a black-box function approximators with sufficient capacity. In comparison, we incorporate self-supervision, deep models, contrastive loss, data augmentation and generative models together into the same theoretical framework, and make an attempt to understand how and what intermediate features emerge in modern SSL architectures with deep models that achieve SoTA.
Overall framework
See Appendix for proofs of all theorems in the main text.
The Covariance Operator
In this paper, we consider three different contrastive losses:
We first note one interesting common property:
With Theorem 1 and Theorem 2, we now present our first main contribution of this paper: the gradient in SimCLR is governed by a positive semi-definite (PSD) covariance operator at any layer :
where the pairwise weight takes the following form for different loss functions:
2 More general loss functions
For more general loss functions in which , we have a corollary:
Feature Emergence through Covariance Operator Based Amplification
The covariance operator in Theorem 3-4 applies to arbitrary data distributions and augmentations. While this conclusion is general, it is also abstract. To understand what feature representations emerge, we study learning under more specific assumptions on the generative process underlying the data.
We make two assumptions under the generative paradigm of (Fig. 2):
The input is generated by two groups of latent variables, class/sample-specific latents and nuisance latents .
Data augmentation changes while preserving .
In this setting, we first show that a linear neuron performs dimensionality reduction within an augmentation preserved subspace. We then consider how nonlinear neurons with local receptive fields (RFs) can learn to detect simple objects. Finally, we extend our analysis to deep ReLU networks exposed to data generated by a hierarchical latent tree model (HLTM), proving that, with sufficient over-parameterization, there exist lucky nodes at initialization whose activation is correlated with latent variables underlying the data, and that SimCLR amplifies these initial lucky representations during learning.
A single linear neuron cannot detect localized objects. We now consider a generative model in which data vectors can be thought of as images of objects of the form where is an important latent semantic variable denoting object identity and is a nuisance latent representing its spatial location. The augmentation procedure scrambles position while preserving object identity (Fig. 3):
We next show both a local RF and nonlinearity can rescue this unfortunate situation.
Thus the learning dynamics amplifies the initial selectivity to the object selective feature vector in a way that cannot be done with a linear neuron. Note this argument also holds with bias terms and initial selectivity for more than one pattern. Moreover, with a local RF, the probability of weak initial selectivity to some local object sensitive features is high, and we may expect amplification of such weak selectivity in real neural network training, as observed in other settings (Williams et al., 2018).
Deep ReLU SSL training with Hierarchical Latent Tree Models (HLTM)
Here we describe a general Hierarchical Latent Tree Model (HLTM) of data, and the structure of a multilayer neural network that learns from this data.
Motivation. The HLTM is motivated by the hierarchical structure of our world in which objects may consist of parts, which in turn may consist of subparts. Moreover the parts and subparts may be in different configurations in relation to each other in any given instantiation of the object, or any given subpart may be occluded in some views of an object.
Data Augmentation. Given a sample , data augmentation involves resampling all (which are in Fig. 4), while fixing the root . This models augmentations as changing part configurations while keeping object identity fixed.
The neural network. We now consider the multi-layer ReLU network that learns from data generated from HLTM (right hand side of Fig. 4). For simplicity let . The neural network has a set of input neurons that are in one to one correspondence with the pixels or visible variables that arise at the leaves of the HLTM, where . For any given object at layer , and its associated parts states at layer , and visible feature values at layer , the input neurons of the neural network receive only the visible feature values as real analog inputs. Thus the neural network does not have direct access to the latent variables and that generate these visible variables.
where the polarity measures how informative is. If then there is no stochasticity in the top-down generation process. If , then there is no information in the downstream latents and the posterior of given the observation can only be uniform. See Appendix for more general cases.
First, we prove that given sufficient over-parameterization (), even at initialization, without any training, we can find some lucky nodes with weak selectivity:
See Appendix for detailed theorem description and proof. is defined in Fig. 5(c) and is a weak threshold that increases monotonically w.r.t. all its arguments. This means that higher polarity, more selectivity in the lower layer and more over-parameterization (larger ) all boost weak initial selectivity of a lucky node.
1.2 Training with constant Jacobian
Here is the polarity between and , which might be far apart in hierarchy. A simple computation (See Lemma and associated remarks in Appendix) shows that is a product of consequent polarities in the tree hierarchy.
If we further assume , then after the gradient update, for the “lucky” node we have:
In Sec. 7, as predicted by our theory, the intermediate layers of deep ReLU networks do learn the latent variables of the HLTM (see Tbl. 6 below and Appendix).
Experiments
We test our theoretical findings through experiments on CIFAR-10 (Krizhevsky et al., 2009) and STL-10 (Coates et al., 2011). We use a simplified linear evaluation protocol: the linear classifier is trained on frozen representations computed without data augmentation. This reuses pre-computed representations and is more efficient. We use ResNet-18 as the backbone and all experiments are repeated 5 times for mean and std. Please check detailed setup in Appendix.
Extended contrastive loss function. When , the covariance operator can still be derived (Sec. 4.2) and remains PSD when . As shown in Tbl. 3, we find that (1) performs better at first 100 epochs but converges to similar performance after 500 epochs, suggesting that might accelerate training, (2) worsens performance. We report a similar observation on ImageNet (Deng et al., 2009), where our default SimCLR implementation achieves 64.6% top-1 accuracy with 60-epoch training; setting yields 64.8%; and setting hurts the performance.
Hierarchical Latent Tree Model (HLTM). We implement HLTM and check whether the intermediate layers of deep ReLU networks learn the corresponding latent variables at the same layer. The degree of learning is measured by the normalized correlations between the ground truth latent variable and its best corresponding node . Tbl. 4 indicates this measure increases with over-parameterization and learning, consistent with our analysis (Sec. 6.1). More experiments in Appendix.
Conclusion and Future Works
In this paper, we propose a novel theoretical framework to study self-supervised learning (SSL) paradigms that consist of dual deep ReLU networks. We analytically show that the weight update at each intermediate layer is governed by a covariance operator, a PSD matrix that amplifies weight directions that align with variations across data points which survive averages over augmentations. We show how the operator interacts with multiple generative models that generate the input data distribution, including a simple 1D model with circular translation and hierarchical latent tree models. Experiments support our theoretical findings.
To our best knowledge, our work is the first to open the blackbox of deep ReLU neural networks to bridge contrastive learning, (hierarchical) generative models, augmentation procedures, and the emergence of features and representations. We hope this work opens new opportunities and perspectives for the research community.
References
Appendix A Background and Basic Setting (Section 3)
Note that many different kinds of layers have this reversible property, including linear layers (MLP and Conv) and (leaky) ReLU nonlinearity. For linear layers, at layer , we have:
For multi-layer ReLU network, for each layer , we have:
In addition to ReLU, other activation function also satisfies this condition, including linear, LeakyReLU and monomial activations. For example, for power activation where , we have:
Remark. Note that the reversibility is not the same as invertible. Specifically, reversibility only requires the transfer function of a backpropagation gradient is a transpose of the forward function.
with respect to weight matrix at layer yields the following gradient at layer :
We prove by induction. Note that our definition of is the transpose of defined in (Tian, 2020). Also our is the gradient before nonlinearity, while (Tian, 2020) uses the same symbol for the gradient after nonlinearity.
Remark on ResNet. Note that the same structure holds for blocks of ResNet with ReLU activation.
A.2 Theorem 1
Now we prove Theorem 1. Note that for deep ReLU networks, is a simple identity matrix and thus:
Here , and .
We consider a more general case where the two towers have different parameters, namely and . Applying Lemma 1 for the branch with input at the linear layer , and using Eqn. 31 we have:
In this case, the gradient (and the weight update, according to gradient descent) of the weight between layer and layer is:
Note that is a function of the current weight , which includes weights at all layers. By the mixed-product property of Kronecker product , we have:
where and .
In SimCLR case, we have so
Appendix B Analysis of SimCLR using Teacher-Student Setting (Section 4)
B.2 The Covariance Operator under different loss functions
Then we compute each terms. Using Theorem 1, we know that:
Since Eqn. 45 holds, will be cancelled out and we have:
B.3 Theorem 3
The conclusion follows since gradient descent is used and . ∎
Note that it is no longer symmetric. We leave detailed discussion to the future work.
B.4 Theorem 4
where the pairwise weight takes the following form for different loss functions:
We consider where there is only a single negative pair (and ). In this case . Let
The constant term with respect to data augmentation. In the following, we first consider the term , which only depends on un-augmented data points . From Lemma 2, we now have a term in the gradient:
Symmetrically, if we swap and since both are sampled from the same distribution , we have:
then it is clear that . We can compute its partial derivatives:
Note that is always bounded.
Note that is a differentiable function with respect to and . We do a Taylor expansion of on the two variables , using intermediate value theorem, we have:
Therefore, we have and we have:
Similarly, for we want to break the term into groups of terms, each is a difference within data augmentation. Using that
we have (here , , and ):
Let so finally we have:
which causes issues with the symmetry trick (Eqn. 67), because the denominator involves many negative pairs at the same time.
However, if we think given one pair of distinct data point , the normalized constant averaged over data augmentation is approximately constant due to homogeneity of the dataset and data augmentation, then Eqn. 67 can still be applied and similar conclusion follows.
The conclusion follows since gradient descent is used and . ∎
Appendix C Hierarchical Latent Tree Models (Section 6)
If function is twice differentiable, and , then we have:
For ReLU activation and , we have:
Note that for the expectation of absolute value , we have:
where the last inequality is due to Cauchy-Schwarz. ∎
C.2 Taxonomy of HLTM
Note that is obvious due to the property of conditional probability. The real condition is . If , then is a square matrix and Def. 3 is equivalent to is double-stochastic. The Def. 3 makes computation of easy for any and .
As we will see in the remark of Lemma 6, symmetric binary HLTM (SB-HLTM) naturally satisfies Def. 3 and is a TR-HLTM. Conversely, TR-HLTM allows categorical latent variables and is a super set of SB-HLTM. See Fig. 9.
For TR-HLTM (Def. 3), for , and , we have:
In general, for any and with , we have:
since , and , the conclusion follows. ∎
TR-HLTM with all latent variables being binary (or binary TR-HLTM) is equivalent to SB-HLTM.
SB-HLTM is TR-HLTM. For symmetric binary HLTM (SB-HLTM) defined in Sec. 6.1, we could define as the following (here ):
And it is easy to verify this satifies the definition of TR-HLTM (Def. 3).
Binary TR-HLTM is SB-HLTM. To see this, note that and provides 4 linear constraints (1 redundant), leaving 1 free parameter, which is the polarity . It is easy to verify that in order to keep a probabilistic measure. Moreover, since , the parameterization is close under multiplication:
C.3 SB-HLTM (Sec. 6.1)
Since we are dealing with binary case, we define the following for convenience:
For standardized Gaussian variable and with (here ):
Then we have lower-bound for the probability:
where is defined in Eqn. 121.
First we compute the precision matrix :
and we can apply the lower-bound of the Mills ratio (The equation below Eqn. 5 in (Steck, 1979)) for standarized 2D Gaussian distribution (note that “” in (Steck, 1979) is here):
where is the 1D Mills ratio and other terms defined as follows:
where is the Cumulative Probability Distribution (CDF) for standard Gaussian distribution.
Remarks. Obviously, is a decreasing function. Following Theorem 2 (Soranzo & Epure, 2014), we have for :
and we could estimate the bound of accordingly. Using Eqn. 120 to estimate the bound would be very difficult numerically since becomes very close to zero once becomes slightly large (e.g., when ).
Note that from Theorem 1 in (Soranzo & Epure, 2014), we have a coarser bound for and (the bound will blow-up when ):
So when is large, we have . Therefore, when is close to and is large (Fig. 11):
Therefore, has the following properties:
Note that when , Eqn. 126 is not accurate and there is no singularity in near , as shown in Fig. 11.
C.4 Lucky node at initialization (Sec. 6.1.1)
According to our setting, for each node , there exists a unique latent variable with that corresponds to it. In the following we omit its dependency on for brevity.
Define the positive and negative set (note that ):
Without loss of generality, assume that . In the following, we show there exists with is greater than some positive threshold. Otherwise the proof is symmetric and we can show is lower than some negative threshold.
Now we want to lower-bound the following probability. For some :
And the probability we want to lower-bound becomes:
Then satisfies the condition of Lemma 8 with :
Therefore, with (note that the residue term is defined in Lemma 8):
which means that with at least probability, there exists at least one node and , so that Eqn. 143 holds.
By Jensen’s inequality, we have (note that is the ReLU activation):
As a side note, using Lemma 3, since ReLU function has Lipschitz constant (empirically it is smaller), we know that:
Combining Eqn. 154 and Eqn. 163, we have a bound for :
Note that and the conclusion follows. ∎
Remark. Intuitively, this means that with large polarity (strong top-down signals), randomly initialized over-parameterized ReLU networks yield selective neurons, if the lower layer also contains selective ones. For example, when , , if then there is a gap between expected activations and , and the gap is larger when the selectivity in the lower layer is higher. With the same setting, if , we could actually compute the range of specified by Theorem 5, using Eqn. 124:
while in practice, we don’t need such a strong over-parameterization.
C.5 Training dynamics (Sec. 6.1.2)
So it reduces to computing . Note that following Eqn. 169:
since . In SB-HLTM, according to Lemma 7, all and we could compute :
Therefore, we have since .
On the other hand, according to Eqn. 103, we have:
Put them together, the covariance operator is:
When , from Lemma 6 and its remark, we have:
and thus the covariance becomes zero as well. ∎
Appendix D Additional Experiments
Normalized correlation between latents in SB-HLTM and hidden nodes in neural networks. We provide additional table (Tbl. 6) that shows that besides the top-most layers, the normalized correlation (NC) between the hidden layer of the deep ReLU network and the intermediate latent variables of the hierarchical tree model is also strong at initialization and will grow over time. In particular, we can clearly see the following several trends:
Top-layer latent is less correlated with the top-layer nodes in deep networks. This shows that top-layer latents are harder to learn (and bottom-layer latents are easier to learn since they are closer to the corresponding network layers).
There is already a non-trivial amount of NC at initialization. Furthermore, NC is higher in the bottom-layer, since they are closer to the input. Nodes in the lowest layer (leaves) would have perfect NC (1.0), since they are identical to the leaves of the latent tree models.
All experiments are repeated 10 times and its mean and standard derivation are reported.
D.2 Experiments Setup
For all STL-10 (Coates et al., 2011) and CIFAR-10 (Krizhevsky et al., 2009) task, we use ResNet18 (He et al., 2016) as the backbone with a two-layer MLP projector. We use Adam (Kingma & Ba, 2015) optimizer with momentum and no weight decay. For CIFAR-10 we use learning rate and for STL-10 we use learning rate. The training batchsize is .