Fluctuation-dissipation relations for stochastic gradient descent

Sho Yaida

Introduction

Equilibration rules the long-term fate of many macroscopic dynamical systems. For instance, as we pour water into a glass and let it be, the stationary state of tranquility is eventually attained. Zooming into the tranquil water with a microscope would reveal, however, a turmoil of stochastic fluctuations that maintain the apparent stationarity in balance. This is vividly exemplified by the Brownian motion (Brown, 1828): a pollen immersed in water is constantly bombarded by jittery molecular movements, resulting in the macroscopically observable diffusive motion of the solute. Out of the effort in bridging microscopic and macroscopic realms through the Brownian movement came a prototype of fluctuation-dissipation relations (Einstein, 1905; Von Smoluchowski, 1906). These relations quantitatively link degrees of noisy microscopic fluctuations to smooth macroscopic dissipative phenomena and have since been codified in the linear response theory for physical systems (Onsager, 1931; Green, 1954; Kubo, 1957), a cornerstone of statistical mechanics.

Machine learning begets another form of equilibration. As a model learns patterns in data, its performance first improves and then plateaus, again reaching apparent stationarity. This dynamical process naturally comes equipped with stochastic fluctuations as well: often given data too gigantic to consume at once, training proceeds in small batches and random selections of these mini-batches consequently give rise to the noisy dynamical excursion of the model parameters in the loss-function landscape, reminiscent of the Brownian motion. It is thus natural to wonder if there exist analogous fluctuation-dissipation relations that quantitatively link the noise in mini-batched data to the observable evolution of the model performance and that in turn facilitate the learning process.

Here, we derive such fluctuation-dissipation relations for the stochastic gradient descent algorithm. The only assumption made is stationarity of the probability distribution that governs the model parameters at sufficiently long time. Our results thus apply to generic cases with non-Gaussian mini-batch noises and nonconvex loss-function landscapes. Practically, the first relation (FDR1) offers the metric for assessing equilibration and yields an adaptive algorithm that sets learning-rate schedule on the fly. The second relation (FDR2) further helps us determine the properties of the loss-function landscape, including the strength of its Hessian and the degree of anharmonicity, i.e., the deviation from the idealized harmonic limit of a quadratic loss surface and a constant noise matrix.

Our approach should be contrasted with recent attempts to import the machinery of stochastic differential calculus into the study of the stochastic gradient descent algorithm (Mandt et al., 2015; Li et al., 2015; Mandt et al., 2017; Li et al., 2017; Smith & Le, 2018; Chaudhari & Soatto, 2017; Jastrzebski et al., 2017; Zhu et al., 2018; An et al., 2018). This line of work all assumes Gaussian noises and sometimes additionally employs the quadratic harmonic approximation for loss-function landscapes. The more severe drawback, however, is the usage of the analogy with continuous-time stochastic differential equations, which is inconsistent in general (see Section 2.3.3). Instead, the stochastic gradient descent algorithm can be properly treated within the framework of the Kramers-Moyal expansion (Van Kampen, 1992; Gardiner, 2009; Risken, 1984; Radons et al., 1990; Leen & Moody, 1993).

The paper is organized as follows. In Section 2, after setting up notations and deriving a stationary fluctuation-dissipation theorem (FDT), we derive two specific fluctuation-dissipation relations. The first relation (FDR1) can be used to check stationarity and the second relation (FDR2) to delineate the shape of the loss-function landscape, as empirically borne out in Section 3. An adaptive scheduling method is proposed and tested in Section 3.3. We conclude in Section 4 with future outlooks.

Fluctuation-dissipation relations

where η>0\eta>0 is a learning rate and a mini-batch loss fB(θ)≡1∣B∣∑α∈Bfα(θ)f^{\mathcal{B}}\left(\bm{\theta}\right)\equiv\frac{1}{\left|{\mathcal{B}}\right|}\sum_{\alpha\in\mathcal{B}}f_{\alpha}\left(\bm{\theta}\right). Note that

and, more generally, higher-point noise tensors

In the next two subsections, we apply this general formula to simple observables in order to derive various stationary fluctuation-dissipation relations. Incidentally, the discrete version of the Fokker-Planck equation can be derived through the Kramers-Moyal expansion, considering the more general nonstationary version of the above equation and performing the Taylor expansion in η\eta and repeated integrations by parts (Van Kampen, 1992; Gardiner, 2009; Risken, 1984; Radons et al., 1990; Leen & Moody, 1993).

Applying the master equation (FDT) to the linear observable,

This is natural because there is no particular direction that the gradient picks on average as the model parameter stochastically bounces around the local minimum or, more generally, wanders around the loss-function landscape according to the stationary distribution.

Performing similar algebra for the quadratic observable ⟨θiθj⟩\left\langle\theta_{i}\theta_{j}\right\rangle yields

In particular, taking the trace of this matrix-form relation, we obtain

More generally, in the case of SGD with momentum μ\mu and dampening ν\nu, whose update equation is given by

a similar derivation yields (see Appendix A)

This first fluctuation-dissipation relation is easy to evaluate on the fly during training, exactly holds without any approximation if sampled well from the stationary distribution, and can thus be used as the standard metric to check if learning has plateaued, just as similar relations can be used to check equilibration in Monte Carlo simulations of physical systems (Santen & Krauth, 2000). [It should be cautioned, however, that the fluctuation-dissipation relations are necessary but not sufficient to ensure stationarity (Odriozola & Berthier, 2011).] Such a metric can in turn be used to schedule changes in hyperparameters, as shall be demonstrated in Section 3.3.

2 Second fluctuation-dissipation relation

Applying the master equation (FDT) on the full-batch loss function and Taylor-expanding it in the learning rate η\eta yields the closed-form expression

where we recalled the equation (4) and introduced

In particular, Hi,j(θ)≡Fi,j(θ)H_{i,j}\left(\bm{\theta}\right)\equiv\mathsfit{F}_{i,j}\left(\bm{\theta}\right) is the Hessian matrix. Reorganizing terms, we obtain

In the case of SGD with momentum and dampening, the left-hand side is replaced by (1−ν)⟨(∇f)2⟩−μ⟨v⋅∇f⟩(1-\nu)\left\langle\left(\bm{\nabla}f\right)^{2}\right\rangle-\mu\left\langle\mathbf{v}\cdot\bm{\nabla}f\right\rangle and C~i1,i2,…,ik\widetilde{\mathsfit{C}}_{i_{1},i_{2},\ldots,i_{k}} by more hideous expressions (see Appendix A).

3 Remarks

3.2 Higher-order relations

Additional relations can be derived by repeating similar calculations for higher-order observables. For example, at the cubic order,

The systematic investigation of higher-order relations is relegated to future work.

3.3 SGD≠\neqSDE

3.4 On stationarity

It is beyond the scope of the present paper to formulate conditions under which stationary distributions exist. Indeed, if the formulation were too generic, there could be counterexamples to such a putative existence statement. A case in point is a model with the unregularized cross entropy loss, whose model parameters keep cascading toward infinity in order to sharpen its softmax output (Neyshabur et al., 2014; 2017) with logarithmically diverging θ2\bm{\theta}^{2} (Soudry et al., 2018). It would be interesting to see if there are any other nontrivial caveats.

Empirical tests

Before proceeding further, let us define the half-running average of an observable O\mathcal{O} as

This is the average of the observable up to the time step tt, with the initial half discarded as containing transient. If SGD drives the distribution of the model parameters to stationarity at long time, then

In order to assess the proximity to stationarity, define

2 Second fluctuation-dissipation relation and shape of loss-function landscape

In order to assess the loss-function landscape information from the relation (FDR2), define

(with the second term nonexistent for SGD without momentum).For the second term, in order to ensure that lim⁡t→∞v⋅∇fB‾(t)=lim⁡t→∞v⋅∇f‾(t)\lim_{t\rightarrow\infty}\overline{\mathbf{v}\cdot\bm{\nabla}f^{\mathcal{B}}}(t)=\lim_{t\rightarrow\infty}\overline{\mathbf{v}\cdot\bm{\nabla}f}(t), we measure the half-running average of v(t)⋅∇fB[θ(t)]\mathbf{v}\left(t\right)\cdot\bm{\nabla}f^{\mathcal{B}}\left[\bm{\theta}\left(t\right)\right] and not v(t+1)⋅∇fB[θ(t)]\mathbf{v}\left(t+1\right)\cdot\bm{\nabla}f^{\mathcal{B}}\left[\bm{\theta}\left(t\right)\right]. Note that (∇f)2\left(\bm{\nabla}f\right)^{2} is a full-batch – not mini-batch – quantity. Given its computational cost, here we measure this first term only at the end of each epoch and take the half-running average over these sparse sample points, discarding the initial half of the run.

3 First fluctuation-dissipation relation and learning-rate schedules

Saturation of the relation (FDR1) suggests the learning stationarity, at which point it might be wise to decrease the learning rate η\eta. Such scheduling is often carried out in an ad hoc manner but we can now algorithmize this procedure as follows:

Here, two scheduling hyperparameters XX and YY are introduced, which control the threshold for saturation of the relation (FDR1) and the amount of decrease in the learning rate, respectively.

Plotted in the figure 4 are results for SGD without momentum, with the Xavier initialization (Glorot & Bengio, 2010) and training through (i) preset training schedule with decrease of the learning rate by a factor of 1010 for each 100100 epochs, (ii) an adaptive scheduler with X=0.01X=0.01 (1%1\% threshold) and Y=0.1Y=0.1 (10%10\% decrease), and (iii) the AMSGrad algorithm (J. Reddi et al., 2018) with the default hyperparameters. The adaptive scheduler attains comparable accuracies with the preset scheduling at long time and outperforms AMSGrad (see Appendix C for additional simulations).

These two scheduling methods span different subspaces of all the possible schedules. The adaptive scheduling method proposed herein has a theoretical grounding and in practice much less dimensionality for tuning of scheduling hyperparameters than the presetting method, thus ameliorating the optimization of scheduling hyperparameters. The systematic comparison between the two scheduling methods for state-of-the-arts architectures, and also the comparison with the AMSGrad algorithm for natural language processing tasks, could be a worthwhile avenue to pursue in the future.

Conclusion

In this paper, we have derived the fluctuation-dissipation relations with no assumptions other than stationarity of the probability distribution. These relations hold exactly even when the noise is non-Gaussian and the loss function is nonconvex. The relations have been empirically verified and used to probe the properties of the loss-function landscapes for the simple models. The relations further have resulted in the algorithm to adaptively set learning-rate schedule on the fly rather than presetting it in an ad hoc manner. In addition to systematically testing the performance of this adaptive scheduling algorithm, it would be interesting to investigate non-Gaussianity and noncovexity in more details through higher-point observables, both analytically and numerically. It would also be interesting to further elucidate the physics of machine learning by extending our formalism to incorporate nonstationary dynamics, linearly away from stationarity (Onsager, 1931; Green, 1954; Kubo, 1957) and beyond (Jarzynski, 1997; Crooks, 1999), so that it can in particular properly treat overfitting cascading dynamics and time-dependent sample distributions.

The author thanks Ludovic Berthier, Léon Bottou, Guy Gur-Ari, Kunihiko Kaneko, Ari Morcos, Dheevatsa Mudigere, Yann Ollivier, Yuandong Tian, and Mark Tygert for discussions. Special thanks go to Daniel Adam Roberts who prompted the practical application of the fluctuation-dissipation relations, leading to the adaptive method in Section 3.3.

References

Appendix A SGD with momentum and dampening

For SGD with momentum μ\mu and dampening ν\nu, the update equation is given by

Just as in the main text, from the assumed stationarity follows the master equation for SGD with momentum and dampening

Note that the relations (26) and (27) are trivially satisfied at each time step if the left-hand side observables are evaluated at one step ahead and thus their being satisfied for running averages has nothing to do with equilibration [the same can be said about the relation (23)]; the only nontrivial relation is the equation (28), which is a consequence of setting ⟨θiθj⟩\left\langle\theta_{i}\theta_{j}\right\rangle constant of time. After taking traces and some rearrangement, we obtain the relation (FDR1’) in the main text.

For the full-batch loss function, the algebra similar to the one in the main text yields

Appendix B Models and simulation protocols

Our multilayer perceptron (MLP) consists of a 784784-dimensional input layer followed by a hidden layer of 200200 neurons with ReLU activations, another hidden layer of 200200 neurons with ReLU activations, and a 1010-dimensional output layer with the softmax activation. The model performance is evaluated by the cross-entropy loss supplemented by the L2L^{2}-regularization term 12λθ2\frac{1}{2}\lambda\bm{\theta}^{2} with the weight decay λ=0.01\lambda=0.01.

Throughout the paper, the MLP is trained on the MNIST data through SGD without momentum. The data are shuffled at each epoch with the mini-batch size ∣B∣=100\left|{\mathcal{B}}\right|=100.

B.2 CNN on CIFAR-10 through SGD with momentum

In order to describe the architecture of our convolutional neural network (CNN) in detail, let us associate a tuple [F,C,S,P;M][F,C,S,P;M] to a convolutional layer with filter width FF, a number of channels CC, stride SS, and padding PP, followed by ReLU activations and a max-pooling layer of width MM. Then, as in the demo at Karpathy (2014), our CNN consists of a (32,32,3)(32,32,3) input layer followed by a convolutional layer with [5,16,1,2;2][5,16,1,2;2], another convolutional layer with [5,20,1,2;2][5,20,1,2;2], yet another convolutional layer with [5,20,1,2;2][5,20,1,2;2], and finally a fully-connected 1010-dimensional output layer with the softmax activation. The model performance is evaluated by the cross-entropy loss supplemented by the L2L^{2}-regularization term 12λθ2\frac{1}{2}\lambda\bm{\theta}^{2} with the weight decay λ=0.01\lambda=0.01.

Throughout the paper (except in Section 3.3 where the adaptive scheduling method is tested for SGD without momentum), the CNN is trained on the CIFAR-10 data through SGD with momentum μ=0.9\mu=0.9 and dampening ν=0\nu=0. The data are shuffled at each epoch with the mini-batch size ∣B∣=100\left|{\mathcal{B}}\right|=100.

Appendix C Additional simulations

Plotted in the figure S1 are the comparisons between Adam (Kingma & Ba, 2014) and AMSGrad (J. Reddi et al., 2018) algorithms with the default hyperparameters α=10−3\alpha=10^{-3}, (β1,β2)=(0.9,0.999)(\beta_{1},\beta_{2})=(0.9,0.999), and ϵ=10−8\epsilon=10^{-8}. The AMSGrad algorithm marginally outperforms the Adam algorithm for the tasks at hand and thus the results with the AMSGrad are presented in the main text.

C.2 Initial accuracy gain with different scheduling hyperparameters

In the figure 4(a) for the MNIST classification task with the MLP, the proposed adaptive method with the scheduling hyperparameters X=0.01X=0.01 and Y=0.1Y=0.1 outperforms the AMSGrad algorithm in terms of accuracy attained at long time and also exhibits a quick initial convergence. In the figure 4(b) for the CIFAR-10 classification task with the CNN, however, while the proposed adaptive method attains better accuracy at long time, its initial accuracy gain is visibly slower than the AMSGrad algorithm. This lag in initial accuracy gain can be ameliorated by choosing another combination of the scheduling hyperparameters, e.g., X=0.1X=0.1 and Y=0.3Y=0.3, at the expense of degradation in generalization accuracy with respect to the original choice X=0.01X=0.01 and Y=0.1Y=0.1. See the figure S2.