Second-Order Information in Non-Convex Stochastic Optimization: Power and Limitations
Yossi Arjevani, Yair Carmon, John C. Duchi, Dylan J. Foster, Ayush Sekhari, Karthik Sridharan
Introduction
This task plays a central role in the study of non-convex optimization: for functions satisfying a weak strict saddle condition (Ge et al. 2015), exact SOSPs (with ) are local minima, and therefore the condition (1) serves as a proxy for approximate local optimality. However, it is NP-Hard to decide whether a SOSP is a local minimum or a high-order saddle point (Murty and Kabadi 1987). Moreover, for a growing set of non-convex optimization problems arising in machine learning, SOSPs are in fact global minima (Ge et al. 2015; Ge et al. 2016; Sun et al. 2018; Ma et al. 2019). Consequently, there has been intense recent interest in the design of efficient algorithms for finding approximate SOSPs (Jin et al. 2017; Allen-Zhu 2018a; Carmon et al. 2018; Fang et al. 2018; Tripuraneni et al. 2018; Xu et al. 2018; Fang et al. 2019).
This restriction typically arises due to computational considerations (when is much cheaper to compute than , as in empirical risk minimization or Monte Carlo simulation), or due to fundamental online nature of the problem at hand (e.g., when represents a routing scheme and represents traffic on a given day). However, for many problems with additional structure, we have access to extra information. For example, we often have access to stochastic second-order information in the form of a Hessian estimator satisfying
In this paper, we characterize the extent to which the stochastic Hessian information (3), as well as higher-order information, contributes to the efficiency of finding first- and second-order stationary points. We approach this question from the perspective of oracle complexity (Nemirovski and Yudin 1983), which measures efficiency by the number of queries to estimators of the form (2)—and possibly (3)—required to satisfy the condition (1).
We provide new upper and lower bounds on the stochastic oracle complexity of finding -stationary points and (-SOSPs. In brief, our main results are as follows.
Finding -stationary points: The elbow effect. We propose a new algorithm that finds an -stationary point () with stochastic gradients and stochastic Hessian-vector products. We furthermore show that this guarantee is not improvable via a complementary lower bound. All previous algorithms achieving complexity require “multi-point” queries, in which the algorithm can query stochastic gradients at multiple points for the same random seed. Moreover, we show that remains a lower bound for stochastic th-order methods for all and hence—in contrast to the deterministic setting—the optimal rates for higher-order methods exhibit an “elbow effect”; see Figure 1.
-stationary points: Improved algorithm and nearly matching lower bound. We extend our algorithm to find -stationary points using stochastic gradient and Hessian-vector products, and prove a nearly matching lower bound.
We first describe our developments for the task of finding -approximate first-order stationary points (satisfying (1) with ), and subsequently extend our results to general . The reader may also refer to Table 1 for a succinct comparison of upper bounds.
Our approach builds on a line of work by Fang et al. 2018; Zhou et al. 2018; Wang et al. 2019; Cutkosky and Orabona 2019 that also develop algorithms with complexity , but require a “multi-point” oracle in which algorithm can query the stochastic gradient at multiple points for the same random seed. Specifically, in the -point variant of this model, the algorithm can query at the set of points and receive
and where the estimator is unbiased and has bounded variance in the sense of (2). The aforementioned works achieve complexity using simultaneous queries, while our new algorithm achieves the same rate using (i.e., is drawn afresh at each query), but using stochastic Hessian-vector products in addition to stochastic gradients. However, we show in Appendix B that under the statistical assumptions made in these works, the two-point stochastic gradient oracle model is strictly stronger than the single-point stochastic gradient/Hessian-vector product oracle we consider here. On the other hand, unlike our algorithm, these works do not require Lipschitz Hessian.
The algorithms that achieve complexity using two-point queries work by estimating gradient differences of the form using and applying recursive variance reduction (Nguyen et al. 2017). Our primary algorithmic contribution is a second-order stochastic estimator for which avoids simultaneous queries while maintaining comparable error guarantees. To derive our estimator, we note that , and use queries to the stochastic Hessian estimator (3) to numerically approximate this integral. More precisely, our estimator (5) only requires stochastic Hessian-vector products, whose computation is often roughly as expensive as that of a stochastic gradient (Pearlmutter 1994). Specifically, our estimator takes the form
For functions with Lipschitz gradient and Hessian, we prove an lower bound on the minimax oracle complexity of algorithms for finding stationary points using only stochastic gradients (2). We formally prove our results for the structured class of zero-respecting algorithms (Carmon et al. 2019a); the lower bounds extend to general randomized algorithms via similar arguments to Arjevani et al. 2019a. This lower bound is an extension of the results of Arjevani et al. 2019a, who showed that for functions with Lipschitz gradient but not Lipschitz Hessian, the optimal rate is using only stochastic gradients (2). Together with our new upper bound, this lower bound reveals that stochastic Hessian-vector products offer an improvement in the oracle complexity for finding stationary points in the single-point query model. This contrasts the noiseless optimization setting, where finite gradient differences can approximate Hessian-vector products arbitrarily well, meaning these oracle models are equivalent.
For algorithms that can query both stochastic gradients and stochastic Hessians, we prove a lower bound of on the oracle complexity of finding an expected -stationary point. This proves that our upper bound is optimal in the leading order term in , despite using only stochastic Hessian-vector products rather than full stochastic Hessian queries.
Notably, our lower bound extends to settings where stochastic higher-order oracles are available, i.e, when the first derivatives are Lipschitz and we have bounded-variance estimators . The lower bound holds for any finite , and thus, as a function of the oracle order , the minimax complexity has an elbow (Figure 1): for the complexity is (Arjevani et al. 2019a) while for all it is . This means that smoothness and stochastic derivatives beyond the second-order cannot improve the leading term in rates of convergence to stationarity, establishing a fundamental limitation of stochastic high-order information. This highlights another contrast with the noiseless setting, where th order methods enjoy improved complexity for every (Carmon et al. 2019a).
As we discuss in Appendix B, for multi-point stochastic oracles (4), the rate is attainable even without stochastic Hessian access. Moreover, our lower bound for stochastic th order oracles holds even when multi-point queries are allowed. Consequently, when viewed through the lens of worst-case oracle complexity, our lower bounds show that even stochastic Hessian information is not helpful in the multi-point setting.
1.2 Second-order stationary points
We incorporate our recursive variance-reduced Hessian-vector product-based gradient estimator into an algorithm that combines SGD with negative curvature search. Under the slightly stronger (relative to (3)) assumption that the stochastic Hessians have almost surely bounded error, we prove that—with constant probability—the algorithm returns an -SOSP after performing stochastic gradient and Hessian-vector product queries.
We prove a minimax lower bound which establishes that the stochastic second-order oracle complexity of finding -SOSPs is . Consequently, the algorithms we develop have optimal worst-case complexity in the regimes and . Compared to our lower bounds for finding -stationary points, proving the lower bound requires a more substantial modification of the constructions of Carmon et al. 2019a and Arjevani et al. 2019a. In fact, our lower bound is new even in the noiseless regime (i.e., ), where it becomes ; this matches the guarantee of the cubic-regularized Newton’s method (Nesterov and Polyak 2006) and consequently characterizes the optimal rate for finding approximate SOSPs using noiseless second-order methods.
2 Further related work
We briefly survey additional upper and lower complexity bounds related to our work and place our results within their context. The works of Monteiro and Svaiter 2013; Arjevani et al. 2019b; Agarwal and Hazan 2018 delineate the second-order oracle complexity of convex optimization in the noiseless setting; Arjevani and Shamir 2017 treat the finite-sum setting.
For functions with Lipschitz gradient and Hessian, oracle access to the Hessian significantly accelerates convergence to -approximate global minima, reducing the complexity from to . However, since the hard instances for first-order convex optimization are quadratic (Nemirovski and Yudin 1983; Arjevani and Shamir 2016; Simchowitz 2018), assuming Lipschitz continuity of the Hessian does not improve the complexity if one only has access to a first-order oracle. This contrasts the case for finding -approximate stationary points of non-convex functions with noiseless oracles. There, Lipschitz continuity of the Hessian improves the first-order oracle complexity from to , with a lower bound of for deterministic algorithms (Carmon et al. 2017; Carmon et al. 2019b). Additional access to full Hessian further improves this complexity to , and for th-order oracles with Lipschitz th derivative, the complexity further improves to (Carmon et al. 2019a); see Figure 1.
3 Paper organization
We formally introduce our notation and oracle model in Section 2. Section 3 contains our results concerning the complexity of finding -first-order stationary points: algorithmic upper bounds (Section 3.1) and algorithm-independent lower bounds (Section 3.2). Following a similar outline, Section 4 describes our upper and lower bounds for finding -SOSPs. We conclude the paper in Section 5 with a discussion of directions for further research. Additional technical comparison with related work is given in Appendix A and B, and proofs are given in Appendix C through Appendix G.
Setup
We study the problem of finding -stationary and -second order stationary points in the standard oracle complexity framework (Nemirovski and Yudin 1983), which we briefly review here.
We consider -times differentiable functions satisfying standard regularity conditions, and define
so that specifies the Lipschitz constants of the th order derivatives with respect to the operator norm. We make no restriction on the ambient dimension .
For a given function , we consider a class of stochastic th order oracles defined by a distribution over a measurable set and an estimator
where are unbiased estimators of the respective derivatives. That is, for all , and for all . For we assume without loss of generality that is a symmetric tensor.
Given variance parameters , we define the oracle class to be the set of all stochastic th-order oracles for which the variance of the derivative estimators satisfies
We consider stochastic th-order optimization algorithms that access an unknown function through multiple rounds of queries to a stochastic th-order oracle . When queried at in round , the oracle performs an independent draw of and answers with . Algorithm queries depend on only through the oracle answers; see e.g. Arjevani et al. 2019a for a more formal treatment.
Complexity of finding first-order stationary points
In this section we focus on the task of finding -approximate stationary points (satisfying ). As prior work observes (Carmon et al. 2017; Allen-Zhu 2018a, cf.), stationary point search is a useful primitive for achieving the end goal of finding second-order stationary points (1). We begin with describing algorithmic upper bounds on the complexity of finding stationary points with stochastic second-order oracles, and then proceed to match their leading terms with general th order lower bounds.
Our algorithms rely on recursive variance reduction (Nguyen et al. 2017): we sequentially estimate the gradient at the points by accumulating cheap estimators of for , where at iteration we reset the gradient estimator by computing a high-accuracy approximation of with many oracle queries. Our implementation of recursive variance reduction, Algorithm 1, differs from previous approaches (Fang et al. 2018; Zhou et al. 2018; Wang et al. 2019) in three aspects.
In Line 9 we estimate differences of the form by averaging stochastic Hessian-vector products. This allows us to do away with multi-point queries and operate under weaker assumptions than prior work (see Appendix B), but it also introduces bias to our estimator, which makes its analysis more involved. This is the key novelty in our algorithm.
Rather than resetting the gradient estimator every fixed number of steps, we reset with a user-defined probability (Line 5); this makes the estimator stateless and greatly simplifies its analysis, especially when we use a varying value of to find second-order stationary points.
We dynamically select the batch size for estimating gradient differences based on the distance between iterates (Line 3), while prior work uses a constant batch size. Our dynamic batch size scheme is crucial for controlling the bias in our gradient estimator, while still allowing for large step sizes as in Wang et al. 2019.
The core of our analysis is the following lemma, which bounds the gradient estimation error and expected oracle complexity. To state the lemma, we let be sequence of queries to Algorithm 1, and let be the sequence of estimates it returns.
For any oracle in and , Algorithm 1 guarantees that
for all . Furthermore, conditional on , and , the execution of Algorithm 1 with reset probability uses at most
stochastic gradient and Hessian-vector product queries in expectation.
We prove the lemma in Appendix C by bounding the per-step variance using the HVP oracle’s variance bound (7), and by bounding the per-step bias relative to using the Lipschitz continuity of the Hessian.
Taking , we are guaranteed that a uniformly selected iterate has expected norm .
To account for oracle complexity, we observe from Lemma 1 that calls to Algorithm 1 require at most oracle queries in expectation. Using , Lemma 1 and (8) imply that . We then choose to out the terms and . This gives the following complexity guarantee, which we prove in Appendix E.1.
For any function , stochastic second-order oracle in , and , with probability at least , Algorithm 2 returns a point such that and performs at most
stochastic gradient and Hessian-vector product queries.
The oracle complexity of Algorithm 2 depends on the Lipschitz parameters of only through lower-order terms in , with the leading term scaling only with the variance of the gradient and Hessian estimators. In the low noise regime where and , the complexity becomes which is simply the maximum of the noiseless guarantees for gradient descent and Newton’s method. We remark, however, that in the noiseless regime , a slightly better guarantee is achievable (Carmon et al. 2017).
In the noiseless setting, any algorithm that uses only first-order and Hessian-vector product queries must have complexity scaling with , but full Hessian access can remove this dependence (Carmon et al. 2019b). We show that the same holds true in the stochastic setting: Algorithm 3, a subsampled cubic regularized trust-region method using Algorithm 1 for gradient estimation, enjoys a complexity bound independent of . We defer the analysis to Appendix E.2 and state the guarantee as follows.
For any function , stochastic second order oracle in , and , with probability at least , Algorithm 3 returns a point such that and performs at most
The guarantee of Theorem 2 constitutes an improvement in query complexity over Theorem 1 in the regime . However, depending on the problem, full stochastic Hessians can be up to times more expensive to compute than stochastic Hessian-vector products.
2 Lower bounds
Having presented stochastic second-order methods with -complexity bound for finding -stationary points, our we next show that this rates cannot be improved. In fact, we show that this rate is optimal even when one is given access to stochastic higher derivatives of any order. We prove our lower bounds for the class of zero-respecting algorithms, which subsumes the majority of existing optimization methods; see Appendix G.1 for a formal definition. We believe that existing techniques (Carmon et al. 2019a; Arjevani et al. 2019a) can strengthen our lower bounds to apply to general randomized algorithms; for brevity, we do not pursue it here.
The lower bounds in this section closely follow a recent construction by Arjevani et al. 2019a, who prove lower bounds for stochastic first-order methods. To establish complexity bounds for th-order methods, we extend the ‘probabilistic zero-chain’ gradient estimator introduced in Arjevani et al. 2019a to high-order derivative estimators.The most technically demanding part of our proof is a careful scaling of the basic construction to simultaneously meet multiple Lipschitz continuity and variance constraints. Deferring the proof details to Appendix G.1, our lower bound is as follows.
A construction of dimension realizes this lower bound.
For second-order methods (with ), Theorem 3 specializes to the oracle complexity lower bound
which is tight in that it matches (up to numerical constants) the convergence rate of Algorithm 2 in the regime where dominates both the upper bound in Theorem 1 and expression (10). The lower bound (10) is also tight when the second-order information is not available or reliable ( is infinite or very large, respectively): Standard SGD matches the term (Ghadimi and Lan 2013), while more sophisticated variants based on restarting (Fang et al. 2019) and normalized updates with momentum (Cutkosky and Mehta 2020) match the term (the former up to logarithmic factors)—neither of these algorithms requires stochastic second derivative estimation.
Theorem 3 implies that while higher-order methods (with ) might achieve better dependence on the variance parameters than the upper bounds for Algorithm 2 or Algorithm 3, they cannot improve the scaling. This highlights a fundamental limitation for higher-order methods in stochastic non-convex optimization which does not exist in the noiseless case. Indeed, without noise the optimal rate for finding -stationary point with a th order method is Carmon et al. 2019a; we illustrate this contrast in Figure 1.
Altogether, the results presented in this section fully characterize (with respect to dependence on ) the complexity of finding -stationary points with stochastic second-order methods and beyond in the single-point query model. We briefly remark that lower bound in (9) immediately extends to multi-point queries, which shows that even second-order methods offer little benefit once two or more simultaneous queries are allowed.
Complexity of finding second-order stationary points
Having established rates of convergence for finding -stationary points, we now turn our attention to -second order stationary points, which have the additional requirement that , i.e. that is -weakly convex around . This section follows the general organization of the prequel: we first design and analyze an algorithm with improved upper bounds, and then develop nearly-matching lower bounds that apply to a broad class of algorithms.
Our first contribution for this section is an algorithm that enjoys improved complexity for finding -second-order stationary points, and that achieves this using only stochastic gradient and Hessian-vector product queries. To guarantee second-order stationarity, we follow the established technique of interleaving an algorithm for finding a first-order stationary point with negative curvature descent (Carmon et al. 2017; Allen-Zhu 2018a). However, we employ a randomized variant of this approach. Specifically, at every iteration we flip a biased coin to determine whether to perform a stochastic gradient step or a stochastic negative curvature descent step.
Our algorithm estimates stochastic gradients using the HVP-RVR scheme (Algorithm 1), where the value of the restart probability depends on the type of the previous step (gradient or negative curvature). To implement negative curvature descent, we apply Oja’s method (Oja 1982; Allen-Zhu and Li 2017) which detects directions of negative curvature using only stochastic Hessian-vector product queries. For technical reasons pertaining to the analysis of Oja’s method, we require the stochastic Hessians to be bounded almost surely, i.e., a.s.; we let denote the class of such bounded noise oracles. Under this assumption, Algorithm 4—whose description is deferred to the Appendix F---enjoys the following convergence guarantee. The notation hides lower-order terms and logarithmic dependence on the dimension . See the proof in Appendix F for the complete description of the algorithm and the full complexity bound, including lower order terms.
For any function , stochastic Hessian-vector product oracle in , , and , with probability at least Algorithm 4 returns a point such that
stochastic gradient and Hessian-vector product queries.
Similar to the case for finding -stationary points (see discussion preceding Theorem 2), using full stochastic Hessian information allows us to design an algorithm (Algorithm 5) which removes the dependence on from the theorem above. Moreover, estimating negative curvature directly from empirical Hessian estimates saves us the need to use Oja’s method, which means that we do not need the additional boundedness assumption on the stochastic Hessian used by Algorithm 4. We defer the complete description and analysis for Algorithm 5 to Appendix F.2, and state its complexity guarantee below.
For any function , stochastic second order oracle in , , and , with probability at least Algorithm 5 returns a point such that
2 Lower bounds
We now develop lower complexity bounds for the task of finding -stationary points. To do so, we prove new lower bounds for the simpler sub-problem of finding a -weakly convex point, i.e., a point such that (with no restriction on ). Lower bounds for finding -SOSPs follow as the maximum (or, equivalently, the sum) of lower bounds we develop here and the lower bounds for finding -stationary points given in Theorem 6. To see why this is so, let and be hard instances for finding -stationary and -weakly-convex points respectively, and consider the “direct sum” ; this is a hard instance for finding -SOSPs that inherits all the regularity properties of its constituent functions.
The basic construction we use here is a modification of the zero-chain introduced in Carmon et al. 2019a (see (75) in Appendix G) in which large is possible only when essentially none of the entries of is zero. Given , we define the hard function
where (as in Carmon et al. 2019a) and .
Our design for the function guarantees that any query whose last coordinate is zero has significant negative curvature, while maintaining the original chain structure which guarantees that zero-respecting algorithms require many queries before “discovering” the last coordinate. We complete the construction by specifying a collection of stochastic derivative estimators similar to those in Section 4.2, except for that we choose the stochastic gradient estimator to be exactly equal to , so that the lower bound holds even for ; Appropriately scaling allows us to tune the Lipschitz constants of its derivatives and the variance of the estimators, thereby establishing the following complexity bounds (see Appendix G.2 for a full derivation).
Let and be fixed. If , then there exists and such that for any stochastic th-order zero-respecting algorithm, the number of queries to required to obtain a -weakly convex point with constant probability is at least
A construction of dimension realizes the lower bound.
Theorem 6 is new even in the noiseless case (in which ), where it specializes to
For the class , the lower bound (13) further simplifies to , which is attained by the th-order regularization method given in Cartis et al. 2017. Together, these results characterize the deterministic complexity of finding -weakly convex points with noiseless th-order methods.
Returning to the stochastic setting, the bound in Theorem 6, when combined with Theorem 3, implies the following oracle complexity lower bound bound for finding -SOSPs with zero-respecting stochastic second-order methods ():
Our lower bound matches the terms in the upper bound given by Theorem 4, but does not match the mixed term appearing in the upper bound. Young’s inequality only gives . Overall, the rates match whenever or .
Theorem 6 is suggestive of another “elbow” phenomenon: In the stochastic regime, the rate does not improve beyond for , while the optimal rate in the noiseless regime, , continues improving for all . Indeed, when high-order noise moments are assumed finite, the term can longer be disregarded. This, in turn, implies that for sufficiently small , one cannot improve over -scaling, as seen by (12). However, we are not yet aware of an algorithm using stochastic third-order information or higher that can achieve the complexity bound.
Conclusion
This paper provides a fairly complete picture of the worst-case oracle complexity of finding stationary points with a stochastic second-order oracle: for -stationary points we characterize the leading term in exactly and for ()-SOSPs we characterize the leading term in for a wide range of parameters. Nevertheless, our results point to a number of open questions.
Our upper and lower bounds (in Theorem 5 and Theorem 6) resolve the optimal rate to find an -stationary point for , i.e., when is second-order smooth and the algorithm can query stochastic gradient and Hessian information. Furthermore, Theorem 3 shows that higher order information () cannot improve the dependence of the rate on the first-order stationarity parameter . However, our lower bound for dependence on scales as for , but scales as for . The weaker lower bound for leaves open the possibility of a stronger upper bound using third-order information or higher.
For statistical learning and sample average approximation problems, it is natural to consider problem instances of the form . For this setting, a more powerful oracle model is the global oracle, in which samples are drawn i.i.d. and the learner observes the entire function for each . Global oracles are more powerful than stochastic th order oracles for every , and lead to improved rates in the convex setting (Foster et al. 2019). Is it possible to beat the elbow for such oracles, or do our lower bounds extend to this setting?
Our lower bounds show that stochastic higher-order methods cannot improve the oracle complexity attained with stochastic gradients and Hessian-vector products. Furthermore, in the multi-point query model, stochastic second-order information does not even lead to improved rates over stochastic first-order information. However, these conclusions could be artifacts of our worst-case point of view—are there natural families of problem instances for which higher-order methods can adapt to additional problem structure and obtain stronger instance-dependent convergence guarantees? Developing a theory of instance-dependent complexity that can distinguish adaptive algorithms stands out as an exciting research prospect.
Acknowledgements
We thank Blake Woodworth and Nati Srebo for helpful discussions. YA acknowledges partial support from the Sloan Foundation and Samsung Research. JCD acknowledges support from the NSF CAREER award CCF-1553086, ONR YIP N00014-19-2288, Sloan Foundation, NSF HDR 1934578 (Stanford Data Science Collaboratory), and the DAWN Consortium. DF acknowledges the support of TRIPODS award 1740751. KS acknowledges support from NSF CAREER Award 1750575 and a Sloan Research Fellowship.
References
Appendix A Detailed comparison with existing rates
Appendix B Comparison: multi-point queries and mean-squared smoothness
Stochastic first-order methods that utilize variance reduction (Lei et al. 2017; Fang et al. 2018; Zhou et al. 2018) employ the following mean-squared smoothness (MSS) assumption on the stochastic gradient estimator:
Algorithms that take advantage of the MSS structure rely on the following additional simultaneous query assumption (which is a special case of (4) for ):
In empirical risk minimization problems, represents the datapoint index and possibly data augmentation parameters, and the value of is typically part of the query, which means that assumption (16) indeed holds. In certain online learning settings, however, the assumption can fail. For example, the variable could represent the instantaneous power demands in an electric grid, and testing two grid configurations for the same grid state might be impractical.
We observe that assuming access to both an MSS gradient estimator and simultaneous two-point queries is stronger than assuming a bounded variance stochastic Hessian-vector product estimator. This holds because the former allows us to simulate the latter with finite differencing. Formally, we have the following.
Let have -Lipschitz Hessian, let satisfy (15), and assume we have access to a two-point query oracle as in (16). Then, for any and every unit-norm vector , the Hessian-vector product estimator
which implies the bound on the bias. To bound the variance, we note that
We conclude from Observation 1 that Algorithm 2, which only requires stochastic Hessian-vector products, attains complexity under assumptions no stronger than previous algorithms. In fact, we show now that our assumptions are in fact strictly weaker than prior work. That is, while an MSS gradient estimator implies a bounded variance Hessian estimator, the opposite is not true in general. This is simply due to the fact that in our oracle model, and can be completely unrelated. Consider for example the case where is uniform on and
Clearly is not MSS, even though has zero variance.
There is, however, an important setting where bounded variance for does imply that is MSS. Suppose that the derivative of exists, and has the form
That is, the Hessian estimator is the Jacobian of the gradient estimator. In this case, bounded variance for the Hessian estimator implies mean-squared smoothness.
Taking the squared norm, applying Jensen’s inequality, and substituting the variance bound (3) gives the MSS property (15). ∎
The property (18) holds for empirical risk minimization, where we have the more general relation for any ; That is, all the stochastic derivative estimators are themselves the derivatives of a single stochastic function. Therefore, by Observation 1 and Observation 2, in empirical risk minimization settings, mean-square smoothness is essentially equivalent to bounded variance of the stochastic Hessian estimator.
Appendix C Variance-reduced gradient estimator (HVP-RVR)
In this section we prove Lemma 1. First, we formally describe the protocol in which our optimization algorithms query the gradient estimator HVP-RVR-Gradient-Estimator described in Algorithm 1, and define some additional notation.
Given a function and a stochastic second-order oracle in , the optimization algorithm interacts with HVP-RVR-Gradient-Estimator by sequentially querying points with reset probabilities , to obtain estimates for for each time ; that is,
where are measurable mappings modeling the optimization algorithm and is an independent sequence of random seeds. This level of formalism is not used within the proof, but we include it here for clarity. That is, Lemma 1 holds for any sequence of queries where , and are adapted to the filtration
but is independent of and .
Lemma 1 is an immediate consequence of Lemma 2 and Lemma 3, proven below, which respectively establish the estimator’s error and complexity bounds.
Given a function , a stochastic oracle in , and initial points and , let denote the sequence of gradient estimates at respectively, returned by HVP-RVR-Gradient-Estimator under the protocol (19). Then, for all ,
whence the result follows by a simple induction whose basis is
Moreover, conditional on , we have from the definition of the gradient estimator that
where and respectively denote the values of and (defined on Line 9) during the call to Algorithm 1.
We may therefore decompose the error conditional on as
where is due to and is due to Young’s inequality.
The facts that is independent from , that , and that is unbiased give
for every . Consequently, the scaling (22) and Hessian estimator variance bound imply
where the equality above is due to the fact that are i.i.d., as well as .
Substituting back through equations (25), (24), (23), (21) and (20), we have
as required; the second inequality follows from algebraic manipulation and the fact that is independent of by assumption. ∎
The following lemma bounds the number of oracle queries made per call to the gradient estimator.
where the final inequality follows from . ∎
Appendix D Supporting technical results
In order to find the negative curvature direction at a given point or to build a cubic regularized sub-model, Algorithm 3 and Algorithm 5 estimate the Hessian by computing an empirical average of the stochastic Hessian queries to the oracle. The following lemma is a standard result which bounds the expected error for the empirical Hessian.
This is an immediate consequence of Lemma 5 below, using and . ∎
We drop the normalization by throughout this proof. We first symmetrize. Observe that by Jensen’s inequality we have
where the second inequality follows by Jensen. We now apply the matrix Khintchine inequality (Mackey et al. 2014, Corollary 7.4), which implies that
Putting all the developments so far together and taking expectation with respect to , we have
To obtain the final result we normalize by . ∎
D.2 Descent lemma for stochastic gradient descent
The following lemma characterizes the effect of gradient descent update step used by Algorithm 2 and Algorithm 4.
Given a function , a point , and gradient estimator at x, define
Then, for any , the point satisfies
Since, the gradient of is -Lipschitz, we have
where uses that , is due to the Cauchy-Schwarz inequality, is given by an application of the AM-GM inequality and holds because . Finally, follows by invoking Jensen’s inequality for the function to upper bound . Rearranging the terms in (26), we get,
D.3 Descent lemma for cubic-regularized trust-region method
The following lemmas establish properties for the updates step involving constrained minimization of the cubic regularized model in used in Algorithm 3 and Algorithm 5.
Since is -Lipschitz, we have
where follows by the definition of the operator norm and follows by observing that . Rearranging the terms, we have
Under the same setting as Lemma 7, the point satisfies
Since one of the two cases ( or ) must hold, we have,
Rearranging the terms, and using the fact that , we have
Finally, using the fact that for any , , we have
where and are taken with respect to the randomness over and .
For the ease of notation, let and denote the error in the gradient estimator and the hessian estimator at respectively, i.e.
We prove the desired statement by combining the following two results.
First, plugging , and in to Lemma 7, we have
Taking expectations on both the sides, we get,
where the last inequality follows from an application of Jensen’s inequality.
Similarly, plugging , in Lemma 8, we get
Raising both the sides with the exponent of , we get
Taking expectations on both the sides and rearranging the terms implies that
where the last inequality follows from an application of the Jensen’s inequality.
The final statement follows from the above inequality by using the definition of and .
D.4 Stochastic negative curvature search
The following lemma establishes properties of the negative curvature search step used in Algorithm 4 and Algorithm 5.
where is an independent Rademacher random variable and is an arbitrary unit vector such that . Then, the point satisfies
where and are taken with respect to the randomness in and .
In the second case, Taylor expansion for at implies that
Taking expectation on both the sides gives the desired statement:
The following lemma establishes properties of Oja’s method (), as used in Algorithm 4.
, and .
if , then and .
Moreover, when invoked as above, the procedure uses at most
queries to the stochastic Hessian-vector product oracle.
Appendix E Upper bounds for finding ϵ\epsilon-stationary points
In the following, we first show that Algorithm 2 returns a point such that, . We then bound the expected number of oracle queries used throughout the execution. In the proof, we show convergence to a -stationary point. A simple change of variable, i.e. running Algorithm 2 with , returns a point that enjoys the guarantee that .
where the last inequality follows from the fact that . Next, taking expectation on both the sides (with respect to the stochasticity of the oracle and the algorithm’s internal randomization), we get
Using Lemma 2, we have for all . Dividing both the sides by , and plugging in the value of the parameters and , we get,
Thus, for chosen uniformly at random from the set , we have
Finally, Markov’s inequality implies that with probability at least ,
Algorithm 2 queries the stochastic oracle in only when it invokes HVP-RVR in Line 5 to compute the gradient estimate at time . Let denote the total number of oracle calls made up until time . Invoking Lemma 3 to bound the expected number of stochastic oracle calls for each , and ignoring all the mutiplicative constants, we get
where is given by plugging in the update rule from Line 6 and by dropping multiplicative constants, is given by rearranging the terms, plugging in the value of and using that (to simplify the ceiling operator) under the assumption , and follows by observing that
as a consequence of Lemma 2 and the bound in (33). Next, note that since we assume , and since we have , the parameter is equal to (as this is smaller than ). Thus, plugging the value of and in the bound (35), we get,
Using Markov’s inequality, we have that with probability at least ,
The final statement follows by taking a union bound with failure probabilities for (34) and (36). ∎
E.2 Proof of Theorem 2
In the following, we first show that Algorithm 3 returns a point , such that with probability at least , . We then bound, with probability at least , the total number of oracle queries made up until time .
Note that, using Lemma 2 and Lemma 4, we have for all ,
Thus, for each , invoking Lemma 9 and plugging in the bounds from (37), and using the value of , we get
Telescoping this inequality from to , we have that
where the equality follows because is sampled uniformly at random from the set . Next, using the fact that, , rearranging the terms, and plugging in the value of , we get
Thus, with probability at least ,
Algorithm 3 queries the stochastic oracle in Line 6 and Line 7 only to compute the respective Hessian and gradient estimates. Let and denote the total number of stochastic oracle queries made by Line 6 and Line 7 till time respectively. Further, Let denote the total number of oracle queries made till time .
In what follows, we first bound and . Then, we invoke Markov’s inequality to deduce that the desired bound on holds with probability at least .
Bound on . Since the algorithm queries the stochastic Hessian oracle times per iteration, . Plugging the values of , and as specified in Algorithm 3, and ignoring multiplicative constant, we get,
where the first inequality above follows from the fact that under the natural choice for the precision parameter and using the identity for .
Bound on . Invoking Lemma 3 for each , we get
where follows by observing due to the update rule in Line 8 and is given by plugging in the value of for the natural choice of parameter . Next, note that since , and since we assume , the parameter is equal to (which is smaller than ). Thus, plugging the value of and in the bound (40), we get
where the second equality follows by using that to simplify the term .
Adding (41) and (39), the total number of oracle queries made by Algorithm 3 till time is bounded, in expectation, by
Using Markov’s inequality, we get that, with probability at least ,
The final statement follows by taking a union bound for the failure probability of (38) and (42).
Appendix F Upper bounds for finding (ϵ,γ)(\epsilon,\gamma)-second-order-stationary points
To begin, note that, for any , there are two scenarios: (a) either and is set using the update rule in Line 8, or, (b) and we set using Line 11, respectively. We analyze the two cases separately below.
Taking expectation on both the sides, while conditioning on the event that , we get
where the last inequality follows using Lemma 2.
Let denote the event that succeeds at time , in the sense that the event in Lemma 11 holds: if then , and otherwise, satisfies .
Then, using Lemma 12, we are guaranteed that
In particular, we are guaranteed by Lemma 11 that
Combining the two cases ( and ) from (43) and (44) above, we get
Using that and that , we have
Telescoping this inequality for from to and using the bound , we get
where follows because is sampled uniformly at random from and follows from Lemma 14. Rearranging the terms, we get
For any , there are two scenarios, either (a) and we go through Line 8, or (b) and Line 18 is executed. Thus,
We denote the two terms on the right hand side above by and , respectively. We bound them separately as follows.
Bound on . Using Lemma 3 with the fact that , we get
where is given by plugging in . The inequality follows by using the bound on from Lemma 14.
Bound on . Using Lemma 3 with the fact that , we get
where follows by plugging in the update rule from Line 8 (when ), follows by rearranging the terms and using the bound on from Lemma 14, and is follows from the choices of (in particular, our assumption that implies that ) and , as well as the following bound for :
where the last inequality is uses Lemma 2 and Lemma 13.
Combining the bounds from (50) and (51) in (49), we have
Using the law of total probability with the observation that Algorithm 4 enters Line 11 only if , we get
where denotes the number of oracle queries made by , the last inequality follows by bounding as in (47). Note that Lemma 11 implies that for ,
Plugging in the value of from Algorithm 4 and from (54), and using Markov’s inequality, we get that, with probability at least ,
The final statement follows by taking a union bound for the failure probability of the claims in (48) and (55). ∎
Under the setting of Theorem 4, we are guaranteed that
Taking conditional expectations, this further implies that
Otherwise, using a third-order Taylor expansion, and following the same reasoning as the proof of Lemma 10, we have
Combining this bound with the earlier inequalities (and being rather loose with constants), we conclude that
Under the same setting as Theorem 4, the point returned by Algorithm 4 satisfies
Starting from (46) in the proof of Theorem 4, we have
Telescoping this inequality for from to and using that , we get
where the last inequality follows from Lemma 14. Rearranging the terms, we get
where the last inequality uses that . ∎
For the values of the parameters and specified in Algorithm 4,
Since, and , we have that
Thus, using the fact that for all , we get
Consequently, by plugging in the values of and , we have
where the first inequality is due to (56). Similarly, we have that
Together, the above two bounds imply that
The bound on follows similarly. ∎
F.2 Full statement and proof for Algorithm 5
Before we delve into the proof, first note that using Lemma 2, we have for all ,
Further, using Lemma 4 with our choice of and , we have, for all ,
To begin the proof, we observe that for any , there are two scenarios: (a) either and the algorithm goes through Line 12, or, (b) and the algorithm goes through Line 18. We analyze the two cases separately below.
Case 1: . In this case, we set using the update rule in Line 12. Invoking Lemma 9 with the bound in (57) and , we get
Combining the two cases ( or ) from (58) and (59) above, we get
Telescoping the inequality above for from 0 to , and using the bound , we get
where the inequality in follows from Lemma 15. The inequality in is given by ignoring the (non-negative) terms and on the right-hand side and using the fact that . Finally, follows by recalling the definition of as samples uniformly at random from the set . Rearranging the terms, we get
Let us first introduce some notation to count the number of oracle calls made in each iteration of the algorithm.
On Line 13 and Line 19, Algorithm 5 queries the stochastic oracle through the subroutine HVP-RVR-Gradient-Estimator. Let denote the total number of oracle queries resulting from either line at iteration .
Let and denote the total number of oracle calls made by Line 11 and Line 15 at iteration to compute and respectively.
Define , and by , and respectively. In what follows, we give separate bounds for , and . The final statement on the total number of oracle calls follows by an application of Markov’s inequality.
For any , there are two scenarios, either (a) and we update through Line 12, or (b) and we update through Line 18 orLine 21. Thus, using the law of total expectation
We denote the two terms on the right hand side above by and , respectively. We bound them separately in as follows.
Bound on . Using Lemma 3 with the fact that , we get
where holds because when , we either have (if we follow the update rule in Line 18) or (if we follow Line 21). The inequality uses the bound on from Lemma 15 and follows from plugging in the value of .
Bound on . Using Lemma 3 with the definition , we get
where is given by the update rule from Line 12 and the fact that HVP-RVR-Gradient-Estimator uses parameter in this case, and follows by using the bound on from Lemma 15. The inequality follows because for the choice of parameters and and the assumed range of in the theorem statement, . Finally, the inequality is given by plugging in the value of and using that .
Plugging the bound in (63) and (64) back in (62), we get
For each , Algorithm 5 samples an independent Bernoulli with bias and executes Line 11 if . For every such pass through Line 11, the algorithm queries the stochastic Hessian oracle times. Thus,
where follows by plugging in the values of and as specified in Algorithm 5 (using that to simplify), and using the bound on from Lemma 15 .
The algorithm executes Line 15 only if , which happens with probability . For every such pass through Line 15, the algorithm queries the stochastic Hessian oracle times. Consequently,
where follows by plugging in the values of as specified in Algorithm 5, and using the bound on from Lemma 15.
Adding together all the bounds above (from (65), (66), and (67)), we have that the total number of oracle queries by Algorithm 5 till time is bounded in expectation by
Using Markov’s inequality, this implies that with probability at least ,
The final statement follows by union bound, using the failure probabilities for (61) and (68). ∎
For the values of the parameters and specified in Algorithm 5, we have
Under the assumption that , we have that
Thus, using the fact that for any , we get
Thus, plugging in the value of and , we get
where the first inequality is due to (69). Similarly, we have that
Together, the above two bounds imply that
The bound on follows similarly. ∎
Appendix G Lower bounds
A collection of derivative estimators for a function forms a probability- zero-chain if
We note that the constant is used here for compatibility with the analysis in Arjevani et al. 2019a. Any non-negative constant less than would suffice in its place. The next lemma formalizes the idea that any zero-respecting algorithm interacting with a probabilistic zero-chain must wait many rounds to activate all the coordinates.
The proof of Lemma 16 is a simple adaptation of the proof of Lemma 1 of Arjevani et al. 2019a to high-order zero-respecting methods—we provide it here for completeness. The proof idea is that any zero-respecting algorithm must activate coordinates in sequence, and must wait on average at least rounds between activations, leading to a total wait time of rounds.
Let denote the oracle responses for the th query made at the point , and let be the natural filtration for the algorithm’s iterates, the oracle randomness, and the oracle answers up to time . We measure the progress of the algorithm through two quantities:
Note that is the largest non-zero coordinate in , and that and . Thus, for any zero-respecting algorithm
for all . Moreover, observe that with probability one,
where the first inequality follows by the zero-chain property. Further, using the -zero chain property, it follows that conditioned on , with probability at least ,
Combining (73) and (74), we have that conditioned on ,
Thus, denoting the increments , we have via the Chernoff method,
Thus, for all ; combined with (72), this yields the desired result. ∎
where the component functions and are
We start by collecting some relevant properties of .
, where .
Parts 1 and 2 follow from Lemma 3 in Carmon et al. 2019a and its proof; Part 3 is proven in Section G.1.1; Part 4 follows from Observation 3 in Carmon et al. 2019a and Part 5 is the same as Lemma 2 in Carmon et al. 2019a. ∎
The derivative estimators we use are defined as
The estimators form a probability- zero-chain, are unbiased for , and satisfy
where the final inequality is due to Lemma 17.3, establishing the variance bound in (78). ∎
for some scalars and to be determined. The relevant properties of scale as follows
We bound the variance of the scaled derivative estimators as
where the last inequality follows by Lemma 18. Our goal now is to meet the following set of constraints:
Since is the only degree of freedom which can be tuned to meet though (not necessarily activate) the -constraint for and the -constraints for , we are forced to set
Lastly, we activate the -constraint by setting
where uses whenever , implying the desired bound. Lastly, we note that one can obtain tight lower complexity bounds for deterministic oracles by setting . Following the same chain of inequalities as in (G.1), in this case we get a lower oracle-complexity bound of
where the penultimate inequality is due to Lemma 1 of Carmon et al. 2019a. Therefore, for a fixed , we have
where follows from the definition of the operator norm, follows by the chain-like structure of , and follows from (86), concluding the proof.
G.2 Proof of Theorem 6
In this section we prove Theorem 6 following the schema outlined in Section 4.2. We start by collecting all the relevant properties of and from the construction in (11).
The functions and satisfy the following properties:
The function is non-negative and its first- and second-order derivatives are bounded by
The function and its first- and second-order derivatives are bounded by
Parts 1-4 are immediate. Part 5 follows from Lemma 1 of Carmon et al. 2019a and by noting that
Using these basic properties of and , we establish the following properties of the construction (analogous to Lemma 17).
The function satisfies the following properties:
, with .
We prove the individual parts of the lemma one by one:
The proof follows along the same lines of Lemma 3 of Carmon et al. 2019a together with the derivative bounds stated in Lemma 19.4.
The claim follows using the same calculation as in Section G.1.1, with the derivative bounds replaced by those in Lemma 19.4, mutatis mutandis.
The claim follows Observation 3 in Carmon et al. 2019a, mutatis mutandis.
The following facts can be verified by a straightforward calculation:
for all .
for all .
Otherwise, if nothing is assumed on , then the same chain of inequalities, using , can be used to bound the minimal value of .
We employ similar derivative estimators to the proof of Theorem 3, only this time we provide a noiseless estimate for the gradient. Formally, we set
where we employ the same notation as in Lemma 16. The proof now proceeds along the same lines of the proof of Theorem 3. The estimators have variance bounded as
which can established the same fashion as Lemma 18 by invoking Lemma 20.3 and Lemma 20.4.
for scalars and to be determined. The relevant properties of are as follows:
for any . The variance of the scaled derivative estimators can be bounded as
where the last inequality is by (90). Our goal now is to meet the following set of constraints:
.
Since is the only degree of freedom which can be tuned to meet (though not necessarily activate) the -constraints for , and the -constraint for , we are forced to have
This constraint holds w.l.o.g. as also bounds the absolute value of the Hessian eigenvalues (in other words, any point is trivially -weakly convex). Lastly, we activate the -constraint, by setting
where uses that whenever , implying the desired result (note that this bound does not depend on and .).
If , we obtain the following lower complexity bound for noiseless oracles (where is effectively set to one), assuming (this holds without loss of generality, as we discuss above). As before, we set . The -constraint is satisfied under the same condition stated in (96). Thus, letting
it follows that our construction is -Lipschitz for any . Following the same chain of inequalities as in (G.2) yields an oracle complexity lower bound of
Note that this bound does not depend on .