Catastrophic Fisher Explosion: Early Phase Fisher Matrix Impacts Generalization

Stanislaw Jastrzebski, Devansh Arpit, Oliver Astrand, Giancarlo Kerg, Huan Wang, Caiming Xiong, Richard Socher, Kyunghyun Cho, Krzysztof Geras

Introduction

The exact mechanism behind implicit regularization effects in training of deep neural networks (DNNs) remains an extensively debated topic despite being considered a critical component in their empirical success (Neyshabur, 2017; Zhang et al., 2017; Jiang et al., 2020b). For instance, it is commonly observed that using a moderately large learning rate in the early phase of training results in better generalization (LeCun et al., 2012; Salakhutdinov, 2014; Jiang et al., 2020b; Bjorck et al., 2018).

Recent work suggests that the early phase of training of DNNs might hold the key to understanding some of these implicit regularization effects. In particular, the learning rate in the early phase of training has a dramatic effect on the local curvature of the loss function (Jastrzebski et al., 2020; Cohen et al., 2021). It has also been found that when using a small learning rate, the local curvature of the loss surface increases along the optimization trajectory until optimization is close to instability.That is, a small further increase in the local curvature is not possible without divergence. Interestingly, when training using gradient descent with a learning rate η\eta, the largest eigenvalue of the Hessian of the training loss was observed to reach the critical value of 2η\frac{2}{\eta} (Cohen et al., 2021), at which training oscillates along the eigenvector corresponding to the largest eigenvalue of the Hessian.

These observations lead to a natural question: does the instability, and the corresponding dramatic change in the local curvature, in the early phase of training influence generalization? We investigate this question through the lens of the Fisher Information Matrix (FIM), a matrix that can be seen as approximating the local curvature of the loss surface (Martens, 2020; Thomas et al., 2020). Achille et al. (2019); Jastrzebski et al. (2019); Golatkar et al. (2019); Lewkowycz et al. (2020) independently suggest that effects of the early phase of training on the local curvature critically influence the final generalization, but did not directly test this proposition.

Implicit and explicit regularization of the Fisher Information Matrix

Fisher Penalty

where (x,y)(\bm{x},y) is a mini-batch of size BB, y^i\hat{y}_{i} is sampled from pθ(y∣xi)p_{\bm{\theta}}(y|\bm{x}_{i}), α\alpha is a hyperparameter. Importantly, the equation does not involve target labels. Finally, we compute the gradient of the second term only every 10 optimization steps, and in a given iteration use the most recently computed gradient. We refer to this regularizer as Fisher penalty (FP).

Catastrophic Fisher Explosion

Why Fisher Information Matrix

The benefit of the FIM is that it can be efficiently regularized during training. In contrast, Wen et al. (2018); Foret et al. (2021) had to rely on certain approximations to efficiently regularize curvature. Equally importantly, the FIM is related to the gradient norm, and as such its effect on learning is more interpretable. We will leverage this connection to argue that FP slows down learning on noisy examples in the dataset.

Concurrent work on a different mechanism

Barrett & Dherin (2021); Smith et al. (2021) concurrently argue that the implicit regularization effects in SGD can be expressed as a form of a gradient norm penalty. The key difference is that we link the mechanism of implicit regularization to the fact optimization is close to instability in the early phase. Supporting this view, we observe that FP empirically performs better than gradient penalty, and that FP works best when applied from the start of training.

We run experiments in two settings: (1) ResNet-18 with Fixup (He et al., 2016; Zhang et al., 2019) trained on the ImageNet dataset (Deng et al., 2009), (2) ResNet-26 initialized as in (Arpit et al., 2019) and trained on the CIFAR-10 and CIFAR-100 datasets (Krizhevsky, 2009). We train each architecture using SGD, with various values of η\eta, SS, and random seed.

Results

Fisher Penalty

We use a similar setting as in the previous section, but we include larger models. We run experiments using Wide ResNet (Zagoruyko & Komodakis, 2016) (depth 44 and width 3, with or without BN layers), SimpleCNN (without BN layers), DenseNet (L=40, K=12) (Huang et al., 2017) and VGG-11 (Simonyan & Zisserman, 2015). We train these models on either the CIFAR-10 or the CIFAR-100 datasets. Due to larger computational cost, we replace ImageNet with the TinyImageNet dataset (Le & Yang, 2015) (with images scaled to 32×3232\times 32 resolution) in these experiments.

Fisher Penalty improves generalization

Table 1 summarizes the results of the main experiment. First, we observe that a suboptimal learning rate (10-30x lower than the optimal) leads to dramatic overfitting. We observe a degradation of up to 9% in test accuracy, while achieving approximately 100% training accuracy (see Table 6 in the Supplement).

Fisher penalty closes the gap in test accuracy between the small and optimal learning rate, and even achieves better performance than the optimal learning rate. A similar performance was observed when minimizing ∥gr∥2\left\|g_{r}\right\|^{2}. We will come back to this observation in the next section.

We also investigate whether similar conclusions hold in large batch size training. In experiments with CIFAR-10 and SimpleCNN, we find that we can close the generalization gap due to training with a large batch size by using Fisher Penalty. We provide further details in Supplement E.

In the second experimental setting, we apply FP to a network trained with the optimal learning rate η∗\eta^{*}. According to Table 2 (see Table 7 for training accuracies), Fisher Penalty improves generalization in 4 out of 5 settings. The gap between the baseline and FP is relatively small in 3 out of 5 settings (below 2%), which is natural given that we already regularize training implicitly by using the optimal η\eta..

Geometry and generalization in the early phase of training

1 Fisher Penalty Reduces Memorization

To study whether the above happens in practice, we compare FP to GPx, GPr, and mixup (Zhang et al., 2018). While mixup is not the state-of-the-art approach for learning with noisy labels, it is competitive among approaches that do not require additional data nor multiple stages of training. In particular, it is a component in several state-of-the-art methods (Li et al., 2020; Song et al., 2020). For gradient norm based regularizers, we evaluate 6 different hyperparameter values spaced uniformly on a logarithmic scale, and for mixup we evaluate β∈{0.2,0.4,0.8,1.6,3.2,6.4}\beta\in\{0.2,0.4,0.8,1.6,3.2,6.4\}. We experiment with the Wide ResNet and VGG-11 models. We describe remaining experimental details in Supplement I.3.

To test whether FP reduces the speed with which the noisy examples are learned, we track the training and validation accuracy, and the gradient norm. We compute these metrics separately on the noisy and clean examples. Figure 5 summarizes the results for VGG-11, and we show results for ResNet-32 in Figure 11 in the Supplement.

We observe that FP limits the ability of the model to memorize data more strongly than it limits its ability to learn from clean data. Figure 5 confirms applying FP results in training accuracy on noisy examples being lower for the same accuracy on clean examples, compared to the baseline.

Taken together, the results suggest implicit or explicit regularization of FP improves generalization at least in part by strengthening the bias of SGD to learn clean examples before learning noisy examples (Arpit et al., 2017). We also note that related conclusions about the effect using large learning rates were reached by Jastrzebski et al. (2017); Li et al. (2019).

Results

Related Work

Implicit regularization effects are critical to the empirical success of DNNs (Neyshabur, 2017; Zhang et al., 2017). Much of it is attributed to the choice of hyperparameters in SGD (Keskar et al., 2017; Smith & Le, 2018; Li et al., 2019), low complexity bias induced by gradient descent (Xu, 2018; Jacot et al., 2018; Arora et al., 2019; Hu et al., 2020), the cross-entropy loss function (Poggio et al., 2018; Soudry et al., 2018), or the importance of the early phase of training (Achille et al., 2019; Jastrzebski et al., 2020; Fort et al., 2020a; Frankle et al., 2020; Golatkar et al., 2019; Lewkowycz et al., 2020). However, developing a mechanistic understanding of how SGD implicitly regularizes DNNs remains a largely unsolved problem.

Many prior works have proposed explicit regularizers aimed at finding low curvature solutions (Hochreiter & Schmidhuber, 1997). Chaudhari et al. (2017) proposed a Langevin dynamics based algorithm. Wen et al. (2018); Izmailov et al. (2018); Foret et al. (2021) propose finding wide minima through approximations involving averaging gradients or parameters at the neighborhood of the current parameter state. In contrast, our focus is on elucidating a mechanistic link behind the implicit regularization effects in SGD and the early phase of training. To corroborate our hypothesis, we develop Fisher Penalty, an efficient and novel explicit regularizer that aims to reproduce the regularization effect of training with a large learning rate.

Our work contributes to a better understanding of the role of the FIM in training and generalization of deep neural networks. The FIM was also used to define complexity measures such as the Fisher-Rao norm (Karakida et al., 2019; Liang et al., 2019), and can be seen as approximating the local curvature of the loss surface (Martens, 2020; Thomas et al., 2020). Most notably, the FIM defines the distance metric used in natural gradient descent (Amari, 1998).

Chatterjee (2020); Fort et al. (2020b) show that SGD avoids memorization by extracting commonalities between examples due to following gradient descent directions shared between examples. Other works have also connected the implicit regularization effects of using large learning rates with preventing memorization (Arpit et al., 2017; Li et al., 2019; Jastrzebski et al., 2017). Liu et al. (2020) was first to note the difference in gradient norms between noisy and clean examples that emerges in the early phase of training. Our work is complementary to these findings. We make a direct connection between memorization, the instability in the early phase of training, and implicit regularization effects arising from using large learning rates in SGD.

Conclusion

The dramatic instability and changes in the curvature that happen in the early phase of training motivated us to probe its importance for generalization (Jastrzebski et al., 2020; Cohen et al., 2021). We investigated if these effects might explain some of the implicit regularization effects attributed to SGD such as the poorly understood generalization benefit of using large learning rates.

Developing theory that is fully consistent with our findings is an interesting topic for the future. Notably, existing theoretical works generally do not connect implicit regularization effects in SGD to the fact optimization is close to instability in the early phase of training (Li et al., 2017; Chaudhari & Soatto, 2018; Smith et al., 2021).

Another exciting topic for the future is connecting these findings to shortcut learning (Geirhos et al., 2020). The tendency of SGD to learn the simplest patterns in the datasets can be detrimental to the broader generalization of the model (Nam et al., 2020). We hope that by better understanding implicit regularization effects in SGD, our work will contribute to developing optimization methods that better optimize for both in and out of distribution generalization.

Limitations

Acknowledgments

GK acknowledges support in part by the FRQNT Strategic Clusters Program (2020-RS4-265502 - Centre UNIQUE - Union Neurosciences & Artificial Intelligence - Quebec). KC was supported by Samsung Advanced Institute of Technology (Next Generation Deep Learning: from pattern recognition to AI) and Samsung Research (Improving Deep Learning using Latent Structure). We thank Catriona C. Geras for proofreading the paper.

References

Apppendix

Appendix A Additional results

In this section, we present the additional experimental results for Section 3. Figure 7 shows the experiments with varying batch size for CIFAR-100 and CIFAR-10. The conclusions are the same as discussed in the main text in Section 3. We also show the training accuracy for all the experiments performed in Figure 2(c) and Figure 7. They are shown in Figure 8 and Figure 9 respectively. Most runs in all these experiments reach training accuracy ∼\sim99% and above.

A.2 Fisher Penalty

Figure 10 complements Figure 4 for the other two models on the CIFAR-10 and CIFAR-100 datasets. The figures are in line with the results of the main text.

Lastly, in Table 7 we report the final training accuracy reached by runs reported in Table 2 in the main text.

A.3 Fisher Penalty Reduces Memorization

In this section, we include additional experimental results for Section 4.1. Figure 11 is the same as Figure 5, but for ResNet-50. Finally, we show additional metrics for the experiments involving 25% noisy examples. Figure 12 shows the cosine between the mini-batch gradients computed on the noisy and clean data. In Table 9 and Table 8 we show training accuracy on the noisy and clean examples in the final epoch of training.

In this section, we present additional experimental results for Section 5. The experiment on CIFAR-10 is shown in Figure 13. The conclusions are the same as discussed in the main text in Section 5.

Appendix C Approximations in Fisher Penalty

where NN and MM are the minibatch size and the number of samples from pθ(y∣xn)p_{\theta}(y|\bm{x}_{n}), respectively. This greatly improves the computational efficiency. With N=BN=B and M=1M=1, we end up with the following learning objective function:

Specifically, we augment the loss function with the norm of the gradient computed on the first example in the mini-batch as follows

We apply this penalty in each optimization step. We tune the hyperparameter α\alpha, checking 10 values equally spaced between 10−410^{-4} and 10−210^{-2} on a logarithmic scale.

Appendix D A closer look at the surprising effect of learning rate on the loss geometry in the early phase of training

We also found similar to hold when varying the batch size, see Section E, which further shows that the observed effects cannot be explained by the difference in learning speed incurred by using a small learning rate.

To summarize, both the published evidence of Jastrzebski et al. (2020); Lewkowycz et al. (2020); Cohen et al. (2021), as well as our additional experiments, are inconsistent with the hypothesis that the results in this paper can be explained by differences in training speed between experiments using large and small learning rates.

Appendix E Catastrophic Fisher Explosion holds in training with large batch size

In this section, we show preliminary evidence that the conclusions transfer to large batch size training. Namely, we show that (1) catastrophic Fisher explosion also occurs in large batch size training, and (2) Fisher Penalty can improve generalization and close the generalization gap due to using a large batch size (Keskar et al., 2017).

Next, we run a variant of one of the experiments in Table 1. Instead of using a suboptimal (smaller) learning rate, we use a suboptimal (larger) batch size. Specifically, we train SimpleCNN on the CIFAR-10 dataset (without augmentation) with a 10x larger batch size while keeping the learning rate the same. Using a larger batch size results in 3.24%3.24\% lower test accuracy (76.9476.94% compared to 73.7%73.7\% test accuracy, c.f. with Table 1).

Taken together, the results suggest that Catastrophic Fisher explosion holds in large batch size training; using a small batch size improves generalization by a similar mechanism as using a large learning rate, which can be introduced explicitly in the form of Fisher Penalty.

Appendix G Relationship between Fisher Penalty and gradient norm penalty

Appendix H Fisher Penalty Reduces Memorization

We present here a short argument that Fisher Penalty can be seen as reducing the training speed of examples that are both labeled randomly and for which the model makes a random prediction.

Appendix I Additional Experimental Details

Here, we describe additional details for experiments in Section 3.

In the experiments with batch size, for CIFAR-10, we use batch sizes 100, 500 and 700, and ϵ=1.2\epsilon=1.2. For CIFAR-100, we use batch sizes 100, 300 and 700, and ϵ=3.5\epsilon=3.5. These thresholds are crossed between 2 and 7 epochs across different hyperparameter settings. The remaining details for CIFAR-100 and CIFAR-10 are the same as described in the main text. The optimization details for these datasets are as follows.

I.2 Fisher Penalty

Here, we describe the remaining details for the experiments in Section 4. We first describe how we tune hyperparameters in these experiments. In the remainder of this section, we describe each setting used in detail.

In all experiments, we refer to the optimal learning rate η∗\eta^{*} as the learning rate found using grid search. In most experiments, we check 5 different learning rate values uniformly spaced on a logarithmic scale, usually between 10−210^{-2} and 10010^{0}. In some experiments, we adapt the range to ensure that it includes the optimal learning rate. We tune the learning rate only once for each configuration (i.e. we do not repeat it for different random seeds).

DenseNet on the CIFAR-100 dataset

We use the DenseNet (L=40, k=12) configuration from (Huang et al., 2017). We largely follow the experimental setting in (Huang et al., 2017). We use the standard data augmentation (where noted) and data normalization for CIFAR-100. We hold out random 5000 examples as the validation set. We train the model using SGD with a momentum of 0.9, a batch size of 128, and a weight decay of 0.0001. Following (Huang et al., 2017), we train for 300 epochs and decay the learning rate by a factor of 0.1 after epochs 150 and 225. To reduce variance, in testing we update Batch Normalization statistics using 100 batches from the training set.

Wide ResNet on the CIFAR-100 dataset

We train Wide ResNet (depth 44 and width 3, without Batch Normalization layers). We largely follow experimental setting in (He et al., 2016).We use the standard data augmentation and data normalization for CIFAR-100. We hold out random 5000 examples as the validation set. We train the model using SGD with a momentum of 0.9, a batch size of 128, and a weight decay of 0.0010. Following (He et al., 2016), we train for 300 epochs and decay the learning rate by a factor of 0.1 after epochs 150 and 225. We remove Batch Normalization layers. To ensure stable training we use the SkipInit initialization (De & Smith, 2020).

VGG-11 on the CIFAR-100 dataset

We adapt the VGG-11 model (Simonyan & Zisserman, 2015) to CIFAR-100. We do not use dropout nor Batch Normalization layers. We hold out random 5000 examples as the validation set. We use the standard data augmentation (where noted) and data normalization for CIFAR-100. We train the model using SGD with a momentum of 0.9, a batch size of 128, and a weight decay of 0.0001. We train the model for 300 epochs and decay the learning rate by a factor of 0.1 after every 40 epochs starting from epoch 80.

SimpleCNN on the CIFAR-10 dataset

We also run experiments on the CNN example architecture from the Keras example repository (Chollet & others, 2015)Accessible at https://github.com/keras-team/keras/blob/master/examples/cifar10_cnn.py., which we change slightly. Specifically, we remove dropout and reduce the size of the final fully-connected layer to 128. We train it for 300 epochs and decay the learning rate by a factor of 0.1 after the epochs 150 and 225. We train the model using SGD with a momentum of 0.9, and a batch size of 128.

Wide ResNet on the TinyImageNet dataset

We train Wide ResNet (depth 44 and width 3, with Batch Normalization layers) on TinyImageNet Le & Yang (2015). TinyImageNet consists of a subset of 100,000 examples from ImageNet that we downsized to 32×\times32 pixels. We train the model using SGD with a momentum of 0.9, a batch size of 128, and a weight decay of 0.0001. We train for 300 epochs and decay the learning rate by a factor of 0.1 after epochs 150 and 225. We do not use validation in TinyImageNet due to its larger size. To reduce variance, in testing we update Batch Normalization statistics using 100 batches from the training set.

I.3 Fisher Penalty Reduces Memorization

Here, we describe additional experimental details for Section 4.1. We use two configurations described in Section I.2: VGG-11 trained on CIFAR-100 dataset, and Wide ResNet trained on the CIFAR-100 dataset. We tune the regularization coefficient α\alpha in the range {0.01,0.1,0.31,10}\{0.01,0.1,0.31,10\}, with the exception of GPx for which we use the range {10,30,100,300,1000}\{10,30,100,300,1000\}. We tuned the mixup coefficient in the range {0.4,0.8,1.6,3.2,6.4}\{0.4,0.8,1.6,3.2,6.4\}. We removed weight decay in these experiments. We use validation set for early stopping, as commonly done in the literature.