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 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 ( and ) 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— achieved without BN vs. 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 top-1 accuracy.
Background
In this section, we adopt the notation of . Recall that denotes an image and and are two views of obtained from two independent transformations and sampled from a distribution . These views are used as input to an encoder network to obtain representations and ; and projections and . 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 are projections from views of all images (including but not ), and 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 ) and a target network (parameterized by ). As a part of the online network, it further defines a predictor network that is used to predict target projections using online projections as inputs. Accordingly, the parameters of the online projection are updated following the gradients of the prediction loss
In turn, the target network weights are updated as an exponential moving average of the online network’s weights, i.e. , with being a decay parameter. As is a function of and is a function of , BYOL’s loss can be seen as a measure of similarity between the views and 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 of dimensions , GN first splits channels into equally-sized groups, then normalizes activations with the mean and standard deviation computed over disjoint slices of size . The number of groups thus trades off between normalization over all channels (, equivalent to LN), and normalization over a single one (, 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 is normalized to get a new weight matrix which is directly used in place of during training. Only the normalized weights are used to compute convolution outputs but the loss is differentiated with respect to non-normalized weights
where is the input dimension (product of input channel dimension and kernel spatial dimension); we set . 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- 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 and trainable, and initialize them as
We use the exact same hyperparameters as for vanilla BYOL (i.e., base learning rate of , weight decay of and decay rate of ), except that we increase the number of warmup epochs from to . After epochs, this representation achieves % top-1 accuracy in the linear evaluation setting compared to 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- 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 instead of , the base learning rate set to instead of and the target update rate, set to instead of ; we also set the number of groups for GN to . With this setup, BYOL (+GN +WS) achieves 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 top-1 accuracy when removing BN and changing the initialization. Moreover, BYOL achieves a competitive 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: .