Understanding self-supervised Learning Dynamics without Contrastive Pairs

Yuandong Tian, Xinlei Chen, Surya Ganguli

Introduction

Self-supervised learning (SSL) has emerged as a powerful method for learning useful representations without requiring expensive target labels (Devlin et al., 2018). Many state-of-the-art SSL methods in computer vision employ the principle of contrastive learning (Oord et al., 2018; Tian et al., 2019; He et al., 2020; Chen et al., 2020a; Bachman et al., 2019) whereby the hidden representations of two augmented views of the same object (positive pairs) are brought closer together, while those of different objects (negative pairs) are encouraged to be further apart. Minimizing differences between positive pairs encourages modeling invariances, while contrasting negative pairs is thought to be required to prevent representational collapse (i.e., mapping all data to the same representation).

However, some recent SSL work, notably BYOL (Grill et al., 2020) and SimSiam (Chen & He, 2020), have shown the remarkable capacity to learn powerful representations using only positive pairs, without ever contrasting negative pairs. These methods employ a dual pair of Siamese networks (Bromley et al., 1994) (Fig. 1): the representation of two views are trained to match, one obtained by the composition of an online and predictor network, and the other by a target network. The target network is not trained via gradient descent; and either employs a direct copy of the online network (e.g., SimSiam (Chen & He, 2020)), or a momentum encoder that slowly follows the online network in a delayed fashion through an exponential moving average (EMA) (e.g., MoCo (He et al., 2020; Chen et al., 2020b) and BYOL (Grill et al., 2020)). Compared to contrastive learning, these non-contrastive SSL methods do not require large batch size (e.g., 4096 in SimCLR (Chen et al., 2020a)) or memory queue (e.g., MoCo (He et al., 2020; Chen et al., 2020b)) to provide negative pairs. Therefore, they are generally more efficient and conceptually simple while maintaining state-of-the-art performance.

Since the entire procedure in non-contrastive SSL encourages the online+predictor network and the target network to become similar to each other, this overall scheme raises several fundamental unsolved theoretical questions. Why/how does it avoid collapsed representations? What is the nature of the learned representations? How do multiple design choices and hyperparameters interact nonlinearly in the learning dynamics? While there are interesting theoretical studies of contrastive SSL (Arora et al., 2019; Lee et al., 2020; Tosh et al., 2020), any theoretical understanding of the nonlinear learning dynamics of non-contrastive SSL remains open.

In this paper, we make a first attempt to analyze the behavior of non-contrastive SSL training and the empirical effects of multiple hyperparameters, including (1) Exponential Moving Average (EMA) or momentum encoder, (2) Higher relative learning rate (αp\alpha_{p}) of the predictor, and (3) Weight decay η\eta. We explain all these empirical findings with an exceedingly simple theory based on analyzing the nonlinear learning dynamics of simple linear networks. Note that deep linear networks have provided a useful tractable theoretical model of nonconvex loss landscapes (Kawaguchi, 2016; Du & Hu, 2019; Laurent & Brecht, 2018) and nonlinear learning dynamics (Saxe et al., 2013, 2019; Lampinen & Ganguli, 2018; Arora et al., 2018) in these landscapes, yielding insights like dynamical isometry (Saxe et al., 2013; Pennington et al., 2017, 2018) that lead to improved training of nonlinear deep networks. Despite the simplicity of our theory, it can still predict how various hyperparameter choices affect performance in an extensive set of real-world ablation studies. Moreover, the simplicity also enables us to provide conceptual and analytic insights into why performance patterns vary the way they do. Specifically, our theory accounts for the following diverse empirical findings:

Essential part of non-contrastive SSL. The existence of the predictor and stop-gradient is absolutely essential. Removing either of them leads to representational collapse in BYOL and SimSiam.

EMA. While the original BYOL needs EMA to work, they later confirmed that EMA is not necessary (i.e., the online and target networks can be identical) if a higher αp\alpha_{p} is used. This is also confirmed with SimSiam, as long as the predictor is updated more often or has larger learning rate (or larger αp\alpha_{p}). However, the performance is slightly lower.

Weight Decay. Table 15 in BYOL (Grill et al., 2020) indicates that no weight decay may lead to unstable results. A recent blogpost (Fetterman & Albrecht, 2020) also mentions using weight decay leads to stable learning in BYOL.

Finally, motivated by our theoretical analysis, we propose a new method DirectPred that directly sets the predictor weights based on principal components analysis of the predictor’s input, thereby avoiding complicated predictor dynamics and initialization issues. We show that this simple DirectPred method nevertheless yields comparable performance in CIFAR-10 and outperforms gradient training of the linear predictor by +5%+5\% Top-1 accuracy in linear evaluation protocol on both STL-10 and ImageNet (60 epochs). On the standard ImageNet benchmark (300 epochs), DirectPred achieves 72.4%/91.0%72.4\%/91.0\% Top-1/Top-5, 2.5%2.5\% higher than BYOL with linear predictor (69.9%/89.6%69.9\%/89.6\%) and comparable with default BYOL setting with 2-layer predictor (72.5%/90.8%72.5\%/90.8\%).

Two-layer linear model

while the dynamics of WaW_{a} is obtained differently, via an exponential moving average (EMA) of WW. We will analyze this combined dynamics for WW, WpW_{p} and WaW_{a}, in the presence of additional weight decay, in the limit of large batch sizes and small discrete time learning rates. This limit can be well approximated by the gradient flow (see Supplementary Material (SM) for all derivations):

We note that since SimSiam is an ablation of BYOL that removes the EMA computation, the underlying dynamics of SimSiam can also be obtained from Lemma 1 simply by setting Wa=WW_{a}=W, inserting this relation into Eqn. 2 and Eqn. 3, and ignoring Eqn. 4. Importantly, the stop-gradient on the target branch is still there.

Overall Eqns. 2-4 constitute our starting point for analyzing the combined roles of relative learning rates αp\alpha_{p} and β\beta, weight decay rate η\eta and various ablations in determining the performance of both BYOL and SimSiam.

We first derive two very general results (see SM).

Completely independent of the particular dynamics of WaW_{a} in Eqn. 4, the update rules (Eqn. 2 and Eqn. 3) possess the invariance

where CC is a symmetric matrix that depends only on the initialization of WW and WpW_{p}.

This theorem implies that for both BYOL and SimSiam, there exists a “balancing” that ensures that any matching between the online and target representations will not be attributable solely to the predictor weights, rendering the online weights useless. Instead what the predictor learns, the online network will also learn, which is important as the online network’s representations are what is used for downstream tasks. We note that similar weight balancing dynamics has been discovered in multi-layer linear networks and matrix factorization (Arora et al., 2018; Du et al., 2018). Our results generalize this to SSL dynamics. Second, a nonzero weight decay could help remove the extra constant CC due to initialization, further balancing the predictor and online network weights and possibly leading to better performance on downstream tasks (Tbl. 2).

If the minimal eigenvalue λmin⁡(H(t))\lambda_{\min}(H(t)) over time is bounded below, inf⁡t≥0λmin⁡(H(t))≥λ0>0\inf_{t\geq 0}\lambda_{\min}(H(t))\geq\lambda_{0}>0, then W(t)→0W(t)\rightarrow 0.

Thus we have proven analytically in this simple setting that removing the stop-gradient leads to representational collapse, as observed in more complex settings in SimSiam (Chen & He, 2020). Similarly, with Wa=WW_{a}=W and no predictor (Wp=In2W_{p}=I_{n_{2}}), then the dynamics Eqn. 3 also reduces to a similar form and W(t)→0W(t)\rightarrow 0 (see SM).

How multiple factors affect learning dynamics

The learning dynamics in Eqns. 2-4 constitute a set of high dimensional coupled nonlinear differential equations that can be difficult to solve analytically in general. Therefore, to obtain analytic insights into the functional roles of the relative learning rates αp\alpha_{p} and β\beta and weight decay η\eta, we make a series of simplifying assumptions. Intriguingly, under these simplifying assumptions we obtain a rich set of analytic predictions, which we then test experimentally in more realistic scenarios. We find, nicely, that these predictions still qualitatively hold even when our simplifying assumptions required for obtaining analytic results do not.

Thus we obtain a reduced dynamics for WW, WpW_{p} and τ\tau. By not enforcing the stronger SimSiam constraint that Wa=WW_{a}=W, we can still model EMA dynamics. Intuitively, τ=τ(t)\tau=\tau(t) is a dynamic parameter that depends on how quickly W=W(t)W=W(t) grows over time. If WW is constant, then W˙=0\dot{W}=0 and τ\tau stabilizes to 11. On the other hand, if WW grows rapidly, then τ\tau becomes small. While Assumption 1 is a simplification, as we shall see, it still reveals interesting verifiable predictions about the functional role of EMA.

Many previous studies of deep learning dynamics made simplifying isotropic assumptions about data (Tian, 2017; Brutzkus & Globerson, 2017; Du et al., 2019; Bartlett et al., 2018; Safran & Shamir, 2018). Since our fundamental goal is to obtain the first analytic understanding of the dynamics of non-contrastive SSL methods, it is useful to first achieve this in the simplest possible isotropic setting. Interestingly, we will find that our final conclusions generalize to non-isotropic real world settings.

We enforce symmetry in WpW_{p} by initializing it to be a symmetric matrix, and then symmetrizing the flow for WpW_{p} in Eqn. 2 (see SM).

This symmetry assumption was motivated by both fixed point analysis and empirical findings. First, the fixed point of Eqn. 2 under Assumption 1 and 2 and η>0\eta>0 is always a symmetric matrix and in numerical simulation the asymmetric part Wp−Wp⊺W_{p}-W_{p}^{\intercal} eventually vanishes (See Appendix for the proof and numerical simulations). Moreover, during BYOL training without a symmetry constraint on the predictor, WpW_{p} gradually moves towards symmetry (Fig. 2).

Second, a set of experiments reveal that whether the predictor is symmetric or not has a dramatic effect in terms of both performance and interaction with EMA. In our STL-10 experiment, enforcing symmetric WpW_{p} in the presence of EMA improves performance on downstream tasks (Tbl. 3). In contrast, in the absence of EMA, a symmetric WpW_{p} fails while an asymmetric WpW_{p} works reasonably well. Similar behavior holds on ImageNet: a symmetric one layer linear predictor WpW_{p} in SimSiam (i.e. without EMA) achieves performance no better than random guessing (Top-1/5: 0.1%/0.5%0.1\%/0.5\%), while an asymmetric WpW_{p} achieves a Top-1/5 accuracy of 68.1%/88.2%68.1\%/88.2\%. Our theory will explain this as well as show how to obtain good performance with a symmetric predictor without EMA by increasing its relative learning rate αp\alpha_{p}.

This dynamics reveals that the eigenspace of WpW_{p} will gradually align with that of FF under certain conditions (see SM for derivation):

Under Eqn. 7, the commutator [F,Wp]:=FWp−WpF[F,W_{p}]:=FW_{p}-W_{p}F satisfies:

If inf⁡t≥0λmin⁡[K(t)]=λ0>0\inf_{t\geq 0}\lambda_{\min}[K(t)]=\lambda_{0}>0, then the commutator

For symmetric WpW_{p}, when WpW_{p} and FF commute they can be simultaneously diagonalized. Thus this shows that the eigenspace of WpW_{p} gradually aligns with that of FF.

To test this prediction, we performed extensive experiments showing that training BYOL using ResNet-18 on STL-10 yields eigenspace alignment, as demonstrated in Fig. 2.

This decoupled dynamics constitutes a dramatically simplified set of 33 dimensional nonlinear dynamical systems for BYOL learning, and two dimensional nonlinear systems (obtained by constraining τ=1\tau=1) for SimSiam. As expected, each mode’s dynamics is equivalent to the 33 dimensional dynamics obtained by setting n1=n2=1n_{1}=n_{2}=1 in Eqns. 2-4 and making the replacements W2=sjW^{2}=s_{j}, Wp=pjW_{p}=p_{j}, and Wa/W=τW_{a}/W=\tau (see SM). Thus the decoupled dynamics in Eqns 11- 13 reduce to the scalar case of BYOL dynamics in Eqns. 2-4 after a change of variables and the condition in Thm. 3 reveals when this decoupled regime is reachable.

Non-symmetric WpW_{p}. When Assumption 3 is absent, the analysis is much more convoluted. One possible way is to decompose Wp=A+BW_{p}=A+B where A=A⊺A=A^{\intercal} is symmetric and B=−B⊺B=-B^{\intercal} is skew-symmetric. We leave it for future work.

2 Analysis of decoupled dynamics

The simplified three (two) dimensional dynamics of BYOL (SimSiam) yields significant insights. First, there is clearly a collapsed fixed point at pj(t)=sj(t)=0p_{j}(t)=s_{j}(t)=0 and τ\tau taking any value. We wish to understand conditions under which pjp_{j} and sjs_{j} can avoid this collapsed fixed point and grow from small random initial conditions. Since sjs_{j} is an eigenvalue of WW⊺WW^{\intercal}, we are particularly interested in conditions under which sjs_{j} achieves large final values, corresponding to a non-collapsed online network, that are moreover sensitive to the statistics of the data, governed by σ2\sigma^{2}.

Exact integral. First, an important observation, similar to Theorem 1, is that the dynamics possesses an exact integral of motion, obtained by multiplying Eqn. 11 by 2αp−1pj2\alpha_{p}^{-1}p_{j}, subtracting, Eqn. 12 and integrating over time yielding

where cj=αp−1pj2(0)−sj(0)c_{j}=\alpha_{p}^{-1}p_{j}^{2}(0)-s_{j}(0) is fixed by initial conditions. In absence of weight decay (η=0\eta=0), this integral reveals that the initial condition encoded in cjc_{j} is never forgotten and the dynamics of pjp_{j} and sjs_{j} are confined to parabolas of the form sj(t)=pj2(t)+cjs_{j}(t)=p_{j}^{2}(t)+c_{j}, as can be seen by the blue flow lines in Fig. 3(left). With weight decay (η>0\eta>0) over time the initial condition is forgotten and the dynamics approaches the invariant parabola sj=αp−1pj2s_{j}=\alpha_{p}^{-1}p_{j}^{2} as can been seen by the approach of the blue flow lines to the black dashed parabola in Fig. 3 right and middle. We discuss these two cases in turn. First we note that in both cases, since the EMA computation is often very slow (Grill et al., 2020), corresponding to small β\beta, the dynamics of τ\tau in Eqn. 13 is slow relative to that of pjp_{j} and sjs_{j}. Therefore to understand the combined dynamics, we can search for the fixed points that pjp_{j} and sjs_{j} will rapidly approach at fixed τ\tau. Over time τ\tau will then either slowly approach 11 (BYOL) or be always equal to 11 (SimSiam), and sjs_{j} and pjp_{j} will follow their τ\tau-dependent fixed points.

No weight decay. When η=0\eta=0, Eqns. 11 and 12 at a fixed value of τ\tau yield a branch of collapsed fixed points given by sj=0s_{j}=0 and pjp_{j} taking any value, and a branch of non-collapsed fixed points, with pj=τ/(1+σ2)p_{j}=\tau/(1+\sigma^{2}) and sjs_{j} taking any value (horizontal and vertical red/green lines in Fig. 3,left). A sufficient criterion on initial conditions to avoid the collapsed branch is sj(0)>pj2(0)/αps_{j}(0)>p^{2}_{j}(0)/\alpha_{p} corresponding to lying above the dashed black parabola in Fig. 3,left. This restricted initial condition reveals why a fast predictor (large αp\alpha_{p}) is advantageous (Obs#1): larger αp\alpha_{p} leads to a smaller basin of attraction of the collapsed branch by flattening the dashed parabola. Indeed both BYOL and SimSiam have noted that a fast predictor can help avoid collapse. On the other hand, αp\alpha_{p} cannot be infinitely large (Obs#2): since sj(+∞)=sj(0)+αp−1(pj2(+∞)−pj2(0))s_{j}(+\infty)=s_{j}(0)+\alpha^{-1}_{p}(p^{2}_{j}(+\infty)-p_{j}^{2}(0)), very large αp\alpha_{p} implies that sjs_{j}, the final value of the online network characterizing the learned representation, does not grow even if pjp_{j} does. This is consistent with results which show that optimizing the predictor too often doesn’t work in SimSiam (Chen & He, 2020), and directly setting an “optimal” predictor fails as well (Tbl. 1). The online network needs to grow along with the predictor and that cannot happen if the predictor is too fast.

Advantage of weight decay. In the non-collapsed branch of fixed points without weight decay (vertical red line in Fig. 3,left), the predictor pjp_{j} takes the exact value τ/(1+σ2)\tau/(1+\sigma^{2}), which models the invariance to augmentation correctly: a large data augmentation variance σ2\sigma^{2} should lead to a small magnitude of the learned representation. Ideally, we want sjs_{j} to have the same property. With weight decay η>0\eta>0 in Eqn. 14, memory of the initial condition cjc_{j} fades away, yielding convergence to some point on the invariant parabola sj=αp−1pj2s_{j}=\alpha_{p}^{-1}p_{j}^{2}. (Obs#3): Therefore, by tying the online network to the predictor, weight decay allows sjs_{j} to also model invariance to augmentations correctly if the predictor does, regardless of the random initial condition cjc_{j}.

Because weight decay forces convergence to the invariant parabola sj=αp−1pj2s_{j}=\alpha^{-1}_{p}p_{j}^{2}, we next focus on dynamics along this parabola (i.e. cj=0c_{j}=0 in Eqn. 14). In this case, Eqn. 13 has a solution:

with initial condition τ(0)=0\tau(0)=0. Inserting the invariant sj=αp−1pj2s_{j}=\alpha_{p}^{-1}p^{2}_{j} into Eqn. 11, the dynamics of pjp_{j} is given by:

We first analyze the fixed points where p˙j=0\dot{p}_{j}=0 at fixed τ\tau.

When the weight decay 0<η≤τ24(1+σ2)0<\eta\leq\frac{\tau^{2}}{4(1+\sigma^{2})}, pjp_{j} has has three fixed points (Fig. 4(b)):

where both pj0∗p^{*}_{j0} and pj+∗p^{*}_{j+} are stable and pj−∗p^{*}_{j-} is unstable, as shown in Fig. 4(b). The basin of attraction of the collapsed fixed point pj0∗=0p^{*}_{j0}=0 is pj<pj−∗p_{j}<p^{*}_{j-} while the basin of attraction of the useful non-collapsed fixed point pj+∗p^{*}_{j+} is pj>pj−∗p_{j}>p^{*}_{j-}, yielding an important constraint on initial conditions to avoid collapse. Note that pj−∗p^{*}_{j-} is a decreasing function of τ\tau and increasing function of η\eta (see SM). This means that with larger η\eta, pj−∗p^{*}_{j-} moves right and the basin of collapse expands (Obs#4). When η>τ24(1+σ2)\eta>\frac{\tau^{2}}{4(1+\sigma^{2})} there is only one stable fixed point pj0∗=0p^{*}_{j0}=0 (Fig. 4(c)). Under such strong weight decay collapse is unavoidable (Obs#5).

We now discuss the dynamics. First we define the quantity Δj:=pj[τ−(1+σ2)pj]−η\Delta_{j}:=p_{j}[\tau-(1+\sigma^{2})p_{j}]-\eta, which must satisfy two criteria. Note that Eqn. 16 can be written as p˙j=pjΔj\dot{p}_{j}=p_{j}\Delta_{j}, so Δj\Delta_{j} must at some point be positive to drive pj(t)p_{j}(t) to any positive non-collapsed fixed point pj+∗p^{*}_{j+}. Second, for eigenspace alignment in Theorem 3 to remain stable (even if the alignment has already happened), K(t)K(t) must be positive definite (PD) in Eqn. 9. Using the eigen-space alignment conditions and the invariance sj=αp−1pj2s_{j}=\alpha^{-1}_{p}p_{j}^{2}, the positive definite condition on K(t)K(t) can be written as

This criterion and the criterion Δj>0\Delta_{j}>0 yield interesting insights into the roles of various hyperparameters choices.

First (Obs#6), larger predictor learning rate αp\alpha_{p} can play an advantageous role by loosening the upper bound in Eqn. 17, making it easier to satisfy. Second (Obs#7), increasing η\eta also has the same effect.

Role of EMA. Without EMA, τ≡1\tau\equiv 1 and (Eqn. 17) may not hold initially when pjp_{j} is small. The reason is Δj\Delta_{j} is to leading order linear in pjp_{j} when τ=1\tau=1 while the right hand side is to leading order sj∼pj2s_{j}\sim p_{j}^{2}, so the left hand side has a larger contribution from pjp_{j} than the right.

EMA resolves this as follows. When the training begins, sjs_{j} is often quite small, and τ\tau remains small since WW changes rapidly. When pjp_{j} grows to the fixed point pj+∗∼τ/(1+σ2)p^{*}_{j+}\sim\tau/(1+\sigma^{2}), the growth of sjs_{j} stops, making τ\tau larger. This in turns sets a higher fixed point goal for pjp_{j}. This process continues until the feature is stabilized and τ=1\tau=1 (Fig. 5 for details).

Therefore, EMA can serve as an automatic curriculum (Obs#8): it sets an initial small goal of τ1+σ2\frac{\tau}{1+\sigma^{2}} for pjp_{j} so Δj\Delta_{j} need only be small and positive to both drive pjp_{j} larger and satisfy Eqn. 17. Then EMA gradually sets a higher goal for pjp_{j} by increasing τ\tau, so that pjp_{j} and sjs_{j} can grow, while keeping the eigenspaces of WpW_{p} and FF aligned.

As a trade-off, a very slow EMA schedule (β\beta small) yields a slow training procedure (Obs#9) (See Fig. 5). Also small τ\tau leads to larger pj−∗p^{*}_{j-} and more eigen modes can be trapped in the collapsed basin (Obs#10).

3 Summarizing the effects of hyperparameters

We summarize the positive and negative effects of multiple hyperparameters in Tbl. 4. We next provide additional ablations and experiments to further justify our reasoning.

Different weight decay ηp\eta_{p} and ηs\eta_{s}. If we set a higher weight decay for the predictor (ηp\eta_{p}) than the online net (ηs\eta_{s}), then pjp_{j} grows slower than sjs_{j} and it is possible that the condition of Theorem 3 can still be satisfied without using EMA. Indeed Tbl. 5 shows this is the case.

Larger learning rate of the predictor αp>1\alpha_{p}>1. Our analysis predicts that one way to make symmetric WpW_{p} work with no EMA is to use αp>1\alpha_{p}>1 (i.e. Theorem 3 is more easily satisfied). Fig. 6 verifies this prediction. Moreover Table 22 in Appendix of BYOL (Grill et al., 2020) also shows that αp>1\alpha_{p}>1 is required to get BYOL working without EMA.

As a reference, Table 22 in Appendix I.2 of BYOL (Grill et al., 2020) also shows a similar trend: the learning rate of the (2-layer) predictor needs to be higher than that of the projector for strong performance in ImageNet, when EMA is absent.

A direct consequence of our theory is a new method for choosing the predictor that avoids gradient descent altogether. Instead, we estimate the correlation matrix FF of predictor inputs and directly set WpW_{p} to be a function of this, thereby avoiding both the need to align the eigenspaces of FF and WpW_{p} through optimization, and the need to initialize WpW_{p} outside the basin of collapse. As we shall see, this exceedingly simple, theory motivated method also yields better performance in practice compared to gradient-based optimization of a linear predictor.

This choice is theoretically motivated by eigenspace-alignment between WpW_{p} and FF (Theorem. 3) and convergence to the invariant parabola sj∝pj2s_{j}\propto p_{j}^{2} in Eqn. 14 with weight decay (η>0\eta>0). Here the estimate correlation matrix F^\hat{F} can be obtained by a moving average:

Hyper-parameter freq. Besides, we also evaluate a hybrid approach by introducing freq, which is how frequently eigen-decomposition is conducted for matrix F^\hat{F} to set WpW_{p}. For example, freq = 5 means that eigen decomposition is run every 5 minibatches. When WpW_{p} is not set by eigen decomposition, it is updated by regular gradient updates. freq = 1 means the eigen-decomposition is performed at every minibatch.

Tbl. 6 shows that directly computing WpW_{p} through DirectPred works better (76.77%76.77\%) than training via gradient descent (74.51%74.51\% in Tbl. 3, regular WpW_{p} with EMA). Additional regularization through ϵ\epsilon yields even better performance (77.38%77.38\%). Different ways to estimate FF (moving average or simple average) yield only small differences.

The performance of DirectPred also remains good over many more training epochs (Tbl. 8). Moreover, if we allow some gradient steps in between directly setting WpW_{p} (i.e., freq > 1), performance becomes even better (80.28%80.28\%). This might occur because the estimated F^\hat{F} may not be accurate enough and SGD can help correct it. This also mitigates the computational cost of eigen-decomposition.

The constant cjc_{j}. What happens if pj=max⁡(sj−cj,0)p_{j}=\sqrt{\max(s_{j}-c_{j},0)} with cj≠0c_{j}\neq 0? If cjc_{j} is small negative, performance is still fine but a positive cjc_{j} leads to very poor performance (Tbl. 7), likely due to many small eigen-values sjs_{j} becoming zero and therefore trapped in the collapsed basin.

Feature-dependent WpW_{p}. Note one of the advantages of using two layer predictors is that WpW_{p} can depend on the input features. We explored this idea by using a few random partitions of the input space, and within each random partition we estimated a different correlation matrix F^\hat{F}. The final F^\hat{F} is the sum of all the correlation matrices. With 66 random partitions, DirectPred achieves 78.20±0.1678.20{\pm}0.16 Top-1 accuracy after 100 epochs, closing performance gap to two-layer predictors (78.85%78.85\% in Tbl. 3). We leave a thorough analysis of the two layer setting to future work.

ImageNet experiments. We conducted additional experiments on ImageNet (Deng et al., 2009), with our own BYOL (Grill et al., 2020) implementation. We used ResNet-50 (He et al., 2016) as the backbone to produce features for a linear probe, followed by a projector and a predictor. The architecture design (e.g., feature dimensions), augmentation strategies (e.g., color jittering, blur (Chen et al., 2020a), solarization, etc.) and linear classification protocol strictly follow BYOL (Grill et al., 2020).

We experimented with two different training settings to study the generalization ability of DirectPred. In the first setting, we employ an asymmetric loss (given two views, only one view is used as the prediction target). The loss is optimized using standard SGD for 60 epochs with a batch size of 256. The second setting follows BYOL more closely, where we use a symmetrized loss, 4096 batch size and LARS optimizer (You et al., 2017), and train for 300 epochs.

The results are summarized in Tbl. 9. Both settings exhibit similar behaviors in comparison, and we take the 300-epoch results as our highlights in the following. As a baseline, the default 2-layer predictor from BYOL (with BatchNorm and ReLU, 4096 hidden dimension, 256 input/output dimension) achieves 72.5% top-1 accuracy, and 90.8% top-5 accuracy with 300-epoch pre-training. This reproduces the accuracy reported in BYOL (Grill et al., 2020). We find DirectPred can match this performance (72.4% top-1, and 91.0% top-5) without any gradient-based training by instead directly setting the (256×\times256) linear predictor weights every mini-batch. In particular for top-5 DirectPred is even 0.2% better. For a fair comparison, we also run BYOL with a learned linear predictor. We find the performance drops to 69.9%, and 89.6% respectively (2.5% gap to our method). The gap is even bigger in 60-epoch settings, up to 5.0% in top-1 (59.4% vs. 64.4%). These experiments demonstrate the success of DirectPred on STL-10 and CIFAR can also generalize and scale to ImageNet.

Discussion

Therefore, remarkably, our theoretical analysis of non-contrastive SSL, primarily centered around a 33 dimensional nonlinear dynamical system, not only yields conceptual insights into the functional roles of complex ingredients like EMA, stop-gradients, predictors, predictor symmetry, diverse learning rates, weight decay and all their interactions, but also predicts the performance patterns of many ablation studies as well as suggests an exceedingly simple DirectPred method that rivals the performance of more complex predictor dynamics in real-world settings.

Two-layer non-linear predictor. With only a linear predictor, our results on ImageNet (Tbl. 9) have already shown strong performance, on par with a default BYOL setting with a 2-layer predictor on ImageNet. One interesting question is how the dynamics changes if the predictor has 2 layers. While we don’t provide a formal analysis and the math can be quite complicated, the intuition here is that the “fat” 2-layer predictor used in practice (e.g., more (4096) hidden dimension than input/output dimensions (256), and a ReLU in between) essentially provides a large pool of initial weight directions to start with, and some of them could be “lucky draws”, that make eigen-space alignment faster. On the other hand, a 1-layer predictor with gradient updates may get stuck in local minima. Therefore, with the same number of epochs, a 2-layer predictor outperforms 1-layer, and is comparable with DirectPred which does not suffer from local minima issues.

Acknowledgements

We thank Lantao Yu for helpful discussions.

References

Appendix A Section 2

Taking partial derivative with respect to WpW_{p} and we get the gradient update rule:

After some manipulation, we finally arrive at the following gradient update rule:

Under the large batch limit, it is the same as Eqn. 41.

Note that the Lemma doesn’t include weight decay. With weight decay η\eta, it is not hard to see that we will arrive at the following slightly altered gradient flow:

The gradient update rules (Eqn. 2 and Eqn. 3) has the following invariance (where the symmetric matrix CC depends on initialization):

Adding them together and multiply both side with e2ηte^{2\eta t}:

This leads to e2ηtWW⊺=αp−1e2ηtWp⊺Wp+Ce^{2\eta t}WW^{\intercal}=\alpha_{p}^{-1}e^{2\eta t}W_{p}^{\intercal}W_{p}+C, or WW⊺=αp−1Wp⊺Wp+e−2ηtCWW^{\intercal}=\alpha_{p}^{-1}W_{p}^{\intercal}W_{p}+e^{-2\eta t}C. ∎

Let H(t)H(t) be dd-by-dd time-varying positive definite (PD) matrices whose minimal eigenvalues are bounded away from 0: inf⁡t≥0λmin⁡(H(t))≥λ0>0\inf_{t\geq 0}\lambda_{\min}(H(t))\geq\lambda_{0}>0, then the following dynamics:

satisfies ∥w(t)∥2≤e−λ0t∥w(0)∥2\|{\bm{w}}(t)\|_{2}\leq e^{-\lambda_{0}t}\|{\bm{w}}(0)\|_{2}, which means that w(t)→0{\bm{w}}(t)\rightarrow 0.

Construct the following Lyapunov function V(w):=12∥w∥22V({\bm{w}}):=\frac{1}{2}\|{\bm{w}}\|^{2}_{2}. For V(w(t))V({\bm{w}}(t)) we have:

Note that H(t)H(t) has eigen-decomposition: H(t)=∑jλj(t)uj(t)uj⊺(t)H(t)=\sum_{j}\lambda_{j}(t){\bm{u}}_{j}(t){\bm{u}}^{\intercal}_{j}(t) with all λj(t)≥λ0\lambda_{j}(t)\geq\lambda_{0} and [u1(t),u2(t),…,ud(t)][{\bm{u}}_{1}(t),{\bm{u}}_{2}(t),\ldots,{\bm{u}}_{d}(t)] forming an orthonormal bases. Therefore:

which leads to V(t)≤e−2λ0tV(0)V(t)\leq e^{-2\lambda_{0}t}V(0). That is ∥w(t)∥2≤e−λ0t∥w(0)∥2\|{\bm{w}}(t)\|_{2}\leq e^{-\lambda_{0}t}\|{\bm{w}}(0)\|_{2}. ∎

If inf⁡t≥0λmin⁡(H(t))≥λ0>0\inf_{t\geq 0}\lambda_{\min}(H(t))\geq\lambda_{0}>0, then W(t)→0W(t)\rightarrow 0.

Remark. Note that if Wa=WW_{a}=W and we choose not to use the predictor (Wp=IW_{p}=I), then no matter whether we choose to use stop-gradient or not, W(t)W(t) always goes to . The theorem above already proved that without stop gradient, it is the case. When there is stop gradient, from Eqn. 3, we have:

Note that X′+ηIX^{\prime}+\eta I is a PD matrix and with similar arguments, W(t)→0W(t)\rightarrow 0.

Appendix B Section 3

Isometric assumptions. Now we use the assumption that X=IX=I and X′=σ2IX^{\prime}=\sigma^{2}I, which leads to

here F=WXW⊺=WW⊺F=WXW^{\intercal}=WW^{\intercal}. If we also have weight decay −ηW-\eta W for WW, then we have:

or using anticommutator {A,B}:=AB+BA\{A,B\}:=AB+BA:

Fig. 7 shows that this assumption is largely correct.

Under this condition, using F=WXW⊺=WW⊺F=WXW^{\intercal}=WW^{\intercal}, the dynamics becomes (Now we also put weight decay for WpW_{p}):

Derivation of Fixed point of Eqn. 2. Given the dynamics Eqn. 62 we now want to check its fixed point:

for some PSD matrix FF. For convenience, let η′=η/αp\eta^{\prime}=\eta/\alpha_{p}. Since FF is always PSD, we have eigendecomposition F=UΛU⊺F=U\Lambda U^{\intercal}. Left-multiplying UU and right-multiplying U⊺U^{\intercal}, we have:

where Wˉp:=U⊺WpU\bar{W}_{p}:=U^{\intercal}W_{p}U. Let Λ′=(1+σ2)Λ+η′I\Lambda^{\prime}=(1+\sigma^{2})\Lambda+\eta^{\prime}I is a diagonal matrix with all positive diagonal element since η′>0\eta^{\prime}>0. Therefore, we have:

and thus Wˉp=τΛ(Λ′)−1\bar{W}_{p}=\tau\Lambda(\Lambda^{\prime})^{-1} is a symmetric matrix and so does Wp=UWˉpU⊺W_{p}=U\bar{W}_{p}U^{\intercal}. When η=0\eta=0 and FF has zero eigenvalues, WpW_{p} can have infinite solutions (or fixed points), and some of them might not be symmetric.

Symmetrization of WpW_{p}. Now we need to assume WpW_{p} is symmetric and also symmetrize its dynamics, which yields (here {A,B}:=AB+BA\{A,B\}:=AB+BA):

Note that the asymmetric dynamic might be interesting and we will leave it later.

Under the dynamics of Eqn. 67, the commutator [F,Wp]:=FWp−WpF[F,W_{p}]:=FW_{p}-W_{p}F satisfies:

If max⁡t≥0λmin⁡[K(t)]=λ0>0\max_{t\geq 0}\lambda_{\min}[K(t)]=\lambda_{0}>0, then the commutator ∥[F(t),Wp(t)]∥F≤e−2λ0t∥[F(0),Wp(0)]∥F→0\|[F(t),W_{p}(t)]\|_{F}\leq e^{-2\lambda_{0}t}\|[F(0),W_{p}(0)]\|_{F}\rightarrow 0, i.e., the eigenspace of WpW_{p} gradually aligns with FF.

Let’s compute the commutator L:=[F,Wp]:=FWp−WpFL:=[F,W_{p}]:=FW_{p}-W_{p}F and its time derivative. First we have:

is a symmetric matrix. We can write the dynamics of L(t)L(t):

where K(t)⊕K(t):=I⊗K(t)+K(t)⊗IK(t)\oplus K(t):=I\otimes K(t)+K(t)\otimes I is the Kronecker sum and is a PSD matrix if KK is PSD.

If inf⁡t≥0λmin⁡(K(t))≥λ0>0\inf_{t\geq 0}\lambda_{\min}(K(t))\geq\lambda_{0}>0 for all tt, then inf⁡t≥0λmin⁡[K(t)⊕K(t)]≥2λ0\inf_{t\geq 0}\lambda_{\min}[K(t)\oplus K(t)]\geq 2\lambda_{0}. Applying Lemma 2 and we have:

This means that WpW_{p} and FF can commute, and the eigen space of WpW_{p} and FF will gradually align. ∎

Remark. Fig. 9 shows numerical simulation of the symmetrized dynamics (Eqn. 67). If K(t)K(t) has negative eigenvalues, then even if WpW_{p} and FF have already approximately aligned, the dynamics is also unstable and might diverge due to noise and/or numerical instability.

Fig. 8 shows a numerical simulation of Eqn. 62 (dynamics with Assumption 1 and Assumption 2 but without the symmetric dynamics). We can clearly see that the asymmetric component converges to zero.

In this case, the time derivatives W˙p\dot{W}_{p} and F˙\dot{F} can all be written as decoupled form: W˙p=UG1U⊺\dot{W}_{p}=UG_{1}U^{\intercal} and F˙=UG2U⊺\dot{F}=UG_{2}U^{\intercal} where G1G_{1} and G2G_{2} are diagonal matrices. In other words, they are both decoupled into each eigen mode, and so does the future value of WpW_{p} and FF. Then UU won’t change over time.

To see why, we consider the general case where we have a symmetric matrix M(t)M(t) with eigen decomposition M(t)=U(t)D(t)U⊺(t)M(t)=U(t)D(t)U^{\intercal}(t). MM follows M˙=U(t)G(t)U⊺(t)\dot{M}=U(t)G(t)U^{\intercal}(t) where G(t)G(t) is an arbitrary diagonal matrix.

To see why U˙=0\dot{U}=0, at each time step we have:

Since U⊺(t)U(t)=IU^{\intercal}(t)U(t)=I, we have U˙⊺U+U⊺U˙=0\dot{U}^{\intercal}U+U^{\intercal}\dot{U}=0 so Q:=U⊺U˙Q:=U^{\intercal}\dot{U} is a skew-symmetric matrix and we have

Since the right hand side is a diagonal matrix, checking each entry and we have qijdj−qijdi=0q_{ij}d_{j}-q_{ij}d_{i}=0 for i≠ji\neq j. If MM has distinctive eigenvalues, then we know qij=0q_{ij}=0 for i≠ji\neq j. QQ is skew-symmetric so qii=0q_{ii}=0. So Q=U⊺U˙=0Q=U^{\intercal}\dot{U}=0 and thus U˙=0\dot{U}=0. If MM has duplicated eigenvalues, then we can show qij=0q_{ij}=0 for any di≠djd_{i}\neq d_{j}. Within high-dimensional eigenspace for duplicated eigenvalues, its eigen-decomposition is not unique and we can always pick the eigenspace within each duplicated eigenspace so that U˙=0\dot{U}=0.

Therefore, we just multiply U⊺U^{\intercal} and UU to Eqn. 67 and the system becomes decoupled. Then after some algebraic manipulation, we arrive at the following:

Multiply Eqn. 79 with 2αp−1pj2\alpha^{-1}_{p}p_{j} and subtract with Eqn. 80, we get:

Therefore, we have integral sj(t)=αp−1pj2(t)+cje−2ηts_{j}(t)=\alpha_{p}^{-1}p_{j}^{2}(t)+c_{j}e^{-2\eta t}. For finite weight decay (η>0\eta>0), we could simply expect sj(t)≈αp−1pj2(t)s_{j}(t)\approx\alpha_{p}^{-1}p_{j}^{2}(t).

On the other hand, the dynamics of τ\tau is:

When FF and WpW_{p} aligns, we have F˙\dot{F} all in the same eigen space.

So the eigenvectors UU won’t change and thus we have:

which has a close form solution when cj=0c_{j}=0. Note that in the case, we have sj=αp−1pj2s_{j}=\alpha_{p}^{-1}p_{j}^{2} and thus s˙j=2αp−1pjp˙j\dot{s}_{j}=2\alpha^{-1}_{p}p_{j}\dot{p}_{j} and we have:

B.2 Section 3.2

Monotonicity of pj−∗p_{j-}^{*} with respect to η\eta and τ\tau. Note that

is the (right) boundary of trivial basin p<pj−∗p<p_{j-}^{*} and determines the size of trivial attractive region towards pj0∗=0p_{j0}^{*}=0. It is dependent on η\eta and τ\tau. It is clear that pj−∗p_{j-}^{*} is a increasing function of η\eta. This means that if the weight decay η\eta is large, so does trivial region (and more eigenvalues will be trapped to trivial solution).

On the other hand, we can compute the derivative of g(x)=x−x2−cg(x)=x-\sqrt{x^{2}-c} for c>0c>0 and x2>cx^{2}>c:

So g(x)g(x) is a decreasing function with respect to xx. Or pj−∗p^{*}_{j-} is a decreasing function with respect to τ\tau.

Appendix C Section 4

Appendix D Analysis of BYOL and SimSiam learning dynamics without isotropic assumptions on data

Also recall that the BYOL learning dynamics, without weight decay, is given by

We first derive exact fixed point solutions to both BYOL and SimSiam learning dynamics in this setting. We then discuss specific models for data distributions and augmentation procedures, and show how the fixed point solutions depend on both data and augmentation distributions. We then discuss how our theory reveals a fundamental role for the predictor in avoiding collapse in BYOL solutions. Finally, we derive a highly reduced three dimensional description of BYOL and SimSiam learning dynamics, assuming decopuled initial conditions, that provides considerable insights into dynamical mechanisms enabling both to avoid collapsed solutions without negative pairs to force apart representations of different objects.

D.2 Illustrative models for data and data augmentation

The above section suggests that the top eigenmodes of Σd[Σs]−1\Sigma^{d}[\Sigma^{s}]^{-1} control the non-collapsed solutions. Here we make this result more concrete by giving illustrative examples of data distributions and data augmentation procedures, and the resulting properties of Σd[Σs]−1\Sigma^{d}[\Sigma^{s}]^{-1}.

Consider for example a multiplicative subspace scrambling model. In this model, data augmentation scrambles a subspace by multiplying by a random Gaussian matrix, while identically preserving the orthogonal complement of the subspace. In applications, the scrambled subspace could correspond to a space of nuisance features, while the preserved subspace could correspond to semantically important features. Indeed many augmentation procedures, including random color distortions and blurs, largely preserve important semantic information, like object identity in images.

Additive scrambling.

We also consider, as an illustrative example, data augmentation procedures which simply add Gaussian noise with a prescribed noise covariance matrix Σn\Sigma^{n}. Under this model, we have Σs=Σx+Σn\Sigma^{s}=\Sigma^{x}+\Sigma^{n} while Σd=Σx\Sigma^{d}=\Sigma^{x}. Thus in this setting, BYOL learns principal eigenmodes of Σd[Σs]−1=Σx[Σx+Σn]−1\Sigma^{d}[\Sigma^{s}]^{-1}=\Sigma^{x}[\Sigma^{x}+\Sigma^{n}]^{-1}. Thus intuitively, dimensions with larger noise variance are attenuated in learned BYOL representations. On the otherhand, correlations in the data that are not attenuated by noise are preferentially learned, but the degree to which they are learned is not strongly influenced by the magnitude of the data correlation (i.e. consider dimensions that lie along small eigenvalues of Σn\Sigma^{n}). Note that in the main paper we focused on the case where Σx=I\Sigma^{x}=I and Σn=σ2I\Sigma^{n}=\sigma^{2}I.

D.3 The importance of the predictor in BYOL and SimSiam.

Here we note that our theory explains why the predictor plays a crucial role in BYOL and SimSiam learning in this simple setting, as is observed empirically in more complex settings. To see this, we can model the removal of the predictor by simply setting Wp=IW_{p}=I in all the above equations. The fixed point solutions then obey W=WΣd[Σs]−1W=W\Sigma^{d}[\Sigma^{s}]^{-1}. This will only have nontrivial, non-collapsed solutions if Σd[Σs]−1\Sigma^{d}[\Sigma^{s}]^{-1} has eigenvectors with eigenvalue 11. Rows of WW consisting of linear combinations of these eigenvectors will then constitute non-collapsed solutions.

Thus overall, in this simple setting, our theory provides conceptual insight into how the introduction of a predictor is crucial for creating new non-collapsed solutions for both BYOL and SimSiam, even though the predictor confers no new expressive capacity in allowing the online network to match the target network.

D.4 Reduction of BYOL learning dynamics to low dimensions

We note again, that the generic existence of these non-collapsed solutions in Fig. 10 depends critically on the presence of a predictor with adjustable weights wpw_{p}. Removing the predictor corresponds to forcing wp=1w_{p}=1, and non-collapsed solutions cannot exist unless λd=λs\lambda_{d}=\lambda_{s}, as demonstrated in Fig. 10 (right). Thus, remarkably, in BYOL in this simple setting, the introduction of a predictor network plays a crucial role, even though it neither adds to the expressive capacity of the online network, nor improves its ability to match the target network. Instead, it plays a crucial role by dramatically modifying the learning dynamics (compare e.g. Fig 10 middle and right panels), thereby enabling convergence to noncollapsed solutions through a dynamical mechanism whereby the online and predictor network cooperatively amplify each others’ weights to escape collapsed solutions ( Fig. 10 (left)).

Overall, this analysis of BYOL learning dynamics provides considerable insight into the dynamical mechanisms enabling BYOL to avoid collapsed solutions, without negative pairs to force apart representations, in what is likely to be the simplest nontrivial setting. Further analysis on this model, in direct analogy to the analysis performed on the equivalent 33 dynamical system (derived under different assumptions) studied in the main paper, can yield similar insights into the dynamics of BYOL and SimSiam under various conditions on learning rates.