BYOL works even without batch statistics

Pierre H. Richemond, Jean-Bastien Grill, Florent Altché, Corentin Tallec, Florian Strub, Andrew Brock, Samuel Smith, Soham De, Razvan Pascanu, Bilal Piot, Michal Valko

Introduction

Self-supervised image representation methods have achieved downstream performance that rivals those of supervised pre-training on ImageNet . Current self-supervised methods rely on image transformations to generate different views from an input image while preserving semantic information. Among the most successful algorithms, contrastive methods use a loss function that balances out two terms: a term associated to the positive pairs (that we refer to as the positive term) encouraging representations from views of the same image to be similar, and a term associated to negative pairs (a negative term) which encourages representations to be spread out.

Taking a different route, other approaches manage to avoid the contrastive paradigm . Among them, BYOL learns its representation by predicting the target network representation of a view from the online representation of another view of the same image. However, such a setup has obvious collapsed equilibria where the representation is constant, and thus can be predicted from any input. This has raised the question of how BYOL could even work without a negative term nor an explicit mechanism to prevent collapse. Experimental reports suggest that the use of batch normalization, BN , in BYOL’s network is crucial to achieve good performance. These reports hypothesise that the BN used in BYOL’s network could implicitly introduce a negative term.

We experimentally confirm the particular importance of BN in BYOL: removing all instances of BN in the network prevents BYOL from learning anything at all in the classic setting, see Section 3.1 and Table 1. However, our experimental results given in Table 2 go against some interpretations proposed notably in . In particular, we refute the following hypotheses:

(H1) BYOL needs BN because BN provides an implicit negative term required to avoid collapse. In Section 3.2, we show that BYOL avoids collapse and achieves 65.7%65.7\% top-1 accuracy on ImageNet under the linear evaluation protocol without using any normalization during training, by using both a better initialization scheme and retaining the additional trainable parameters scaling and bias (γ\gamma and β\beta) introduced by BN.

Therefore, unlike (H1), we hypothesize that the main role of BN is to make the network more robust to cases when the initialization is scaled improperly. Indeed, proper initialization is critical for deep nets and BYOL suffers from a bad initialization in two ways: (i) as for any deep network, it makes optimization difficult and (ii) BYOL’s target network outputs will be ill-conditioned, which initially provides poor targets for the online network.

(H2) BYOL cannot achieve competitive performance without the implicit contrastive effect provided by batch statistics. In Section 3.3, we show that most of this performance gap—65.7%65.7\% achieved without BN vs. 74.3%74.3\% achieved with BN—can be bridged without using batch statistics. Specifically, if we replace BN with a combination of group normalization, GN , and weight standardization, WS , while keeping standard initializations, BYOL achieves 73.9%73.9\% top-1 accuracy.

Background

In this section, we adopt the notation of . Recall that xx denotes an image and v=t(x)v=t(x) and v′=t′(x)v^{\prime}=t^{\prime}(x) are two views of xx obtained from two independent transformations tt and t′t^{\prime} sampled from a distribution T\mathcal{T}. These views are used as input to an encoder network to obtain representations yθ=fθ(v)y_{\theta}=f_{\theta}(v) and yθ′=fθ(v′)y^{\prime}_{\theta}=f_{\theta}(v^{\prime}); and projections zθ=gθ(yθ)z_{\theta}=g_{\theta}(y_{\theta}) and zθ′=gθ(yθ′)z^{\prime}_{\theta}=g_{\theta}(y^{\prime}_{\theta}). We continue by a brief recap of standard contrastive methods and BYOL.

Most contrastive methods use variants of the InfoNCE loss to train their representation,

where the zθiz_{\theta}^{i} are projections from views of all images (including zθ′z^{\prime}_{{\theta}} but not zθz_{{\theta}}), and τ\tau is a temperature parameter. The first term of this loss (the positive term) encourages projections of views of the same image to become similar, while the second term (the negative term) makes projections of views from different images more dissimilar. Such a loss has a strong theoretical underpinning: minimizing this loss is equivalent to maximizing a lower bound on the mutual information between the representation of two views which is tight when the function approximator is sufficiently expressive.

BYOL

BYOL trains its representation using both an online network (parameterized by θ{\theta}) and a target network (parameterized by ξ\xi). As a part of the online network, it further defines a predictor network qθq_{\theta} that is used to predict target projections zξ′z^{\prime}_{\xi} using online projections zθz_{{\theta}} as inputs. Accordingly, the parameters of the online projection are updated following the gradients of the prediction loss

In turn, the target network weights ξ\xi are updated as an exponential moving average of the online network’s weights, i.e. ξ←(1−η)ξ+ηθ\xi\leftarrow(1-\eta)\xi+\eta{\theta}, with η\eta being a decay parameter. As qθ(zθ)q_{\theta}(z_{\theta}) is a function of vv and zξ′z^{\prime}_{\xi} is a function of v′v^{\prime}, BYOL’s loss can be seen as a measure of similarity between the views vv and v′v^{\prime} and therefore resembles the positive term of the InfoNCE loss.

Group normalization (GN)

GN is an activation normalization method, like BN , layer normalization (LN ), and instance normalization (IN ). For an activation tensor XX of dimensions (N,H,W,C)(N,H,W,C), GN first splits channels into GG equally-sized groups, then normalizes activations with the mean and standard deviation computed over disjoint slices of size (1,H,W,C/G)(1,H,W,C/G). The number of groups GG thus trades off between normalization over all channels (G=1G=1, equivalent to LN), and normalization over a single one (G=CG=C, equivalent to IN). Importantly, GN operates independently on each batch element and therefore it does not rely on batch statistics.

Weight standardization (WS)

WS normalizes the weights corresponding to each activation using weight statistics. Each row of the weight matrix WW is normalized to get a new weight matrix W^\widehat{W} which is directly used in place of WW during training. Only the normalized weights W^\widehat{W} are used to compute convolution outputs but the loss is differentiated with respect to non-normalized weights W,W,

where I\mathcal{I} is the input dimension (product of input channel dimension and kernel spatial dimension); we set ε=10−4\varepsilon=10^{-4}. Contrary to BN, LN, and GN, WS does not create additional trainable weights.

Experimental results

While many metrics can be used to evaluate self-supervised representations, we focus on classification accuracy on ImageNet under the standard linear evaluation protocol with a ResNet-5050 architecture, with the same setup as . Unless otherwise specified, we follow the training setup and hyperparameters described in when training BYOL.

In Table 1, we explore the impact of using different normalization schemes in SimCLR and BYOL, by using either BN, LN, or removing normalization in each component of BYOL and SimCLR, i.e., the encoder, the projector (for SimCLR and BYOL), and the predictor (for BYOL only). First, we observe that removing all instances of BN in BYOL leads to performance that is no better than random. Noticeably, this is specific to BYOL as SimCLR still performs reasonably well in this regime. Nevertheless, solely applying BN to the ResNet encoder is enough for BYOL to achieve high performance. Some of these observations differ from the ones initially reported in . Specifically, the authors observed a collapse when removing BN in BYOL’s predictor and projector. This difference could be linked to the use of the SGD optimizer instead of LARS .

From these observations, hypothesizes that BN implicitly introduces a negative contrastive term, which acts as a crucial component to stabilize training (H1). This hypothesis may seem further supported by the performance difference between SimCLR and BYOL when replacing BN (which uses batch statistics) with LN which does not.

However, we observe that BN seems to be mainly useful in the ResNet encoder, for which standard initializations are known to lead to poor conditioning . Also BYOL might be even more affected by improper initialization as it creates its own targets. Rather than (H1), we therefore hypothesize that the main contribution of BN in BYOL is to compensate for improper initialization.

2 Proper initialization allows working without BN

To confirm this assumption, we design the following protocol to mimic the effect of BN on initial scalings and training dynamics, without using or backpropagating through batch statistics. Before training, we compute per-activation BN statistics for each layer by running a single forward pass of the network with BN on a batch of augmented data. We then remove then batch normalization layers, but retain the scale and offset parameters γ\gamma and β\beta trainable, and initialize them as

We use the exact same hyperparameters as for vanilla BYOL (i.e., base learning rate of 0.20.2, weight decay of 1.5⋅10−61.5\cdot 10^{-6} and decay rate of 0.9960.996), except that we increase the number of warmup epochs from 1010 to 5050. After 10001000 epochs, this representation achieves 65.765.7% top-1 accuracy in the linear evaluation setting compared to 74.3%74.3\% for the baseline. These results are reported in Table 2.

Despite its comparatively low performance, the trained representation still provides considerably better classification results than a random ResNet-5050 backbone, and is thus necessarily not collapsed. This confirms that BYOL does not need BN to prevent collapse. It also confirms that one of the effects of BN is to provide better initial scalings and training dynamics, and that, contrary to SimCLR, these are required for BYOL to perform well.

3 Using GN with WS leads to competitive performance

In the previous section, we have shown that BYOL can learn a non-collapsed representation without using BN. Yet, BYOL performs worse in this regime. This only disproves (H1), but BN could still both provide better initial scaling and an implicit contrastive term, responsible for some of the performance. To study this hypothesis, we explore other refined element-wise normalization procedures. More precisely, we apply weight standardization to convolutional and linear parameters by weight standardized alternatives, and replace all BN by GN layers.

To train the network, we use the same hyperparameters as in BYOL except for the weight decay, set to 3⋅10−83\cdot 10^{-8} instead of 1.5⋅10−61.5\cdot 10^{-6}, the base learning rate set to 0.240.24 instead of 0.20.2 and the target update rate, set to 0.9990.999 instead of 0.9960.996; we also set the number of groups for GN to G=16G=16. With this setup, BYOL (+GN +WS) achieves 73.9%73.9\% top-1 accuracy after 1000 epochs.

As neither GN nor WS compute batch statistics, this version of BYOL cannot compare elements from the batch, and therefore it likewise cannot implement a batch-wise implicit contrastive mechanism. Therefore, we experimentally show that BYOL can maintain most of its performance even without a hypothetical implicit contrastive term provided by BN.

Conclusion

Unlike contrastive methods, the loss used in BYOL does not explicitly include a negative term that would encourage its representations to spread apart. Nonetheless, BYOL’s representation does not collapse during training, and BN has been hypothesized to fill the crucial role of an implicit negative term by leaking batch statistics into the gradient. We refute this hypothesis, and show that BYOL can achieve competitive results without using batch statistics. In particular, BYOL achieves 65.7%65.7\% top-1 accuracy when removing BN and changing the initialization. Moreover, BYOL achieves a competitive 73.9%73.9\% top-1 accuracy by replacing BN with a normalization scheme operating element-wise.

Acknowledgement

The authors would like to thank the following people for their help throughout the process of writing this paper, in alphabetical order: Jean-Baptiste Alayrac, Bernardo Avila Pires, Nathalie Beauguerlange, Elena Buchatskaya, Jeffrey De Fauw, Sander Dieleman, Carl Doersch, Mohammad Gheshlaghi Azar, Zhaohan Daniel Guo, Olivier Henaff, Koray Kavukcuoglu, Pauline Luc, Katrina McKinney, Rémi Munos, Aaron van den Oord, Jason Ramapuram, Adria Recasens, Karen Simonyan, Oriol Vinyals and the DeepMind team. We would like to also thank the authors of the following papers for fruitful discussions: .

References