Stochastic modified equations and adaptive stochastic gradient algorithms
Qianxiao Li, Cheng Tai, Weinan E
Introduction
Stochastic gradient algorithms are often used to solve optimization problems of the form
Solving (1) using the standard gradient descent (GD) requires gradient evaluations per step and is prohibitively expensive when . An alternative, the stochastic gradient descent (SGD), is to replace the full gradient by a sampled version, serving as an unbiased estimator. In its simplest form, the SGD iteration is written as
where and are i.i.d uniform variates taking values in . The step-size is the learning rate. Unlike GD, SGD samples the full gradient and its computational complexity per iterate is independent of . For this reason, stochastic gradient algorithms have become increasingly popular in large scale problems.
Many convergence results are available for SGD and its variants. However, most are upper-bound type results for (strongly) convex objectives, often lacking the precision and generality to characterize the behavior of algorithms in practical settings. This makes it harder to translate theoretical understanding into algorithm analysis and design.
In this work, we address this by pursuing a different analytical direction. We derive continuous-time stochastic differential equations (SDE) that can be understood as weak approximations (i.e. approximations in distribution) of stochastic gradient algorithms. These SDEs contain higher order terms that vanish as , but at finite and small they offer much needed insight of the algorithms under consideration. In this sense, our framework can be viewed as a stochastic parallel of the method of modified equations in the analysis of classical finite difference methods Noh & Protter (1960); Daly (1963); Hirt (1968); Warming & Hyett (1974). For this reason, we refer to these SDEs as stochastic modified equations (SME). Using the SMEs, we can quantify, in a precise and general way, the leading-order dynamics of the SGD and its variants. Moreover, the continuous-time treatment allows the application of optimal control theory to study the problems of adaptive hyper-parameter adjustments. This gives rise to novel adaptive algorithms and perhaps more importantly, a general methodology for understanding and improving stochastic gradient algorithms.
We distinguish sequential and dimensional indices by writing a bracket around the latter, e.g. is the coordinate of the vector , the SGD iterate.
Stochastic Modified Equations
We now introduce the SME approximation. Background materials on SDEs are found in Supplementary Materials (SM) B and references therein. First, rewrite the SGD iteration rule (2) as
where is a -dimensional random vector. Conditioned on , has mean 0 and covariance matrix with
Now, consider the Stochastic differential equation
whose Euler discretization , resembles (3) if we set , and . Then, we would expect (5) to be an approximation of (2) with the identification . It is now important to discuss the precise meaning of “an approximation”. The noises that drive the paths of SGD and SDE are independent processes, hence we must understand approximations in the weak sense.
Let , and set . Let denote the set of functions of polynomial growth, i.e. if there exists constants such that . We say that the SDE (5) is an order weak approximation to the SGD (2) if for every , there exists , independent of , such that for all ,
The definition above is standard in numerical analysis of SDEs Milstein (1995); Kloeden & Platen (2011). Intuitively, weak approximations are close to the original process not in terms of individual sample paths, but their distributions. We now state informally the approximation theorem.
The stochastic process , satisfying
is an order 1 weak approximation of the SGD.
The stochastic process , satisfying
is an order 2 weak approximation of the SGD.
The full statement, proof and numerical verification of Thm. 1 is given in SM. C. We hereafter call equations (6) and (7) stochastic modified equations (SME) for the SGD iterations (2). We refer to the second order approximation (7) for exact calculations in Sec. 3 whereas for simplicity, we use the first order approximation (6) when discussing acceleration schemes in Sec. 4, where the order of accuracy is less important.
Thm. 1 allows us to use the SME to deduce distributional properties of the SGD. This result differs from usual convergence studies in that it describes dynamical behavior and is derived without convexity assumptions on or . In the next section, we use the SME to deduce some dynamical properties of the SGD.
The Dynamics of SGD
We start with a case where the SME is exactly solvable. Let , and set with and . Then, the SME (7) for the SGD iterations on this objective is (see SM. D.1)
with . This is the well-known Ornstein-Uhlenbeck process Uhlenbeck & Ornstein (1930), which is exactly solvable (see SM. B.3), yielding the Gaussian distribution
For , descent dominates and when , fluctuation dominates. This two-phase behavior is known for convex cases via error bounds Moulines (2011); Needell et al. (2014). Using the SME, we obtained a precise characterization of this behavior, including an exact expression for . In Fig. 1, we verify the SME predictions regarding the mean, variance and the two-phase behavior.
2 Stochastic Asymptotic Expansion
In general, we cannot expect to solve the SME exactly, especially for . However, observe that the noise terms in the SMEs (6) and (7) are . Hence, we can write as an asymptotic series where each is a stochastic process with initial condition and for . We substitute this into the SME and expand in orders of and equate the terms of the same order to get equations for for . This procedure is justified rigorously in Freidlin et al. (2012). We obtain to leading order (see SM. B.5),
where solves and , where , with denoting the Hessian of , and , . It is then possible to deduce the dynamics of the SGD. For example, there is generally a transition between descent and fluctuating regimes. has a steady state (assuming it is asymptotically stable) with . This means that one should expect a fluctuating regime where the covariance of the SGD is of order . Preceding this fluctuating regime is a descent regime governed by the gradient flow.
Adaptive Hyper-parameter Adjustment
We showed in the previous section that the SME formulation help us better understand the precise dynamics of the SGD. The natural question is how this can translate to designing practical algorithms. In this section, we exploit the continuous-time nature of our framework to derive adaptive learning rate and momentum parameter adjustment policies. These are particular illustrations of a general methodology to analyze and improve upon SGD variants. We will focus on the one dimensional case , and subsequently apply the results to high dimensional problems by local diagonal approximations.
1D SGD iterations with learning rate adjustment can be written as
where is the adjustment factor and is the maximum allowed learning rate. The corresponding SME for (9) is given by (SM. D.1)
where the time-dependent function is minimized over an admissible control set to be specified. To make headway analytically, we now turn to a simple quadratic objective.
1.2 Optimal Control of the Learning Rate
Hence, we may now recast the control problem as
This problem can solved by dynamic programming, using the Hamilton-Jacobi-Bellman equation Bellman (1956). We obtain the optimal control policy (SM. E.3)
With the policy (13), we can solve (12) and plug the solution for back into (13) to obtain the annealing schedule
where . Note that by putting , for small , this expression agrees with the transition time (3.1) between descent and fluctuating phases for the SGD dynamics considered in Sec. 3.1. Thus, this annealing schedule says that maximum learning rate should be used for descent phases, whereas decay on learning rate should be applied after onset of fluctuations. Our annealing result agree asymptotically with the commonly studied annealing schedules (Moulines, 2011; Shamir & Zhang, 2013), but the difference is that we suggest maximum learning rate before the onset of fluctuations. Of course, the key limitation is that our result is only valid for this particular objective. This naturally brings us to the next question: how does one apply the optimal control results to general objectives?
1.3 Application to General Objectives
Since we only assume that the diagonal-quadratic assumption holds locally, the terms , , and must be updated on the fly. There are potentially many methods for doing so. The approach we take exploits the linear relationship . Consequently, we may estimate via linear regression on the fly: for each dimension, we maintain exponential moving averages (EMA) where . For example, . The EMA decay parameter controls the effective averaging window size. We adaptively adjust it so that it is small when gradient variations are large, and vice versa. We employ the heuristic . This is similar to the approach in Schaul et al. (2013). We also clip each to to improve stability. Here, we use for all experiments, but we checked that performance is insensitive to these values. We can now compute by the ordinary-least-squares formula and as the variance of the gradients:
This allows us to estimate the policy (13) as
for . Since quantities are computed from exponentially averaged sources, we should also update our learning rate policy in the same way. The algorithm is summarized in Alg. 1. Due to its optimal control origin, we hereafter call this algorithm the controlled SGD (cSGD)
Alg. 1 can similarly be applied to mini-batch SGD. Let the batch-size be , which reduces the covariance by times and so in the SME is replaced by . However, at the same time estimating from mini-batch gradient sample variances will underestimate by a factor of . Thus the product remains unchanged and Alg. 1 can be applied with no changes.
The additional overheads in cSGD are from maintaining exponential averages and estimating on the fly with the relevant formulas. These are operations and hence scalable. Our current rough implementation runs slower per epoch than the plain SGD. This is expected to be improved by optimization, parallelization or updating quantities less frequently.
1.4 Performance on Benchmarks
Let us test cSGD on common deep learning benchmarks. We consider three different models. M0: a fully connected neural network with one hidden layer and ReLU activations, trained on the MNIST dataset (LeCun et al., 1998); C0: a fully connected neural network with two hidden layers and Tanh activations, trained on the CIFAR-10 dataset (Krizhevsky & Hinton, 2009); C1: a convolution network with four convolution layers and two fully connected layers also trained on CIFAR-10. Model details are found in SM. F.1. In Fig. 3, we compare the performance of cSGD with Adagrad (Duchi et al., 2011) and Adam (Kingma & Ba, 2015) optimizers. We illustrate in particular their sensitivity to different learning rate choices by performing a log-uniform random search over three orders of magnitude. We observe that cSGD is robust to different initial and maximum learning rates (provided the latter is big enough, e.g. we can take for all experiments) and changing network structures, while obtaining similar performance to well-tuned versions of the other methods (see also Tab. 1). In particular, notice that the best learning rates found for Adagrad and Adam generally differ for different neural networks. On the other hand, many values can be used for cSGD with little performance loss. For brevity we only show the test accuracies, but the training accuracies have similar behavior (see SM. F.5).
2 Momentum Parameter
Another practical way of speeding up the plain SGD is to employ momentum updates - an idea dating back to deterministic optimization Polyak (1964); Nesterov (1983); Qian (1999). However, the stochastic version has important differences, especially in regimes where sampling noise dominates. Nevertheless, provided that the momentum parameter is well-tuned, the momentum SGD (MSGD) is very effective in speeding up convergence, particularly in early stages of training (Sutskever et al., 2013).
Selecting an appropriate momentum parameter is important in practice. Typically, generic values (e.g. 0.9, 0.99) are suggested without fully elucidating their effect on the SGD dynamics. In this section, we use the SME framework to analyze the precise dynamics of MSGD and derive effective adaptive momentum parameter adjustment policies.
The SGD with momentum can be written as the following coupled updates
The parameter is the momentum parameter taking values in the range . Intuitively, the momentum term remembers past update directions and pushes along , which may otherwise slow down at e.g. narrow parts of the landscape. The corresponding SME is now a coupled SDE
This can be derived by comparing (16) with the Euler discretization scheme of (17) and matching moments. Details can be found in SM. D.3.
2.2 The Effect of Momentum
If , has a positive eigenvalue and hence diverges exponentially. Since is negative, its value must then decrease exponentially for all , and the descent rate is maximized at . The more interesting case is when . Instead of solving (18), we observe that all eigenvalues of have negative real parts as long as . Therefore, has an exponential decay dominated by , where denotes real part and is the eigenvalue with the least negative real part. Observe that the descent rate is maximized at
and when , becomes complex. Also, from (18) we have , provided the steady state is stable. The role of momentum in this problem is now clear. To leading order in we have for . Hence, any non-zero momentum will improve the initial convergence rate. In fact, the choice is optimal and above it, oscillations set in because of a complex . At the same time, increasing momentum also causes increment in eventual fluctuations, since . In Fig. 4(a), we demonstrate the accuracy of the SME prediction (18) by comparing MSGD iterations.
Armed with an intuitive understanding of the effect of momentum, we can now use optimal control to design policies to adapt the momentum parameter.
2.3 Optimal Control of the Momentum Parameter
For , we have discussed previously that maximizes the descent rate and fluctuations generally help decrease concave functions. Thus, the optimal control is always . The non-trivial case is when . Due to its bi-linearity, directly controlling (18) leads to bang-bang type solutionsBang-bang solutions are control solutions lying on the boundary of the control set and abruptly jumps among the boundary values. For example, in this case it jumps between and repeatedly. that are rarely feed-back laws Pardalos & Yatsenko (2010) and thus difficult to apply in practice. Instead, we notice that the descent rate is dominated by , and the leading order asymptotic fluctuations is , hence we may consider
with . Solving this control problem yields the (approximate) feed-back policy (SM. E.4)
with given in (19). This says that when far from optimum ( large), we set which maximizes average descent rate. When , fluctuations set in and we lower .
As in Sec. 4.1.3, we turn the control policy above into a generally applicable algorithm by performing local diagonal-quadratic approximations and estimating the relevant quantities on the fly. The resulting algorithm is mostly identical to Alg. 1 except we now use (21) to update and SGD updates are replaced with MSGD updates (see S.M. F.4 for the full algorithm). We refer to this algorithm as the controlled momentum SGD (cMSGD).
2.4 Performance on Benchmarks
We apply cMSGD to the same three set-ups in Sec. 4.1.4, and compare its performance to the plain Momentum SGD with fixed momentum parameters (MSGD) and the annealing schedule suggested in Sutskever et al. (2013), with (MSGD-A). In Fig. 5, we perform a log-uniform search over the hyper-parameters , and . We see that cMSGD achieves superior performance to MSGD and MSGD-A (see Tab. 1), especially when the latter has badly tuned . Moreover, it is insensitive to the choice of initial . Just like cSGD, this holds across changing network structures. Further, cMSGD also adapts to other hyper-parameter variations. In Fig. 6, we take tuned (and any ) and vary the learning rate . We observe that cMSGD adapts to the new learning rates whereas the performance of MSGD and MSGD-A deteriorates and must be re-tuned to obtain reasonable accuracy. In fact, it is often the case that MSGD and MSGD-A diverge when is large, whereas cMSGD remains stable.
Related Work
Classical bound-type convergence results for SGD and variants include Moulines (2011); Shamir & Zhang (2013); Bach & Moulines (2013); Needell et al. (2014); Xiao & Zhang (2014); Shalev-Shwartz & Zhang (2014). Our approach differs in that we obtain precise, albeit only distributional, descriptions of the SGD dynamics that hold in non-convex situations.
In the vein of continuous approximation to stochastic algorithms, a related body of work is stochastic approximation theory (Kushner & Yin, 2003; Ljung et al., 2012), which establish ODEs as almost sure limits of trajectories of stochastic algorithms. In contrast, we obtain SDEs that are weak limits that approximate not individual sample paths, but their distributions. Other deterministic continuous time approximation methods include Su et al. (2014); Krichene et al. (2015); Wibisono et al. (2016).
Related work in SDE approximations of the SGD are Mandt et al. (2015, 2016), where the authors derived the first order SME heuristically. In contrast, we establish a rigorous statement for this type of approximations (Thm. 1). Moreover, we use asymptotic analysis and control theory to translate understanding into practical algorithms. Outside of the machine learning literature, similar modified equation methods also appear in numerical analysis of SDEs (Zygalakis, 2011) and quantifying uncertainties in ODEs (Conrad et al., 2015).
The second half of our work deals with practical problems of adaptive selection of the learning rate and momentum parameter. There is abundant literature on learning rate adjustments, including annealing schedules Robbins & Monro (1951); Moulines (2011); Xu (2011); Shamir & Zhang (2013), adaptive per-element adjustments Duchi et al. (2011); Zeiler (2012); Tieleman & Hinton (2012); Kingma & Ba (2015) and meta-learning Andrychowicz et al. (2016). Our approach differs in that optimal control theory provides a natural, non-black-box framework for developing dynamic feed-back adjustments, allowing us to obtain adaptive algorithms that are truly robust to changing model settings. Our learning rate adjustment policy is similar to Schaul et al. (2013); Schaul & LeCun (2013) based on one-step optimization, although we arrive at it from control theory. Our method may also be easier to implement because it does not require estimating diagonal Hessians via back-propagation. There is less literature on momentum parameter selection. A heuristic annealing schedule (referred to as MSGD-A earlier) is suggested in Sutskever et al. (2013), based on the original work of Nesterov (1983). The choice of momentum parameter in deterministic problems is discussed in Qian (1999); Nesterov (2013). To the best of our knowledge, a systematic stochastic treatment of adaptive momentum parameter selection for MSGD has not be considered before.
Conclusion and Outlook
Our main contribution is twofold. First, we propose the SME as a unified framework for quantifying the dynamics of SGD and its variants, beyond the classical convex regime. Tools from stochastic calculus and asymptotic analysis provide precise dynamical description of these algorithms, which help us understand important phenomena, such as descent-fluctuation transitions and the nature of acceleration schemes. Second, we use control theory as a natural framework to derive adaptive adjustment policies for the learning rate and momentum parameter. This translates to robust algorithms that requires little tuning across multiple datasets and model choices.
An interesting direction of future work is extending the SME framework to develop adaptive adjustment schemes for other hyper-parameters in SGD variants, such as Polyak-Ruppert Averaging (Polyak & Juditsky, 1992), SVRG (Johnson & Zhang, 2013) and elastic averaging SGD (Zhang et al., 2015). More generally, the SME framework may be a promising methodology for the analysis and design of stochastic gradient algorithms and beyond.
Appendix A Modified equations in the numerical analysis of PDEs
The method of modified equations is widely applied in finite difference methods in numerical solution of PDEs Hirt (1968); Noh & Protter (1960); Daly (1963); Warming & Hyett (1974). In this section, we briefly demonstrate this classical method. Consider the one dimensional transport equation
We set time and space discretization steps to and and denote for and . The simplest scheme that can exhibit stability is the upwind scheme (Courant et al., 1952), where we approximate (22) by the difference equation
where + . The idea is to now approximate this difference scheme by another continuous PDE, that is not equal to the original equation (22) for non-zero . This can be done by Taylor expanding each term in (23) around . Simplifying and truncating to leading term in , we obtain the modified equation
where is the Courant-Friedrichs-Lewy (CFL) number (Courant et al., 1952). Notice that in the limit with fixed, one recovers the original transport equation, but for finite step sizes, the upwind scheme is really described by the modified equation (24). In other words, this truncated equation describes the leading order, non-trivial behavior of the finite difference scheme.
From the modified equation (24), one can immediately deduce a number of interesting properties. First, the error from the upwind scheme is diffusive in nature, due to the presence of the second order spatial derivative on the right hand side. Second, we observe that if the CFL number is greater than 1, then the coefficient for the diffusive term becomes negative and this results in instability. This is the well-known CFL condition. This places a fundamental limit on the spatial resolution for fixed temporal resolution with regards to the stability of the algorithm. Lastly, the error term is proportional to for fixed , thus it may be considered a first order method.
Now, another possible proposal for discretizing (22) is the Lax-Wendroff (LW) scheme (Lax & Wendroff, 1960):
Comparing with (24), we observe that the LW scheme error is of higher order (), but at the cost of introducing dispersive, instead of diffusive errors due to the presence of the third derivative. These findings are in excellent agreement with the actual behavior of their respective discrete numerical schemes (Warming & Hyett, 1974).
We stress here that if we simply took the trivial leading order, the right hand sides of (24) and (26) disappear and vital information, including stability, accuracy and the nature of the error term will be lost. The ability to capture the effective dynamical behavior of finite difference schemes is the key strength of the modified equations approach, which has become the primary tool in analyzing and improving finite difference algorithms. The goal of our work is to extend this approach to analyze stochastic algorithms.
Appendix B Summary of SDE terminologies and results
Here, we summarize various SDE terminologies and results we have used throughout the main paper and also subsequent derivations. A particular important result is the Itô formula (Sec. B.2), which is used throughout this work for deriving moment equations. For a thorough reference on the subject of stochastic calculus and SDEs, we suggest Oksendal (2013).
Let . An Itô stochastic differential equation on the interval is an equation of the form
The equation (27) is really a “shorthand” for the integral equation
The last integral is defined in the Itô sense, i.e.
where is a sequence of -partitions of and the limit represents convergence in probability. In (27), is known as the drift, and is known as the diffusion matrix. When they satisfy Lipschitz conditions, one can show that (27) (or (29)) has a unique strong solution (Oksendal (2013), Chapter 5). For our purposes in this paper, we consider the special case where do not depend on time and we set so that is a square matrix.
To perform calculus, we need an important result that generalizes the notion of chain rule to the stochastic setting.
B.2 Itô formula
where denotes gradient with respect to the first argument and denotes the Hessian, i.e. . The formula (31) is the Itô formula. If is not a scalar but a vector, then each of its component satisfy (31). Note that if , this reduces to the chain rule of ordinary calculus.
B.3 The Ornstein-Uhlenbeck process
To solve this equation, we change variables . Applying Itô formula, we have
which we can integrate from to to get
This is a path-wise solution to the SDE (32). To infer distributional properties, we do not require such precise solutions. In fact, we only need the distribution of the random variable at any fixed time . Observe that is really a Gaussian process, since the integrand in the Wiener integral is deterministic. Hence, we need only calculate its moments. Taking expectation on (34), we get
To obtain the covariance function, we see that
This can be evaluated by using Itô’s isometry, which says that for any adapted process , we have
and in particular, for fixed , we have
In Sec. 3.1 in the main paper, the solution of the SME is the OU process with . Making these substitutions, we obtain
B.4 Numerical solution of SDEs
By definition, , and are independent for each . Here, is the identity matrix. Hence, we have the Euler-Maruyama scheme
where .
One can show that the Euler-Maruyama method (42) is a first order weak approximation (c.f. Def. 1 in main paper) to the SDE (27). However, it is only a order scheme in the strong sense (Kloeden & Platen, 2011), i.e.
With more sophisticated methods, one can design higher order schemes (both in the strong and weak sense), see Milstein (1986).
B.5 Stochastic asymptotic expansion
Besides numerics, if there exists small parameters in the SDE, we can proceed with stochastic asymptotic expansions Freidlin et al. (2012). This is the case for the SME, which has a small multiplied to the noise term. Let us consider a time-homogeneous SDE of the form
where . The idea is to follow standard asymptotic analysis and write as an asymptotic series
We substitute (45) into (44) and assuming smoothness of and , we expand
and . In general, the equation for are linear stochastic differential equations with time-dependent coefficients depending on and the initial conditions are , for all . Hence, the asymptotic equations can be solved sequentially to obtain an estimate of to arbitrary order in . The equations for higher order terms become messy quickly, but they are always linear in the unknown, as long as all the previous equations are solved. For more details on stochastic asymptotic expansions, the reader is referred to Freidlin et al. (2012).
B.6 Asymptotics of the SME
We now derive the first two asymptotic equations of the SME. we take , ( term can be ignored for first two terms) and . Then, (47) becomes
where is the Hessian of .
In the following analysis, we shall assume that the truncated series approximation
where satisfy (48) and (49), describes the leading order stochastic dynamics of the SGD. Now, let us analyze the asymptotic equations in detail. First, we assume that the ODE (48) has a unique solution with . This is true if for example, is locally Lipschitz. Next, let us define the non-random functions
Both and are matrices for each . Then, (49) becomes the time-inhomogeneous linear SDE
with . Since the drift is linear and the diffusion matrix is constant (i.e. independent of ), is a Gaussian process. Hence we need only calculate its mean and covariance using Itô formula (see B.2). We have
and the covariance matrix satisfies the differential equation
with . This equation is a linearized version of the Riccati equation and there are simple closed-form solutions under special conditions, e.g. or is constant.
Hence, we conclude that the asymptotic approximation is a Gaussian process with distribution
where solves the ODE (48) and solves the ODE (54), with given by (51).
At this point, it is important to discuss the validity of the asymptotic approximation (55), and the SME approximation (56) in general. What we prove in Sec. C and is shown in Freidlin et al. (2012) is that for fixed , we can take small enough so that the SME and its asymptotic expansion is a good approximation of the distribution of the SGD iterates. What we did not prove is that for fixed , the approximations hold for arbitrary . In particular, it is not hard to construct systems where for fixed , both the SME and the asymptotic expansion fails when is large enough. To prove the second general statement requires further assumptions, particularly on the distribution of ’s. This is out of the scope of the current work.
Appendix C Formal Statement and proof of Thm. 1
and .
Fix some test function (c.f. Def. 1 in main paper). Suppose further that the following conditions are met:
satisfy a Lipschitz condition: there exists such that
and its partial derivatives up to order belong to .
satisfy a growth condition: there exists such that
and its partial derivatives up to order belong to .
Then, there exists a constant independent of such that for all , we have
That is, the equation (56) is an order weak approximation of the SGD iterations.
The basic idea of the proof is similar to the classical approach in proving weak convergence of discretization schemes of SDEs outlined in the seminal papers by Milstein (Milstein (1975, 1979, 1986, 1995)). The main difference is that we wish to establish that the continuous SME is an approximation of the discrete SGD, instead of the other way round, which is the case dealt by classical approximation theorems of SDEs with finite difference schemes. In the following, we first show that a one-step approximation has order error, and then deduce, using the general result in Milstein (1986), that the overall global error is of order .
In the subsequent proofs we will make repeated use of Taylor expansions in powers of . To simplify presentation, we introduce the shorthand that whenever we write , we mean that there exists a function (c.f. Def. 1 in main text) such that the error terms are bounded by . For example, we write
These results can be deduced easily using Taylor’s theorem with a variety of forms of the remainder, e.g. Lagrange form. We omit such routine calculations. We also denote the partial derivative with respect to by .
First, let us prove a lemma regarding moments of SDEs with small noise.
Let . Consider a stochastic process , satisfying the SDE
All functions above are evaluated at .
A classical result on semigroup expansions (Hille & Phillips (1996), Chapter XI) states that if and its derivatives up to order 6 belong to , then
To ensure its existence we may instead set to be purely imaginary, i.e. where is real. Then, (62) is known as the characteristic function (CF). The important property we make use of is that the moments of are found by differentiating the MGF (or CF) with respect to t. In fact, we have
where . We now expand in powers of using formula (61). We get,
All functions are again evaluated at . Finally, we apply formula (63) to deduce (i)-(iii). ∎
Next, we have an equivalent result for one SGD iteration.
Let . Consider , satisfying the SGD iterations
where All functions above are evaluated at .
From definition (65) and the definition of , the results are immediate. ∎
Now, we will need a key result linking one step approximations to global approximations due to Milstein. We reproduce the theorem, tailored to our problem, below. The more general statement can be found in Milstein (1986).
Let be a positive integer and let the assumptions in Theorem 1 hold. If in addition there exists so that
Then, there exists a constant so that for all we have
See Milstein (1986), Theorem 2 and Lemma 5. ∎
We are now ready to prove theorem 1 by checking the conditions in theorem 2 with . The second condition is implied by Lemma 2. The first condition is implied by Lemma 1 and Lemma 2 with the choice
To illustrate our approximation result, let us calculate, using Monte-Carlo simulations, the weak error of the SME approximation
for v.s. for different and generic polynomial test functions . The results are shown in Fig. 7. We see that we have order weak convergence, even when some conditions of the above theorem are not satisfied (Fig. 7(b)).
The Lipschitz condition (i) is to ensure that the SME has a unique strong solution with uniformly bounded moments Milstein (1986). If we allow weak solutions and establish uniform boundedness of moments by other means (more assumptions on the growth and direction of for large ), then condition (i) is expected to be relaxed although the technical details will be tedious.
The regularity conditions on and in Theorem 1 are inherited from Theorem 2 in Milstein (1986). For smooth objectives, polynomial growth conditions are usually not restrictive. Still, with care, these should be relaxed since in our case the small noise helps to reduce the number of terms containing higher derivatives in various Taylor and Itô-Taylor expansions. Proving a more general version of Theorem 1 will be left as future work.
Appendix D Derivation of SMEs
In this section, we include more detailed derivations of the SMEs used in the main paper. For brevity, we do not include rigorous proofs of approximation statements for SGD variants in Sec. D.2 and D.3, but only heuristic justifications. Proving rigorous statements for these approximations can be done by modifying the proof of Thm. 1.
We start with the example in Sec. 3.1 of the main paper. Let , and set with and . The SGD iterations picks at random between and and performs descent with their respective gradients. Recall that the (second order) SME is given by
Now, , and
D.2 SME for learning rate adjustment
The SGD iterations with learning rate adjustment is
where is the learning rate adjustment factor. is the maximum allowed learning rate. There are two reasons we introduce this hyper-parameter. First, gradients cannot be arbitrarily large since that will cause instabilities. Second, the SME is only an approximation of the SGD for small learning rates, and so it is hard to justify the approximation for large .
In this case, deriving the corresponding SME is extremely simple. Notice that we can define , . Then, the iterations above is simply
And hence the SME for SGD with learning rate adjustments is
D.3 SME for SGD with momentum
First let us consider the constant momentum parameter case. The SGD with momentum is the paired update
To derive and SME, notice that we can write the above as
Recall that since we are looking at first order weak approximations, it is sufficient to compare to the Euler-Maruyama discretization (Sec. B.4). We observe that the above can be seen as an Euler-Maruyama discretization of the coupled SDE
with the usual choice of . Hence, this is the first order SME for the SGD with momentum having a constant momentum parameter . For time-varying momentum parameter, we just replace by to get
Appendix E Solution of optimal control problems
We first introduce some basic terminologies and results on optimal control theory to pave way for our solutions to optimal control problems for the learning rate and momentum parameter. For simplicity, we restrict to one state dimension (), but similar equations hold for multiple dimensions. For a more thorough introduction to optimal control theory and calculus of variations, we refer the reader to Liberzon (2012).
Note that additional path constraints can also be added and (78) can also be made time-inhomogeneous, but for our purposes it is sufficient to consider the above form.
There are two principal ways of solving optimal control problems: either dynamic programming through the Hamilton-Jacobi-Bellman (HJB) equation (Bellman, 1956), or using the Pontryagin’s maximum principle (PMP) (Pontryagin, 1987). In this section, we will only discuss the HJB method as this is the one we employ to solve the relevant control problems in this paper.
E.2 Dynamic programming and the HJB equation
Notice that if there exists a solution to (80), then the value of the minimum cost is . The dynamic programming principle allows us to derive a recursion on the function , in the form of a partial differential equation (PDE)
This is known as the Hamilton-Jacobi-Bellman equation (HJB). Note that this PDE is solved backwards in time. The derivation of this PDE can be found in most references on optimal control, e.g. in Liberzon (2012). The main idea is the dynamic programming principle: for any the -portion of the optimal trajectory must again be optimal.
After solving the HJB (82), we can then obtain the optimal control as function of the state process and , given by
In some cases, we find that the optimal control is independent of time and is strictly of a feed-back control law, i.e.
In summary, to solve the optimal control problem (80), we first solve the HJB PDE (82), and then solve for the optimal control (83), and lastly (if necessary) solve the optimally controlled state process by substituting the solution of (83) into (78). Sometimes, the optimal control (83) can be solved without fully solving the HJB for , e.g. when and one can infer the sign of . This is the case for the two control problems we encounter in this paper. The solution to (83) is the most important for all practical purposes since it gives a way to adjust the control parameters on-the-fly, especially when we have a feed back control law.
E.3 Solution of the learning rate control problem
Now, let us apply the HJB equations (Sec. E.2) to solve the learning rate control problem. Recall from Sec. 4.1.2 that we wish to solve
This is of the form (80) with , and . Thus, we write the HJB equation
First, it’s not hard to see that for , for all . This is because, the lower the , the closer we are to the optimum and hence the minimum cost achievable in the same time interval should be less. Similarly, holds for if one reverses all previous statements (in this case is negative). Hence, we can calculate the minimum
Notice that this solution is a feed-back control policy. We can now substitute where
And therefore, we get from (88) the effective annealing schedule
E.4 Solution of the momentum parameter control problem
We shall consider the case , since for the optimal control is trivially . The momentum parameter control problem is
This is of the form (80) with , and . The HJB equation is
Again, it is easy to see that for all and so
This minimization problem has no closed form solution. However, observe that and is minimized at . Now, if , we have and so is increasing in for (one can check this by differentiation and showing that the derivative is always positive). Hence, and it is sufficient to consider in the minimization problem (95).
Next, observe that is decreasing in and negative if
At the same time, is negative and decreasing for . Thus, the product is positive and increasing for and hence we must have
Note that this is only a bound, but for small , we can take this as an approximation of , so long as it is less than . Hence, we arrive at
One can of course follow the steps in Sec. E.3 to calculate and hence in the form of an annealing schedule. We omit these calculations since they are not relevant to applications.
Appendix F Numerical experiments
In this section, we provide model and algorithmic details for the various numerical experiments considered in the main paper, as well as a brief description of the commonly applied adaptive learning rate methods that we compare the cSGD algorithm with.
In Sec. 4 from the main paper, we consider three separate models for two datasets.
where the activation function is the commonly used Rectified Linear Unit (ReLU)
Each regularization strength is set to be 1 divided by the dimension of the trainable parameter.
C0: fully connected NN on CIFAR-10
The CIFAR-10 dataset (Krizhevsky & Hinton, 2009) consists of 60000 small pixels of RGB natural images belonging to ten separate classes. We split the dataset into 50000 training samples and 10000 test samples. Our first model for this dataset is a deeper fully connected neural network
where we use a tanh activation function between the hidden layers
C1: convolutional NN on CIFAR-10
F.2 Adagrad and Adam
Here, we write down for completeness the iteration rules of Adagrad (Duchi et al., 2011), and Adam (Kingma & Ba, 2015) optimizers, which are commonly applied tools to tune the learning rate. For more details and background, the reader should consult the respective references.
Adagrad. The Adagrad modification to the SGD reads
where is the running sum of gradients for . The tunable hyper-parameters are the learning rate and the initial accumulator value . In this paper we consider only the learning rate hyper-parameter as this is equivalent to setting the initial accumulator to a common constant across all dimensions.
Adam. The Adam method has similar ideas to momentum. It keeps the exponential moving averages
The hyper-parameters are the learning rate and the EMA decay parameters .
Note that for both methods above, one can also introduce a regularization term to the denominator to prevent numerical instabilities.
F.3 Implementation of cSGD
Recall from Sec. 4.1 that the optimal control solution for learning rate control of the quadratic objective is given by
The idea is to perform a local quadratic approximation
This is equivalent to a local linear approximation of the gradient, i.e. for
This effectively decouples the control problems of identical one-dimensional control problems, so that we may apply (111) element-wise. We note that this approximation is only assumed to hold locally and the parameters must be updated. There are many ways to do this. Our approach uses linear regression on-the-fly via exponential moving averages (EMA). For each trainable dimension , we maintain the following exponential averages
The decay parameter controls the effective averaging window size. In practice, we should adjust so that it is small when variations are large, and vice versa. This ensures that our local approximations adapts to the changing landscapes. Since local variations is related to the gradient, we use the following heuristic
which is similar to the one employed in Schaul et al. (2013) for maintaining EMAs. The additional clipping to the range is to make sure that there are enough samples to calculate meaningful regressions, and at the same time prevent too large decay values where the contribution of new samples vanish. In the applications presented in this paper, we usually set and , but results are generally insensitive to these values.
With the EMAs (114), we compute by the ordinary-least-squares formula and as the variance of the gradients:
This allows us to estimate the policy (111) as
for . Since our averages are from exponentially averaged sources, we should also update our learning rate policy in the same way:
F.4 Implementation of cMSGD
We wish to apply the momentum parameter control
where . We proceed in the same way as in Sec. F.3 by keeping the relevant EMA averages and performing linear regression on the fly. The only difference is the application of the momentum parameter adjustment, which is
F.5 Training accuracy for C1
For completeness we also provide in Fig. 8 the training accuracies of C1 with various hyper-parameter choices and methods tested in this work. These complements the plots of test accuracies in Fig. 3,5,6 in the main paper. We see that cSGD and cMSGD display the same robustness in terms of test and training accuracies.