Online Stochastic Gradient Descent with Arbitrary Initialization Solves Non-smooth, Non-convex Phase Retrieval

Yan Shuo Tan, Roman Vershynin

Introduction

This problem is well motivated by practical concerns, having applications to Coherent Diffraction Imaging (CDI), Electron Microscopy, and X-ray Crystallography, and as such has been a topic of study from at least the early 1980s. We refer the reader to the survey papers for a comprehensive account of the contexts in which the problem arises, as well as the techniques that practitioners employ to solve it.

Over the last decade, the phase retrieval problem has also garnered substantial attention from the optimization and machine learning communities, the reason being that it can be formulated as a relatively benign non-convex optimization problem. In other words, it can be solved by minimizing the least squares objective:

Many papers have attempted to study how to optimize this objective given distributional assumptions on the sampling vectors . One popular approach is to use a two-step procedure: First, a spectral technique is used to obtain an initial estimate x(0)\textbf{x}^{(0)} so that its distance from a global minimum, ∥x(0)−x∗∥2\lVert\textbf{x}^{(0)}-\textbf{x}^{*}\rVert_{2}, is bounded above by a small constant. Next, the estimate is refined to arbitrary precision using an iterative method such as gradient descent or stochastic gradient descent (SGD). This procedure is well-supported by theoretical guarantees for both real and complex signals. Here we note that one can also construct loss functions for phase retrieval which are different from (1.2). Running variations of gradient descent or SGD on these functions also leads to provable guarantees (see for instance ).

Nonetheless, this state of affairs is not entirely satisfactory from an optimization theory perspective, since the use of a spectral initialization diminishes the novelty of being able to provably minimize a non-convex objective using first-order methods. The spectral initialization essentially allows the first-order method to begin within a “basin of convexity”, thus artificially escaping the difficulties of non-convexity. Such a deus ex machina may not be available when trying to optimize other non-convex functions.

More recently, the authors of showed that vanilla gradient descent converges in O(log⁡d)O(\log d) iterations, again given N=Ω(d⋅polylog(d))N=\Omega(d\cdot\textnormal{polylog}(d)) measurements. Up to log factors, this matches the running time guarantees for the original two-stage method. While an important step, their analysis still requires full gradient updates and does not apply to SGD. In the high-dimensional setting, SGD is particularly advantageous because it can be applied in an online, streaming fashion. This lowers the space complexity of the algorithm from O(Nd)O(Nd) to O(d)O(d), and allows progress toward the solution to be made even before the analyst gains access to the full data sample.

Moreover, with respect to the non-smooth amplitude least squares objective

the question of even gradient descent convergence remains open. This objective is especially interesting because numerical simulations have shown gradient descent and SGD with respect to it to succeed with fewer measurements than are necessary for the alternative objective (1.1) (See and .)

In this paper, we prove that for a real signal vector and real sampling vectors, online stochastic gradient descent for the non-smooth objective (1.2) converges to a global minimum from arbitrary initializations given Ω(dlog⁡d)\Omega(d\log d) Gaussian measurements. We believe our work to be among the first in establishing convergence of SGD in the non-smooth, non-convex regime. Furthermore, it will be readily apparent that our analysis framework generalizes easily to other single index models, and we conjecture that similar techniques will also work low-rank matrix sensing models in general.

We perform SGD with respect to (1.2), using a single data point per iteration. More formally, we form a sequence of signal estimates x(0),x(1),x(2),…\textbf{x}^{(0)},\textbf{x}^{(1)},\textbf{x}^{(2)},\ldots with the update rule:

We typically choose η=1d\eta=\frac{1}{d}. This is the same as in previous work that analyzed SGD as part of the two-step approach (see ), and allows for a clear geometric interpretation. At each step kk, we receive the datum (a(k),b(k))\left(\textbf{a}^{(k)},b^{(k)}\right); the solution set to the corresponding equation ∣⟨a(k),x⟩∣=b(k)\lvert\langle\textbf{a}^{(k)},\textbf{x}\rangle\rvert=b^{(k)} is then the union of two parallel hyperplanes. Taking an SGD step projects the current iterate x(k−1)\textbf{x}^{(k-1)} onto the closer of these hyperplanes. This iterative projection is strongly reminiscent of the randomized Kaczmarz algorithm for solving linear systems. For a more in-depth discussion of this connection, we again refer the reader to .

We make the following assumptions for the rest of the paper.

Assume the following for each positive integer kk:

(Fresh measurements) At step kk of the algorithm, we use a sampling vector a(k)\textbf{a}^{(k)} that is fully independent of the previous measurements a(1),…,a(k−1)\textbf{a}^{(1)},\ldots,\textbf{a}^{(k-1)}.

(No noise) We have bk=∣⟨a(k),x∗⟩∣b_{k}=\lvert\langle\textbf{a}^{(k)},\textbf{x}^{*}\rangle\rvert.

The following is the main result of the paper.

where σ=sign(⟨x(k),x∗⟩)\sigma=\textnormal{sign}(\langle\textbf{x}^{(k)},\textbf{x}^{*}\rangle). Furthermore, there is some constant DD, such that if d≥Dd\geq D, we may choose η0=1\eta_{0}=1.

This theorem tells us that in TT steps, SGD brings us to a point in the “basin of convexity” around a global minimum, within which we get linear convergence. As a consequence of this theorem, it is easy to see that we can get an ϵ\epsilon-relative-error estimate using O(dlog⁡(d/ϵ))O(d\log(d/\epsilon)) measurements (and the same number of steps). In fact, the analysis in tells us that within the “basin of convexity”, linear convergence still holds if we resample from amongst O(d)O(d) measurements. By doing this instead, our final sample complexity is O(dlog⁡d)O(d\log d).

The novelty of this result is not simply the convergence of the algorithm, but rather its convergence with near optimal sample and time complexity. Indeed, it is well known that as we let the step size decrease to zero, SGD under our assumptions approximates gradient flow for the population loss function. It is easy to show that gradient flow converges to a global minimum. With smaller steps, however, the algorithm takes a longer time to converge.

In addition, our chosen step size 1d\frac{1}{d} is the smoothness parameter for SGD if we were optimizing a least squares system objective under the same assumptions on a(k)\textbf{a}^{(k)}. It is interesting that although our objective is no longer smooth, 1d\frac{1}{d} nonetheless remains the right scaling.

The proof of the result is somewhat complicated and makes use of several new ideas. The first idea is to think of the sequence of SGD iterates as a Markov chain on a two-dimensional summary state space Y\mathcal{Y}. The two coordinates are the squared Euclidean norm of the iterate, r2=∥x∥2r^{2}=\lVert\textbf{x}\rVert^{2}, and the correlation with the signal, s=⟨x,x∗∥x∗∥⟩s=\langle\textbf{x},\frac{\textbf{x}^{*}}{\lVert\textbf{x}^{*}\rVert}\rangle. Just as how in thermodynamics state variables such as temperature, pressure, and volume suffice to determine the evolution of a thermodynamic system, so too in our case do the state variables r2r^{2} and ss suffice to determine the progress of the SGD algorithm.

The state space is obviously independent of the dimension dd. In fact, one can show that the distribution of the update in the state space is effectively independent of the dimension up to overall scaling (see Theorem 2.1). As such, as dd tends to infinity, the stochastic dynamics of the Markov chain when initialized at a fixed y∈Y\textbf{y}\in\mathcal{Y} approximates the solution of an ODE system whose corresponding vector field is given by the rescaled drift of the process.

Unfortunately, we are hit with the curse of dimensionality: if we take a random initialization x(0)∼N(0,I)\textbf{x}^{(0)}\sim\mathcal{N}(0,\textbf{I}), then it is well-known that with high probability,

This implies that the initial correlation with the signal decays as the dimension increases, and in fact, at the corresponding point (r02,s0)(r_{0}^{2},s_{0}) in the state space, the drift and the fluctuations have the same magnitude, and it is no longer appropriate to approximate the stochastic dynamics with a deterministic process.

Overcoming this is the most difficult part of the proof. To do so, we use a small-ball probability argument. In other words, we show that the distribution of sK=s(x(K))s_{K}=s(\textbf{x}^{(K)}) is anti-concentrated away from 0 when KK is large enough. This involves comparing the process s0,s1,…s_{0},s_{1},\ldots with a more well-understood process s^0,s^1,…\hat{s}_{0},\hat{s}_{1},\ldots via stochastic dominance. The distribution of this new sequence can in turn be controlled via recursive inequalities bounding 4th moments from above and 2nd moments from below. We conclude by applying the Paley-Zygmund inequality.

2 Related Work

There is already a large body of work on phase retrieval, and it is impossible to give a full account of the literature. We have already mentioned survey papers for how phase retrieval arises in various engineering problems. On the theoretical side, we have already discussed the two-step non-convex optimization approach, and will further mention here the convex relaxation approaches pioneered in the papers .

2.2 SGD as a Markov chain

It has long be observed that constant step-size SGD can be thought of as a Markov chain, and there seems to be a resurgence of interest in this view of SGD. For instance, uses this approach to analyze the limiting distribution of SGD iterates for strongly convex functions. Furthermore, uses this interpretation to see how SGD can be used as a sampling algorithm. Both these works have drawn inspiration from the recent body of work on Langevin algorithms for sampling from log-concave distributions. The idea of analyzing SGD through diffusion approximation is also present in .

2.3 Non-convex optimization and first-order methods

The Kaczmarz method is a classical method in numerical analysis for solving large scale overdetermined linear systems. A randomized version of it was first analyzed by . In our earlier work , we proposed adapting the method to the setting of phase retrieval, where it coincides with SGD under the Gaussian measurement setting,

Stochastic first-order methods have emerged as the optimization method of choice for modern machine learning. In particular, deep neural networks are trained almost exclusively using SGD and variants like ADAM. The loss functions for these models, however, are non-convex functions, for which there has traditionally been little theory on how first-order methods behave.

Unsurprisingly, there has been a concerted push over the last few years to address this issue. One line of work studies how gradient descent or SGD can be made to escape saddle points quickly (see ). Another line of work has focused on identifying regimes for shallow and deep neural networks for which gradient descent or SGD can be shown to converge (see for instance ).

3 Notation

Outline of proof

We start by making some simplifying assumptions. First, note that the algorithm and our guarantee (1.4) are both invariant with respect to scaling and rotation. As such, we may assume without loss of generality that x∗=e1\textbf{x}^{*}=\textbf{e}_{1}, the first coordinate basis vector. We also only analyze the case where dd is large enough so that η0=1\eta_{0}=1 and the step size is set to be η=1d\eta=\frac{1}{d}. The extension to smaller dd will be obvious.

The reason we choose to use r2r^{2} instead of rr is for the convenience of obtaining formulas for the stochastic update, as will be evident later. We further define θ=θ(x)≔arccos⁡(s/r)\theta=\theta(\textbf{x})\coloneqq\arccos(s/r). This is the smaller angle between x and x∗\textbf{x}^{*}.

Note that we can track the progress of SGD purely in terms of the state variables. Indeed, we have

so that the error of the kk-th step estimate x(k)\textbf{x}^{(k)} is equal to Ψ(π(x(k)))\Psi(\pi(\textbf{x}^{(k)})). Note that −x∗-\textbf{x}^{*} and x∗\textbf{x}^{*} are mapped onto (1,−1)(1,-1) and (1,1)(1,1), so that Ψ\Psi is uniquely minimized at these values. We hence wish to show that r2r^{2} and ss coordinates of our iterates converge to 11 and ±1\pm 1 respectively.

The sequence y(0),y(1),y(2),…\textbf{y}^{(0)},\textbf{y}^{(1)},\textbf{y}^{(2)},\ldots is a Markov chain on (Y,B(Y))(\mathcal{Y},\mathcal{B}(\mathcal{Y})) whose transition kernel has the random mapping representation

This theorem tells us that the state space sequence y(0),y(1),y(2),…\textbf{y}^{(0)},\textbf{y}^{(1)},\textbf{y}^{(2)},\ldots suffices not just to track our progress, as discussed earlier in the section, but also to determine its own dynamics. We hence no longer need to concern ourselves with the original SGD sequence, and instead work with this object for the rest of the paper. Henceforth, we let {Fk}k\{\mathcal{F}_{k}\}_{k} denote the filtration defined by this sequence.

2 Doob decomposition and continuous time limit as d→∞→𝑑d\to\infty

Let us try to understand the random mappings (2.3) and (2.4) better. It is well-known that (u,v)(u,v) converges in distribution to a standard 2-dimensional Gaussian N(0,I2)\mathcal{N}(0,\textbf{I}_{2}) as the ambient dimension dd tends to infinity. Therefore, the only essential dependence of the update formula (2.2) on dd is through the overall 1d\frac{1}{d} scaling. If we think of the indices k=1,2,…k=1,2,\ldots as a time variable, rescale time by a factor of 1d\frac{1}{d}, we can think of the sequence as being generated by an Euler discretization of a continuous-time process.

While we do not actually take this approach in our rigorous analysis, is it instructive to see what intuition this gives us. To do this, we do a Doob decomposition of the process {y(k)}k=0∞\left\{\textbf{y}^{(k)}\right\}_{k=0}^{\infty}, separating it into a drift term and a fluctuation term. Denote the drift terms using

Letting (αj,βj)(\alpha_{j},\beta_{j}) denote the random mapping used in the jj-th step of the Markov chain, we have

We now try to do a heuristic comparison of the relative magnitudes of the two terms. Suppose kk is small enough so that we have y(j)≈y(0)\textbf{y}^{(j)}\approx\textbf{y}^{(0)} for j=1,…,kj=1,\ldots,k. Then the drift can be approximated by

so that the fluctuation term has standard deviation approximately equal to

Therefore, for any fixed y(0)\textbf{y}^{(0)}, we see that the drift dominates the fluctuations as dd tends to infinity. This means that the continuous time limit of the process trajectory should be an integral curve associated to the vector field on the state space Y\mathcal{Y} defined by (αˉ,βˉ)\left(\bar{\alpha},\bar{\beta}\right). While this picture is incomplete, it offers a good first approximation, and the next step we take is to analyze the solutions to this first order ODE system.

3 Drift in continuous time limit

Miraculously, it is actually possible to derive a closed form formula for the vector field. We state it in the following lemma.

Studying the vector field plot in Figure 1, it is obvious that y∗≔(1,1)\textbf{y}^{*}\coloneqq(1,1) and −y∗=(−1,1)-\textbf{y}^{*}=(-1,1) are the only attracting fixed points, with basins of attraction the sets Y+≔Y∩{s>0}\mathcal{Y}_{+}\coloneqq\mathcal{Y}\cap\{s>0\} and Y−≔Y∩{s<0}\mathcal{Y}_{-}\coloneqq\mathcal{Y}\cap\{s<0\} respectively. While this assures us that the system has the right qualitative long-term behavior, the visualization alone is not sufficient to give quantitative bounds on convergence rates. This analysis turns out to be somewhat tricky. Given an integral curve yˉ(t)=(rˉt2,sˉt)\bar{\textbf{y}}^{(t)}=(\bar{r}_{t}^{2},\bar{s}_{t}) starting from an arbitrary initialization yˉ(0)\bar{\textbf{y}}^{(0)}, we will analyze its convergence rate by breaking it into three separate phases, as depicted in the figure.

To demarcate the phases, we define two stopping times as follows. We let τˉ1\bar{\tau}_{1} be the earliest time tt for which ∣rˉt2−1∣≤0.1\left\lvert\bar{r}_{t}^{2}-1\right\rvert\leq 0.1, and we let τˉ2\bar{\tau}_{2} be the earliest time tt for which the Lyapunov function Ψ\Psi defined in (2.1) satisfies Ψ(yˉ(t))≤0.2\Psi(\bar{\textbf{y}}^{(t)})\leq 0.2. Phase 1 is then the portion of the curve traversed between time and time τˉ1\bar{\tau}_{1}, Phase 2 the portion traversed between time τˉ1\bar{\tau}_{1} and time τˉ2\bar{\tau}_{2}, with Phase 3 the remainder of the curve traversed after τˉ2\bar{\tau}_{2}.

Let us compute the duration of Phase 1, which is the same as bounding τˉ1\bar{\tau}_{1}. To do this, we solve (2.6) to get rˉt2−1=e−t(rˉ02−1)\bar{r}_{t}^{2}-1=e^{-t}(\bar{r}_{0}^{2}-1), so that τˉ1≲log⁡(rˉ02)∨1\bar{\tau}_{1}\lesssim\log(\bar{r}_{0}^{2})\vee 1. The second phase is trickier due to the unwieldiness of (2.7). As such, we compute more amenable lower and upper bounds for the expression as follows.

There is a constant bˉmax\bar{b}_{max} such that we have

Furthermore, for any ϵ>0\epsilon>0 small enough, there is some η=η(ϵ)>0\eta=\eta(\epsilon)>0 and some constant bˉmin=bˉmin(ϵ)>0\bar{b}_{min}=\bar{b}_{min}(\epsilon)>0 such that

where D≔{(r2,s)∈Y  ⁣: ∣s∣≤1−ϵ,∣r2−1∣≤η}\mathcal{D}\coloneqq\{(r^{2},s)\in\mathcal{Y}~{}\colon~{}\lvert s\rvert\leq 1-\epsilon,\lvert r^{2}-1\rvert\leq\eta\}.

One can show that η(0.1)≥0.1\eta(0.1)\geq 0.1, and therefore, we have the bound dsˉtdt≥bˉminsˉ\frac{d\bar{s}_{t}}{dt}\geq\bar{b}_{min}\bar{s} for τˉ1≤t≤τˉ2\bar{\tau}_{1}\leq t\leq\bar{\tau}_{2}. Solving this gives sˉt=sˉτˉ1ebˉmin(t−τˉ1)\bar{s}_{t}=\bar{s}_{\bar{\tau}_{1}}e^{\bar{b}_{min}(t-\bar{\tau}_{1})}, and we have the estimate τˉ2−τˉ1≲log⁡(1/∣sˉτˉ1∣)∨1\bar{\tau}_{2}-\bar{\tau}_{1}\lesssim\log(1/\lvert\bar{s}_{\bar{\tau}_{1}}\rvert)\vee 1.

Finally, Phase 3 corresponds to portion of the integral curve that lies within the “basin of convexity” around y∗\textbf{y}^{*}. Indeed we compute:

Here, the inequality in the third line comes from a relative bound on the error term 1π(2θˉt−sin⁡(2θˉt))\frac{1}{\pi}\left(2\bar{\theta}_{t}-\sin(2\bar{\theta}_{t})\right) provided by the geometry of the basin region. As such, we also get linear convergence Ψ(yˉ(t))≤Ψ(yˉ(τˉ2))e−c(t−τˉ2)\Psi(\bar{\textbf{y}}^{(t)})\leq\Psi(\bar{\textbf{y}}^{(\bar{\tau}_{2})})e^{-c(t-\bar{\tau}_{2})}.

Putting everything together, we see that for any ϵ>0\epsilon>0, if we would like Ψ(yˉt)≤ϵ\Psi(\bar{\textbf{y}}_{t})\leq\epsilon, it suffices for

4 Discretizing the drift

We now examine what this means for the Markov chain (rk2,sk)=y(k)=y(k,d)(r^{2}_{k},s_{k})=\textbf{y}^{(k)}=\textbf{y}^{(k,d)}, where for clarity, we have made the dependence on dd in (2.2) explicit as a component of the indexing. We have argued that the integral curve yˉ(t)\bar{\textbf{y}}^{(t)} is the limit of the trajectories {t↦y(⌊t/d⌋,d)}\{t\mapsto\textbf{y}^{(\lfloor t/d\rfloor,d)}\} as dd tends to infinity. The number of iterations kk needed for convergence of y(k)\textbf{y}^{(k)} to ±y∗\pm\textbf{y}^{*} is thus the quantity (2.8) scaled by a factor of dd.

There may also some additional dependence on dd implicit in the definitions of rˉ02\bar{r}_{0}^{2} and sˉτˉ1\bar{s}_{\bar{\tau}_{1}}. The first has to do with the conditioning of the problem, and it is reasonable to assume that rˉ02≲1\bar{r}_{0}^{2}\lesssim 1. On the other hand, if we take a random initialization, measure concentration inflicts us with the curse of dimensionality shown in (1.5), so that log⁡(1/∣sˉτˉ1∣)≳log⁡d\log(1/\lvert\bar{s}_{\bar{\tau}_{1}}\rvert)\gtrsim\log d. The final time complexity estimate is therefore O(dlog⁡(d/ϵ))O(d\log(d/\epsilon)).

5 Accounting for fluctuations

Comparing the time complexity estimate for the drift process given in the last section with the guarantee given by Theorem 1.2 suggests to us that the fluctuations do not ultimately affect the rate of convergence for the random process y(k)\textbf{y}^{(k)}. This is true, and indeed, we may make our heuristic arguments rigorous and show that the fluctuation “error term” is dominated by the drift over the discretized versions of Phase 1 and Phase 3.

The situation is more complicated for Phase 2. While the length of the phase for the random process remains the same, the underlying dynamics of the random process is very different from that of the deterministic drift process. This is because when ∣sk∣≲1/d\lvert s_{k}\rvert\lesssim 1/\sqrt{d}, as would be the case under random initialization, the magnitude of the fluctuations are of the same order as the drift, thereby invalidating the approximation argument.

There is no hope of upper bounding the fluctuations of y(k)\textbf{y}^{(k)} around iteration k≈dτˉ1k\approx d\bar{\tau}_{1}, so we change course and instead aim at bounding the cumulative variance of the fluctuations from below. We seek a small ball probability bound for the horizontal marginal sdτˉ1+k′s_{d\bar{\tau}_{1}+k^{\prime}} for k′≍dlog⁡dk^{\prime}\asymp d\log d. More precisely, we would like

for some universal constants cc, η\eta and δ\delta. Once sks_{k} is of a constant distance away from zero, we may return to using the approximation argument and treat the fluctuations as an error term.

6 Small ball probability bound through diffusion approximation

Here, dB(t)d\textbf{B}^{(t)} is standard, two-dimensional Brownian motion, while for each y, Σ(y)\bm{\Sigma}(\textbf{y}) is a positive semidefinite matrix reflecting the fluctuation covariance at state y.

It is not easy to compute a closed form solution to this SDE. To perform a heuristic analysis, we instead solve a simplified form of the equation for the ss-marginal:

7 Outline for rest of paper

In this section so far, we have sketched a heuristic proof for SGD convergence, where the main idea was to consider the continuous time limit of the state space Markov chain, and solve the resulting differential equation or stochastic differential equation. In our rigorous analysis, we will not adopt this approach, and instead solve finite difference equations coming from the Markov chain while obtaining non-asymptotic control over the fluctuation process.

For the rest of this paper, we work with a fixed initialization y(0)\textbf{y}^{(0)}, and let y(1),y(1),y(2),…\textbf{y}^{(1)},\textbf{y}^{(1)},\textbf{y}^{(2)},\ldots be the sequence of iterates generated by repeated applying the Markov kernel (2.2). For each kk, we will use rk2r_{k}^{2} and sks_{k} to denote the coordinates of y(k)\textbf{y}^{(k)}. We will continue to do a multi-phase analysis of convergence, and as such, define analogues of the stopping times τˉ1\bar{\tau}_{1} and τˉ2\bar{\tau}_{2}. We set

Here, γ1\gamma_{1} and γ2\gamma_{2} are constants to be determined later.

We shall set T=τ2bT=\tau_{2b} in Theorem 1.2. As such, we wish to prove that with high probability, τ2b≲dlog⁡(d∥x(0)∥∨1)\tau_{2b}\lesssim d\log\left(d\lVert\textbf{x}^{(0)}\rVert\vee 1\right), and that linear convergence in expectation occurs after τ2b\tau_{2b}. The second statement follows easily from the theory we established in , while the first is, as mentioned before, the main result of this paper. Our strategy is to bound τ2b\tau_{2b} by bounding τ1\tau_{1}, τ2a−τ1\tau_{2a}-\tau_{1}, and τ2b−τ2a\tau_{2b}-\tau_{2a} separately.

In Section 3, we bound τ1\tau_{1}, and also establish uniform control over ∣rk2−1∣\lvert r_{k}^{2}-1\rvert, which is needed for the rest of the proof. In Sections 4 and 5, we bound τ2a−τ1\tau_{2a}-\tau_{1}. This is the most difficult part of the proof, and relies on developing a nuanced notion of stochastic dominance with which we compare the sequence {sk}k\{s_{k}\}_{k} with a carefully constructed sequence {s^k}k\{\hat{s}_{k}\}_{k}. We then apply a small ball probability argument to the latter. In Section 6, we bound τ2b−τ2a\tau_{2b}-\tau_{2a} by approximating the sequence with that obtained by removing the fluctuations. Finally, we will complete the proof of Theorem 1.2 in Section 7, before concluding and discussing the broader implications of the result in Section 8.

In this section, we will show that the sequence of squared norms, {rk2}k\{r_{k}^{2}\}_{k}, quickly converges to a small interval of width O(log⁡d/d)O(\log d/\sqrt{d}) around the value 1, and thereafter remains within this interval for at least Cdlog⁡dCd\log d iterations. In subsequent sections, we will show that this is sufficient time for the {sk}k\{s_{k}\}_{k} sequence also to converge. The reason why such uniform control is necessary is because the formula for the horizontal update (2.4) taken from a point y depends on r(y)r(\textbf{y}). By showing that {rk2}k\{r_{k}^{2}\}_{k} concentrates uniformly, we thereby also obtain control over the drift and fluctuations of {sk}k\{s_{k}\}_{k}, which enables us to do an essentially univariate analysis of this latter sequence.

There exists universal constants C1C_{1} and C2C_{2} such that τ1≤C2d⋅(log⁡d+log⁡∣r02−1∣)\tau_{1}\leq C_{2}d\cdot(\log d+\log\lvert r_{0}^{2}-1\rvert) with probability at least 1−C2/log⁡d1-C_{2}/\log d.

Next, we may obtain a recursive bound for the variance of the iterates using the law of total variance. We have

It easy to check that this quantity is bounded by C/dC/d whenever (3.2) holds. In this case, we may apply Chebyshev’s inequality to conclude that

with probability at least 1−δ1-\delta. We may also similarly bound rk2−1r_{k}^{2}-1 from below. Choosing δ=C/log⁡d\delta=C/\log d gives us the probability bound we want. ∎

After rkr_{k} has contracted to a value close to 1, we need to show that it remains close to 1 throughout the time scale needed for the algorithm to converge. Although it is clear from the formula (2.3) that the increments are subexponential, a naive union bound is not tight enough for our purposes. To overcome this, we make use of a maximal Bernstein inequality.

Let X1,X2,…,XMX_{1},X_{2},\ldots,X_{M} be a martingale difference sequence that is adapted to a filtration {Gt}\{\mathcal{G}_{t}\}. Suppose there is a constant K>0K>0 such that the following pointwise inequality holds almost surely for any time tt and 0<λ≤1/2K0<\lambda\leq 1/2K:

Denote St=∑i=1tXiS_{t}=\sum_{i=1}^{t}X_{i}, for t=1,2…t=1,2\ldots. Then for all ϵ>0\epsilon>0, we have the uniform tail bound

This result is an easy consequence of combining two classical arguments: the supermartingale inequality and an exponential martingale inequality. This argument also appears with more sophistication in the exponential line-crossing method recently developed by . The proof details are deferred to Appendix F. ∎

The tail bound on the right hand side of (3.4) is exactly the Bernstein tail at time MM. See .

Let W1,W2,…,WMW_{1},W_{2},\ldots,W_{M} be a real-valued stochastic process adapted to a filtration {Gt}\{\mathcal{G}_{t}\}. Suppose that there is some 0<ρ<10<\rho<1 such that for each time tt, we have

Furthermore, assume that there is a constant K>0K>0 such that for any tt, the following pointwise inequality holds almost surely for any time tt and 0<λ≤1/2K0<\lambda\leq 1/2K:

Then for all ϵ>0\epsilon>0, we have the uniform tail bound

Since X1,X2,…X_{1},X_{2},\ldots form a martingale difference sequence satisfying the assumptions of Lemma 3.2, its partial sums are bounded, and we may apply an easy combinatorial result (Lemma E.1) to bound the tails of WtW_{t}. ∎

We are now ready to state the guarantee that the radius remains uniformly close to 1.

For any constant C1C_{1}, M≤C1dlog⁡dM\leq C_{1}d\log d, we have

with probability at least 1−1d1-\frac{1}{d}, where C2C_{2} is a universal constant depending only on C1C_{1}. Furthermore, for M≤d2/log⁡dM\leq d^{2}/\log d, there is a constant C3C_{3} such that

For ease of notation, assume τ1=0\tau_{1}=0. Set Wk≔rk2−1W_{k}\coloneqq r_{k}^{2}-1 for k=0,1,2,…k=0,1,2,\ldots. Recall also that {Fk}k\{\mathcal{F}_{k}\}_{k} is the filtration generated by the stochastic updates. We want to show that these satisfy the assumptions of Lemma 3.4. By Lemma 2.2, we see that (3.5) is satisfied with ρ=1−1d\rho=1-\frac{1}{d}, and we would like to use Lemma B.1 to verify (3.6). However, Lemma B.1 is only valid when y lies in a bounded subset of Y\mathcal{Y}, and we don’t assume this a priori.

with probability at least 1−1d1-\frac{1}{d}. Meanwhile, on the same event, we have τ>M\tau>M, so that the bound also holds for the original sequence, thereby giving us (3.8). The proof of (3.9) is similar and is hence omitted. ∎

Let us denote τr≔min⁡{k ⁣:∣rτ1+k2−1∣≥C2log⁡d/d}\tau_{r}\coloneqq\min\{k\colon\lvert r_{\tau_{1}+k}^{2}-1\rvert\geq C_{2}\log d/\sqrt{d}\}. The previous lemma tells us that τr≥C1dlog⁡d\tau_{r}\geq C_{1}d\log d with probability at least 1−1/d1-1/d. In order to obtain the control over {sk}k\{s_{k}\}_{k} promised at the start of this section, we would ideally like to condition on this event, but doing so would destroy the Markov nature of the process y(k)\textbf{y}^{(k)} and invalidate the update formula (2.4).

Phase 2a: Stochastic dominance argument

We have broken up the analysis of “Phase 2” of the SGD process into two sub-phases. The first sub-phase concerns the portion of the process in which ∣sk∣=o(1)\lvert s_{k}\rvert=o(1), so that the drift and fluctuations are of comparable magnitudes. The goal of this section and the next is to bound the length of this phase by proving the following theorem.

There is some 0<γ1<1/20<\gamma_{1}<1/2 and success probability p>0p>0, such if we use this value of γ1\gamma_{1} in (2.11), for dd large enough, τ2a−τ1≤Cdlog⁡d\tau_{2a}-\tau_{1}\leq Cd\log d with probability at least pp.

For ease of notation and making use of the strong Markov property, we may assume again that τ1=0\tau_{1}=0. Our strategy for proving the theorem is to show that for iterations satisfying τ1≤k≤τ2a\tau_{1}\leq k\leq\tau_{2a}, a subsequence of the horizontal marginals of the process, {sBk}k\{s_{Bk}\}_{k}, stochastically dominate another process {s^k}k\{\hat{s}_{k}\}_{k} in magnitude. The second process will be constructed so that we will have precise control over its second and fourth moments, which will allow us to apply the Paley-Zygmund inequality to get a small ball probability argument. In this section, we will focus on constructing the comparison process and fleshing out the stochastic dominance argument.

In order to get control over moments, the process {s^k}k\{\hat{s}_{k}\}_{k} will be constructed to have approximately Gaussian increments. The comparison will thus be based on normal approximation using the Berry-Essen theory. It may be possible to obtain appropriate Berry-Essen bounds for martingale sequences, but the theory on this topic seems to be incomplete at the time of writing. We shall bypass this by showing that over epochs of an appropriate length BB, the update sk→sk+Bs_{k}\to s_{k+B} is well-approximated by the update we would get had we done a batch update, summing up BB independent steps taken from y(k)\textbf{y}^{(k)}. Since this is a sum of i.i.d. random variables, the classical Berry-Eseen bound can then be applied to this latter update.

More precisely, Let PP denote the random mapping defined by (1.3). It will be helpful to use this notation in this section, as we will have to deal with a few different stochastic sequences. Fix the epoch length to be B=d2/3log⁡dB=d^{2/3}\log d. We define a Markov kernel on the state space via the random mapping QQ:

where αk(y)\alpha_{k}(\textbf{y}) and βk(y)\beta_{k}(\textbf{y}) for k=1,2,…,Bk=1,2,\ldots,B are independent realizations of the random variables defined in (2.3) and (2.4). Denoting σ(y)2≔Var{β(y)}\sigma(\textbf{y})^{2}\coloneqq\text{Var}\{\beta(\textbf{y})\}, we note here that the drift and variance of this batch update satisfy the formulas

Fix y(m)\textbf{y}^{(m)} for τ1≤m≤τ1+τr\tau_{1}\leq m\leq\tau_{1}+\tau_{r}. With probability at least 1−1/d21-1/d^{2}, we have

For ease of notation, re-index so that m=0m=0. First, observe that we have the decomposition

Applying the maximal Bernstein bound for martingale sequences (Lemma 3.2) with M=BM=B and ϵ=Blog⁡d\epsilon=\sqrt{B\log d}, we get

Plugging (4) into the above equation, and using the bound (4.3), we get

We can simplify this recursive bound using some combinatorics. Applying Lemma E.2 with ρ=bˉmax/d\rho=\bar{b}_{max}/d, xt=∣βˉ(y(t))∣x_{t}=\lvert\bar{\beta}(\textbf{y}^{(t)})\rvert, and ξ=∣s0∣+(CBlog⁡d)/d\xi=\lvert s_{0}\rvert+(C\sqrt{B\log d})/d, we get

Plugging this back into (4), for any t≤Bt\leq B, we get

Fix y(m)\textbf{y}^{(m)} for τ1≤m≤τ1+τr\tau_{1}\leq m\leq\tau_{1}+\tau_{r}. For dd large enough, with probability at least 1−1/d21-1/d^{2}, there is a coupling such that we have

For ease of notation, we again re-index so that m=0m=0. Couple the random mappings QQ and PBP^{B} by using the same draws for u(t),v(t)u^{(t)},v^{(t)}, t=1,2…t=1,2\ldots. First condition on the probability 1−1/d21-1/d^{2} event promised by Lemma 4.2. Now write

where for each tt, At\mathcal{A}_{t} is the event that

We further condition on the probability 1−1/d21-1/d^{2} events in which we have

By Lemma 4.2, we see that the first term is bounded as follows:

The second term may be bounded similarly. To bound the third term, we use Lemma F.1. Observe that {1At}\{\textbf{1}_{\mathcal{A}_{t}}\} is a sequence of Bernoulli variables. If we condition on the high probability event promised by Lemma 4.2, we also have

In Lemma F.1, set M=BM=B, ϵ=1/2\epsilon=1/2, and θ\theta to be the right hand side of (4.4). This gives

for dd large enough. On this event, the third term is therefore also bounded by CB2log⁡dd2(∣s0∣+log⁡dB)\frac{CB^{2}\log d}{d^{2}}\left(\left\lvert s_{0}\right\rvert+\sqrt{\frac{\log d}{B}}\right).

By adjusting constants if necessary, we can make sure that the total error probability is bounded by 1/d21/d^{2}. ∎

In order to compare {sBk}k\{s_{Bk}\}_{k} with our (as yet undefined) reference sequence {s^k}k\{\hat{s}_{k}\}_{k}, we will use a more nuanced version of the usual notion of stochastic dominance.

Given two real-valued random variables XX and YY defined on the same probability space, we say that XX dominates YY in place up to error δ\delta, denoted X≥δYX\geq_{\delta}Y, if X≥YX\geq Y with probability at least 1−δ1-\delta. In addition, given any two real-valued random variables XX and YY, we say that XX stochastically dominates YY up to error δ\delta, if there is a coupling of XX and YY for which X≥δYX\geq_{\delta}Y. Denote this relation by X⪰δYX\succeq_{\delta}Y. Note that the usual notion of stochastic dominance is equivalent to X⪰0YX\succeq_{0}Y, which we will also denote using X⪰YX\succeq Y.

We first establish stochastic dominance for the conditional distribution of ∣sk+B∣\lvert s_{k+B}\rvert given y(k)\textbf{y}^{(k)} over that obtained from an approximately Gaussian increment. Following this, we will argue that stochastic dominance of individual steps also implies stochastic dominance for the entire process.

and let y be any point in D\mathcal{D}. Define the approximation error term

Then with δ=1/d2+C/B\delta=1/d^{2}+C/\sqrt{B}, we have

Fix y, and set WB≔1σ(y)B∑k=1B(βk(y)−βˉ(y))W_{B}\coloneqq\frac{1}{\sigma(\textbf{y})\sqrt{B}}\sum_{k=1}^{B}\left(\beta_{k}(\textbf{y})-\bar{\beta}(\textbf{y})\right), where σ(y)2≔Var{β(y)}\sigma(\textbf{y})^{2}\coloneqq\text{Var}\{\beta(\textbf{y})\}. The subexponential bound on β(y)\beta(\textbf{y}) implies bounds on the 3rd moments, so by Berry-Esseen, its distribution function FWF_{W} satisfies

where Φ\Phi is the distribution function of a standard normal random variable, and CC is an absolute constant.

Setting δ=C/B\delta=C/\sqrt{B}, we have have WB⪰δgW_{B}\succeq_{\delta}g by Lemma D.5. Since

for δ′=1/d2\delta^{\prime}=1/d^{2}. We may combine this with (4.8) using the transitivity of stochastic dominance (Lemma D.3), to get the bound we want in (4.7). ∎

and 0<κ<10<\kappa<1 a small constant to be determined later, while σ(y)2≔Var{β(y)}\sigma(\textbf{y})^{2}\coloneqq\text{Var}\{\beta(\textbf{y})\} as before. Again setting τ1=0\tau_{1}=0 for notational convenience, we define

Let s^0,s^1,s^2,…\hat{s}_{0},\hat{s}_{1},\hat{s}_{2},\ldots be the Markov process defined by the update rule (4.10). Then for any positive integer kk, we have

For convenience of notation, we define the auxiliary transitional kernel

Consider a fixed y∈D\textbf{y}\in\mathcal{D}. In Lemma 4.5, we showed that s(PB(y))2⪰δK(s(y),y)2s(P^{B}(\textbf{y}))^{2}\succeq_{\delta}K(s(\textbf{y}),\textbf{y})^{2}. If we choose κ\kappa small enough in (4.9), then by the drift lower bound in Lemma 2.3, replacing the kernel KK with LL simply reduces the magnitude of the drift while preserving the variance of the Gaussian increment and the magnitude of the soft-thresholding. Using the first part of Lemma D.6 thus gives the bound K(s(y),y)2⪰L(s(y),y)2K(s(\textbf{y}),\textbf{y})^{2}\succeq L(s(\textbf{y}),\textbf{y})^{2}. This implies through transitivity (Lemma D.3) that

Treating each epoch of updates as a single step, what we have showed is the stochastic dominance of a single step over a Gaussian increment, conditioned on a fixed starting value for y. It is not so easy, however, to conclude that stochastic dominance is preserved when composing multiple steps together. This is because we are working with transition kernels, and not sums of independent random variables. To overcome this, we need to use the second part of Lemma D.6, which tells us that

Let us now prove the claim by induction. The case k=0k=0 is clear since s^0=s0\hat{s}_{0}=s_{0}. Now assume that the statement holds for some kk. Let Gk\mathcal{G}_{k} be the σ\sigma-algebra generated by y(0),…y(k),s^02,…s^k2\textbf{y}^{(0)},\ldots\textbf{y}^{(k)},\hat{s}_{0}^{2},\ldots\hat{s}_{k}^{2}. Condition on Gk\mathcal{G}_{k} as well as the 1−kδ1-k\delta probability event for which

If τ2a≤kB\tau_{2a}\leq kB, then s(k+1)B∧τ2a=skB∧τ2as_{(k+1)B\wedge\tau_{2a}}=s_{kB\wedge\tau_{2a}} and s^(k+1)∧⌊τ2a/B⌋=s^k∧⌊τ2a/B⌋\hat{s}_{(k+1)\wedge\lfloor\tau_{2a}/B\rfloor}=\hat{s}_{k\wedge\lfloor\tau_{2a}/B\rfloor}, so that

on the same event. Otherwise, since y(kB)∈D\textbf{y}^{(kB)}\in\mathcal{D}, we have

Here, the first dominance bound follows from (4.12), while the second follows from (4.11). This means that we can construct a coupling of the update kernels such that s^k+12≤s(k+1)B2\hat{s}_{k+1}^{2}\leq s_{(k+1)B}^{2} with conditional probability at least 1−δ1-\delta.

If τ2a≥(k+1)B\tau_{2a}\geq(k+1)B, then this immediately implies that (4.13) holds. On the other hand, if kB<τ2a<(k+1)BkB<\tau_{2a}<(k+1)B, we have ⌊τ2a/B⌋=k\lfloor\tau_{2a}/B\rfloor=k and skB2<γ12s_{kB}^{2}<\gamma_{1}^{2}, so that

As such, the statement also holds for k+1k+1. ∎

Phase 2a: Small ball probability argument via Paley-Zygmund

Recall that our goal is to obtain a small ball probability bound for sk2s_{k}^{2}, for k≳dlog⁡dk\gtrsim d\log d, thereby proving Theorem 4.1. By the stochastic dominance argument in the last section, it suffices to obtain such a bound for s^k2\hat{s}_{k}^{2}, with the appropriate time rescaling. This is the plan for this section, and we start by observing the following general recursive bounds for moments of adapted sequences.

Let S1,S2,…S_{1},S_{2},\ldots be a process adapted to the filtration {Gt}\{\mathcal{G}_{t}\}. Then we have

The first equation is standard. To prove the second, we expand and use martingale orthogonality to write

Furthermore, the third term can be bounded as follows:

Plugging these into the original equation, we get

In order to apply these bounds to our sequence {s^k}k\{\hat{s}_{k}\}_{k}, we will need estimates of the moments of each increment.

In this proof we shall, for convenience, denote X≔b(s^k)+Bσ(y(kB))gX\coloneqq b(\hat{s}_{k})+\sqrt{B}\sigma(\textbf{y}^{(kB)})g and ϵ≔ϵ(s^k)\epsilon\coloneqq\epsilon(\hat{s}_{k}). Writing out the definition (4.10), we then have s^k+1=ρϵ[X]\hat{s}_{k+1}=\rho_{\epsilon}\left[X\right]. Using the definition of XX and the fact that y(kB)∈D\textbf{y}^{(kB)}\in\mathcal{D}, we compute

We need to show that soft-thresholding XX does not decrease its variance by two much, and will consider two cases depending on the value of s^k\hat{s}_{k}. We shall suppose WLOG that s^k>0\hat{s}_{k}>0.

First, suppose s^k≤1log⁡3d\hat{s}_{k}\leq\frac{1}{\log^{3}d}. Then

Plugging our assumption on s^k\hat{s}_{k} into (4.6), we get

By (5.6), the quantity on the right hand side is a lower bound for Var{X \vline FkB}\text{Var}\left\{X~{}\vline~{}\mathcal{F}_{kB}\right\}, giving us what we want.

Now assume instead that s^k≥1log⁡3d\hat{s}_{k}\geq\frac{1}{\log^{3}d}. Let X′X^{\prime} be an independent copy of XX. Using a well-known formula for variance, we have

where gg and g′g^{\prime} are independent standard normal random variables.

Next, to obtain (5.5), we simply apply Lemma F.2 to get

By our constraint that y(kB)∈D\textbf{y}^{(kB)}\in\mathcal{D}, we have σ(ykB)2≤Cd2\sigma(\textbf{y}^{kB})^{2}\leq\frac{C}{d^{2}}, thereby giving the bound we want. ∎

The second inequality is easier to prove, so we shall start with this. We may use Lemma F.2 to get

Plugging this bound together with (5.5) into (5.2) gives us (5.9).

Next, denoting X≔b(s^k)+Bσ(y(kB))gX\coloneqq b(\hat{s}_{k})+\sqrt{B}\sigma(\textbf{y}^{(kB)})g and ϵ≔ϵ(s^k)\epsilon\coloneqq\epsilon(\hat{s}_{k}) as in the previous lemma, observe that

It remains to compare the relative magnitudes of the two terms. For convenience, we reproduce the definition of ϵ(s^k)\epsilon(\hat{s}_{k}) here:

We will divide the proof into three cases, depending on the magnitude of s^k\hat{s}_{k}. WLOG, assume that s^k≥0\hat{s}_{k}\geq 0.

When s^k≥1d1/3\hat{s}_{k}\geq\frac{1}{d^{1/3}}, then

for dd large enough, and when log⁡4dd2/3≤s^k≤1d1/3\frac{\log^{4}d}{d^{2/3}}\leq\hat{s}_{k}\leq\frac{1}{d^{1/3}}, one may check that

In either case, we may plug the bound for ϵ(s^k)\epsilon(\hat{s}_{k}) into (5) to get

Next, when 0≤s^k≤log⁡4dd2/30\leq\hat{s}_{k}\leq\frac{\log^{4}d}{d^{2/3}}, we have

Furthermore, we can bound the second term on the right hand side via

Combining this with (5.1) and (5.4) gives (5.8). ∎

For convenience, denote A=1−1/log⁡2dA=1-1/\log^{2}d. We solve the recursion in (5.8) to get

Meanwhile, for any T≍CdBlog⁡d≍log⁡dlog⁡(1+κAB/d)T\asymp\frac{Cd}{B}\log d\asymp\frac{\log d}{\log(1+\kappa AB/d)}, we have

On the other hand, we may solve the second recursion (5.9) and apply a similar computation as before to get

Taking the ratio of (5.12) and (5.14), plugging in the time point k=Tk=T, we get

Note that in the last equation, we used the definition of AA and (5.13).

where the inequality holds for dd large enough. On this event, we have τ2a≤TB≤Cdlog⁡d\tau_{2a}\leq TB\leq Cd\log d as we wanted. ∎

Phase 2b: Approximation by drift process

In the previous two sections, we have bounded the duration of Phase 2a of the SGD process, that is, the time it takes for ∣sk∣\lvert s_{k}\rvert to increase to a constant value. The goal of this section is to bound the duration of Phase 2b in which the iterates converge to the “basin of convexity” around y∗\textbf{y}^{*} or −y∗-\textbf{y}^{*}. We will prove the following theorem.

We have τ2b−τ2a≤Cd\tau_{2b}-\tau_{2a}\leq Cd with probability at least 1−1/d1-1/d.

As mentioned in the overall proof outline, the idea is to understand the trajectory of the drift process, and then show that the fluctuations do not affect the trajectory by too much. For convenience, we condition on the event τ2a<∞\tau_{2a}<\infty, and then use the strong Markov property to re-index, setting τ2a=0\tau_{2a}=0.

The drift process {yˉ(k)}k\{\bar{\textbf{y}}^{(k)}\}_{k} is defined via the deterministic update

with the initialization yˉ(0)=y(0)\bar{\textbf{y}}^{(0)}=\textbf{y}^{(0)}. For simplicity, we denote rˉk≔r(yˉ(k))\bar{r}_{k}\coloneqq r(\bar{\textbf{y}}^{(k)}) and sˉk≔s(yˉ(k))\bar{s}_{k}\coloneqq s(\bar{\textbf{y}}^{(k)}).

Let τ=min⁡{k  ⁣: sˉk≥1−(γ2−ϵ)/2}\tau=\min\left\{k~{}\colon~{}\bar{s}_{k}\geq 1-(\gamma_{2}-\epsilon)/2\right\} for some small ϵ>0\epsilon>0. Then for dd large enough, we have τ≤Cd\tau\leq Cd, where CC is a universal constant depending only on γ1\gamma_{1}, γ2\gamma_{2}, and ϵ\epsilon.

First, choose ϵ\epsilon in Lemma 2.3 to be equal to γ2/4\gamma_{2}/4. Let dd be large enough so that ∣r02−1∣≤Clog⁡dd≤η\lvert r_{0}^{2}-1\rvert\leq C\frac{\log d}{\sqrt{d}}\leq\eta, where η=η(ϵ)\eta=\eta(\epsilon) is the required value in Lemma 2.3. By Lemma 2.2, we see that ∣rˉk2−1∣≤η\lvert\bar{r}_{k}^{2}-1\rvert\leq\eta for all k≤τk\leq\tau. This allows us to use Lemma 2.4 to observe that the recursive inequality sˉk+1≥(1+c/d)sˉk\bar{s}_{k+1}\geq(1+c/d)\bar{s}_{k} applies whenever k≤τk\leq\tau, where cc is a constant only depending on γ2\gamma_{2}. By the definition of τ\tau, we have

Recall that by the discussion at the end of Section 3, we have

Using Lemma 3.2, there is also a probability 1−1/d1-1/d event over which we have

We will show that (6.1) holds when conditioned on both of these events. First, for any t≤τt\leq\tau, we may write

and similarly we have (4), which we reproduce here:

Subtracting these two equations, and recalling that sˉ0=s0\bar{s}_{0}=s_{0}, we get

where the bound for the first term on the right hand side comes from (6.3)

We are now in a position to apply Lemma E.2 with xt=1d∣βˉ(y(t))−βˉ(yˉ(t))∣x_{t}=\frac{1}{d}\left\lvert\bar{\beta}(\textbf{y}^{(t)})-\bar{\beta}(\bar{\textbf{y}}^{(t)})\right\rvert, ρ=L/d\rho=L/d, and ξ=Clog⁡d/d\xi=C\log d/\sqrt{d}. Doing so, we get

Set ϵ=γ2/2\epsilon=\gamma_{2}/2. Combining the previous two lemmas, we have

where the last inequality holds for dd large enough. As such, we have τ2b≤τ≤Cd\tau_{2b}\leq\tau\leq Cd. ∎

Linear convergence in Phase 3 and proof of Theorem 1.2

To summarize, we have showed that τ1≤Cd⋅(log⁡d+log⁡∣r02−1∣)\tau_{1}\leq Cd\cdot(\log d+\log\lvert r_{0}^{2}-1\rvert) with probability at least 1−C/log⁡d1-C/\log d (Lemma 3.1), τ2a−τ1≤Cdlog⁡d\tau_{2a}-\tau_{1}\leq Cd\log d with probability at least pp (Theorem 4.1), and τ2b−τ2a≤Cd\tau_{2b}-\tau_{2a}\leq Cd with probability at least 1−1/d1-1/d (Lemma 6.1). Putting all of these together gives

with probability at least p−C/log⁡d−C/dp-C/\log d-C/d, which is larger than p/2p/2 for dd large enough. Unfortunately, this is not good enough for our purposes and we need to do a bit more work to bring down the error probability.

Let us condition on the 1−1/d1-1/d probability event for which (3.9) holds so that we have uniform control over {rk2}\{r_{k}^{2}\} over an appropriate timescale (more precisely, we use the coupling argument explained in the discussion after Lemma 3.5). Define A≔Cdlog⁡(d⋅C3)A\coloneqq Cd\log\left(d\cdot C_{3}\right), where C3C_{3} is the same constant used in (3.9). Then by the strong Markov property, for any k0>0k_{0}>0, we have

If we set t=log⁡(10)log⁡(1−p/2)t=\frac{\log(10)}{\log(1-p/2)}, we see that with probability at least 0.9−C/log⁡d−1/d0.9-C/\log d-1/d,

As mentioned earlier, we will take T=τ2bT=\tau_{2b}.

If we write Yk≔(1−1/2d)−kΨ(y(T+k))⋅1τ>T+kY_{k}\coloneqq(1-1/2d)^{-k}\Psi(\textbf{y}^{(T+k)})\cdot\mathbf{1}_{\tau>T+k}, then the above bound allows us to compute

and we see that {Yk}k\{Y_{k}\}_{k} is a supermartingale with respect to the filtration {FT+k}k\{\mathcal{F}_{T+k}\}_{k}.

By the supermartingale inequality, there is a probability 0.950.95 event over which

On the intersection of this event, and that on which τ=∞\tau=\infty, we have that

for all k≥0k\geq 0 as we wanted. If we total the measure of the excluded bad events, the final success probability is at least 0.8−C/log⁡d−1/d0.8-C/\log d-1/d as promised. ∎

We now discuss some straightforward extensions of this result. First, the lack of a high probability guarantee is a little unfortunate, and results from having to apply the supermartingale inequality. We are unsure whether this can be overcome theoretically, but we can easily modify the algorithm so that convergence holds with high probability. To do this, we use the “majority vote” procedure described in . The price we have to pay is an additional log⁡(1/δ)\log(1/\delta) factor on the number of iterations, where δ\delta is the total error probability we can tolerate.

Second, in the algorithm we presented, we required fresh samples to be used in every update step. Once we are in the linear convergence regime, however, samples can actually be reused so long as we choose uniformly from Ω(d)\Omega(d) of them (see ).

Conclusion and discussion

In this paper, we have analyzed the convergence of constant step-size stochastic gradient descent for the non-convex, non-smooth phase retrieval objective (1.1), for which we assume Gaussian sampling vectors, and use an arbitrary initialization. The main idea was to view the SGD sequence as a Markov chain on a summary state space, and then use the natural 1/d1/d step size scaling to argue that as dd tends to infinity, the process trajectory converges to something we understand. We believe that our proof framework and techniques will have applications beyond the vanilla phase retrieval model, and indeed inform the theory of non-convex optimization in general.

We have analyzed phase retrieval in the noiseless setting, as is customary in the literature. It is easy to see that the arguments still go through for an additive noise model, except that we may now need to either use batch updates, or reduce the step size further, in order to get convergence. It will also be interesting to see whether the results we have can be extended to the setting of non-Gaussian sampling vectors. Our state space argument relies on the rotational symmetry of the sampling vectors, and this clearly will not hold in the non-Gaussian setting. However, it may be possible for this to be overcome using approximations.

2 Extensions to other single index models

Phase retrieval is an example of a single index model with the link function f(t)=∣t∣f(t)=\lvert t\rvert. One can easily check that the state space argument generalizes to models with other link functions. Less clear is how to generalize the other arguments in the paper. Nonetheless, we expect this to be not too difficult, so long as we assume some natural regularity conditions on the link function. Finally, we conjecture that similar ideas can work for analyzing SGD for low-rank matrix sensing, since heuristically what makes everything work is the underlying low-dimensional structure in the problem.

3 Nonconvex optimization

In the introduction to this paper, we have already talked about the growing interest in a theoretical understanding of first-order methods applied to non-convex problems. Here, we reiterate that the main contribution of our paper should be seen as not simply providing a convergence guarantee for SGD, but also one that has close to optimal sample and computational complexity. We are able to achieve this through a careful analysis of the “essential dynamics” of the SGD process, as represented by the summary state space.

Standard proofs of convergence for first order methods applied to convex problems proceed by tracking one of the following three quantities: ∥xk−x∗∥2\lVert\textbf{x}_{k}-\textbf{x}^{*}\rVert^{2}, f(xk)−f(x∗)f(\textbf{x}_{k})-f(\textbf{x}^{*}), or ∥∇f(x∗)∥\lVert\nabla f(\textbf{x}^{*})\rVert. Under our framework, this can be seen as implicitly using a one-dimensional state space. In non-convex optimization, however, it makes sense to use a multi-dimensional state space, whereby we measure “progress” in terms of multiple quantities. The number of such quantities one needs to track for a given problem can then perhaps be used to define a notion of “complexity” for that problem.

Acknowledgements

Y.T. was partially supported by NSF CCF-1740855, and completed part of this manuscript while visiting the Simons Institute for the Theory of Computing. R.V. was partially supported by U.S. Air Force Grant FA9550-18-1-0031. Y.T. would like to thank Xiang Cheng, Jelena Diakonikolas, Michael Jordan, and Yian Ma for helpful discussions.

References

Appendix A State space calculations

Recall that we let x(0),x(1),x(2),…\textbf{x}^{(0)},\textbf{x}^{(1)},\textbf{x}^{(2)},\ldots denote the sequence obtained by iteratively performing the SGD update (1.3) with constant step size η=1d\eta=\frac{1}{d}.

The sequence x(0),x(1),x(2),…\textbf{x}^{(0)},\textbf{x}^{(1)},\textbf{x}^{(2)},\ldots is a Markov chain whose transition kernel has the random mapping representation

a∼Unif(dSd−1)\textbf{a}\sim\textnormal{Unif}(\sqrt{d}S^{d-1}), and A\mathcal{A} is the event that sign(⟨a,x⟩)≠sign(⟨a,x∗⟩)\textnormal{sign}(\langle\textbf{a},\textbf{x}\rangle)\neq\textnormal{sign}(\langle\textbf{a},\textbf{x}^{*}\rangle).

The fact that the sequence is a Markov chain is clear. After dropping indices, we may write the update step (1.3) as

Abusing notation slightly, let us define the state space updates

Our goal is to compute formulas for the distributions for α(x)\alpha(\textbf{x}) and β(x)\beta(\textbf{x}), in particular showing that they depend on x only through y. To this end, we simplify notation, denoting r=r(x)r=r(\textbf{x}), s=s(x)s=s(\textbf{x}), and θ=θ(x)\theta=\theta(\textbf{x}). Furthermore, define x⊥≔Px∗⊥x∥Px∗⊥x∥\textbf{x}^{\perp}\coloneqq\frac{\textbf{P}_{\textbf{x}^{*}}^{\perp}\textbf{x}}{\lVert\textbf{P}_{\textbf{x}^{*}}^{\perp}\textbf{x}\rVert}.

We can decompose x into its components parallel and perpendicular to x∗\textbf{x}^{*}, writing

Let a∼Unif(dSd−1)\textbf{a}\sim\textnormal{Unif}(\sqrt{d}S^{d-1}) be the random vector used to generate Δ(x)\Delta(\textbf{x}). We also have the orthogonal decomposition

where a1a_{1}, a2a_{2} are the marginals of a along x∗\textbf{x}^{*} and x⊥\textbf{x}^{\perp} respectively, and r is defined as the remainder in the decomposition above.

Combining these formulas, we immediately get

To compute the formula for α(x)\alpha(\textbf{x}), we first expand

Now use equations (A.2), (A.4) and (A.5) to write

For the second formula, we note that β(x)=⟨Δ(x),x∗⟩\beta(\textbf{x})=\langle\Delta(\textbf{x}),\textbf{x}^{*}\rangle, and write

Here, the second equality follows from a combination of (A.4), (A.3) and (A.5). ∎

The distribution of the event A=A(θ)\mathcal{A}=\mathcal{A}(\theta) depends only on the angle θ\theta between x and x∗\textbf{x}^{*}. Furthermore, using the formula (A.4), we have the following identities.

Appendix B Properties of subexponential random variables

Subexponential random variables are defined in terms of tail bounds, and can also be equivalently defined as elements of an Orlicz space with Orlicz norm

One may easily check that this is a norm, which allows for easy tail bounds for random variables that are sums of subexponential random variables. We will only state the propreties of subexponential variables needed in our paper, and refer the interested reader to the textbook .

Let D⊂Y\mathcal{D}\subset\mathcal{Y} be a compact domain. Then the subexponential norms of α(y)\alpha(\textbf{y}) and β(y)\beta(\textbf{y}) are uniformly bounded for y∈Y\textbf{y}\in\mathcal{Y}.

Suppose we have r(y)≤Rr(\textbf{y})\leq R for all y∈Y\textbf{y}\in\mathcal{Y}. Then

The first inequality is an application of the triangle inequality, the third follows from the following basic property for subexponential random variables: ∥XY∥ψ1≤∥X∥ψ2∥Y∥ψ2\lVert XY\rVert_{\psi_{1}}\leq\lVert X\rVert_{\psi_{2}}\lVert Y\rVert_{\psi_{2}}. Finally, it is easy to check that ∥u∥ψ2≲1\lVert u\rVert_{\psi_{2}}\lesssim 1. ∎

Appendix C Lemmas for drift

For each fixed rr, βˉ\bar{\beta} is an odd function with respect to ss, so it suffices to prove the statement for s>0s>0. Let us now compute its derivative with respect to ss. Differentiating (2.7) first with respect to θ\theta, we get

Next, note that since s=rcos⁡θs=r\cos\theta, we have

Using the chain rule and then simplifying, we thereby get

From this expression, we can see that for any fixed r>0r>0, the function s↦βˉ(r2,s)s\mapsto\bar{\beta}(r^{2},s) is concave downwards on [0,r][0,r], with βˉ(r2,0)=0\bar{\beta}(r^{2},0)=0, and ∂sβˉ(r2,s)\vlines=0>0\partial_{s}\bar{\beta}(r^{2},s)\vline_{s=0}>0. This implies that the graph of βˉ(r2,s)\bar{\beta}(r^{2},s) as a function of ss lies beneath the line passing through the origin with slope ∂sβˉ(r2,s)\vlines=0\partial_{s}\bar{\beta}(r^{2},s)\vline_{s=0}. When r≥12r\geq\frac{1}{2}, we have

For the lower bound, first note that concavity also implies that βˉ(r2,s)≥ss′βˉ(r2,s′)\bar{\beta}(r^{2},s)\geq\frac{s}{s^{\prime}}\bar{\beta}(r^{2},s^{\prime}) for any 0<s<s′<r0<s<s^{\prime}<r. Now, one may easily check that βˉ(1,1)=0\bar{\beta}(1,1)=0, so that βˉ(1,1−ϵ)>0\bar{\beta}(1,1-\epsilon)>0 for ϵ\epsilon small enough. By continuity, there is some η>0\eta>0 for which

Indeed, since ∂rβˉ(r2,s)=−cos⁡θ\partial_{r}\bar{\beta}(r^{2},s)=-\cos\theta, one may even provide a precise formula if one wishes. Set b‾≔βˉ(1,1−ϵ)2(1−ϵ)\underline{b}\coloneqq\frac{\bar{\beta}(1,1-\epsilon)}{2(1-\epsilon)}. This is the universal constant we want. ∎

Fix ϵ>0\epsilon>0 in the previous lemma, and by making ϵ\epsilon and η\eta smaller if necessary, assume that ∣s∣<r−ϵ/2\lvert s\rvert<r-\epsilon/2 for all (r2,s)∈D(r^{2},s)\in\mathcal{D}, where D≔{(r2,s)∈Y  ⁣: ∣s∣≤1−ϵ,∣r2−1∣≤η}\mathcal{D}\coloneqq\{(r^{2},s)\in\mathcal{Y}~{}\colon~{}\lvert s\rvert\leq 1-\epsilon,\lvert r^{2}-1\rvert\leq\eta\}. Then βˉ\bar{\beta} is Lipschitz continuous on D\mathcal{D} with Lipschitz constant bounded by a universal constant LL depending only on ϵ\epsilon.

The first term is trivially bounded by ∥y−y′∥\lVert\textbf{y}-\textbf{y}^{\prime}\rVert. Next, observe that θ=arccos⁡(s/r2)\theta=\arccos(s/\sqrt{r^{2}}), which is jointly differentiable in r2r^{2} and ss, and so is Lipschitz continuous with respect to these coordinates on a compact set bounded away from r=sr=s. By assumption, D\mathcal{D} is such a compact set. ∎

Appendix D Facts about Kolmogorov distance and stochastic dominance

For any real-valued random variables XX and YY, we have

as we wanted. The second identity can be obtained similarly. ∎

Let XX and YY be real-valued random variables. Then for any 0<δ<10<\delta<1, XX stochastically dominates YY up to error δ\delta if and only if their CDFs satisfy FX≤FY+δF_{X}\leq F_{Y}+\delta.

Let UU be uniformly distributed on $,andonthesameprobabilityspace,define, and on the same probability space, defineU^{\prime}(\omega)=U(\omega)+\delta\mod 1.Then. ThenU^{\prime}isalsouniformlydistributedonis also uniformly distributed on.Next,iteasytocheckthat. Next, it easy to check thatq_{X}(U^{\prime})\sim F_{X}andandq_{Y}(U)\sim F_{Y}(thisisastandardconstructioninprobabilitytheory).Withprobability(this is a standard construction in probability theory). With probability1-\delta,wehave, we have\delta\leq U^{\prime}=U+\delta\leq 1$. Conditioning on the event in which this occurs, we have

Let XX, YY, and ZZ be random variables, δ,δ′>0\delta,\delta^{\prime}>0 such that X⪰δYX\succeq_{\delta}Y and Y⪰δ′ZY\succeq_{\delta^{\prime}}Z. Then X⪰δ+δ′ZX\succeq_{\delta+\delta^{\prime}}Z.

This follows from the characterization in the previous lemma. ∎

Let XX and YY be real-valued random variables such that their CDFs satisfy ∥FX−FY∥∞≤δ\lVert F_{X}-F_{Y}\rVert_{\infty}\leq\delta for some 0<δ<10<\delta<1. Then XX stochastically dominates YY up to error δ\delta.

Fix σ2,ϵ>0\sigma^{2},\epsilon>0. For any b>0b>0, define the random variable Xb≔ρϵ[b+σg]2X_{b}\coloneqq\rho_{\epsilon}[b+\sigma g]^{2}. Then whenever b′>bb^{\prime}>b, we have Xb′⪰XbX_{b^{\prime}}\succeq X_{b}.

Fix σ02\sigma_{0}^{2}, set σ2=Bσ02d2\sigma^{2}=\frac{B\sigma_{0}^{2}}{d^{2}}. Let ϵ(s)\epsilon(s) and b(s)b(s) be defined as in (4.6) and (4.9) respectively . Consider the collection of random variables Ys≔ρϵ(s)[b(s)+σg]2Y_{s}\coloneqq\rho_{\epsilon(s)}[b(s)+\sigma g]^{2} for −1/2<s<1/2-1/2<s<1/2. For dd large enough, whenever 1/2>s′>s>01/2>s^{\prime}>s>0, we have Ys′⪰YsY_{s^{\prime}}\succeq Y_{s}.

Throughout, we let gg denote a standard normal random variable. We start by proving the first statement. Let FbF_{b} denote the CDFs of XbX_{b} and Xb′X_{b^{\prime}} respectively. For any given a>0a>0, we wish to show that Fb(a)≥Fb′(a)F_{b}(a)\geq F_{b^{\prime}}(a), and it suffices to show that ddbFb(a)≤0\frac{d}{db}F_{b}(a)\leq 0 for b>0b>0. In order to do this, we write out FsF_{s} in terms of the Gaussian CDF. We have

Since ϕ\phi is even, and decreases away from , and ∣−a−ϵ−b∣>∣a+ϵ−b∣\lvert-\sqrt{a}-\epsilon-b\rvert>\lvert\sqrt{a}+\epsilon-b\rvert for b>0b>0, the quantity on the right hand side is negative as we wanted.

The second statement is proved similarly. We let FsF_{s} denote the CDF of YsY_{s} and will show that ddsFs(a)≤0\frac{d}{ds}F_{s}(a)\leq 0. As before we compute

For 0<s<log⁡dB0<s<\sqrt{\frac{\log d}{B}}, we have ϵ′(s)=0\epsilon^{\prime}(s)=0, and it is clear that this quantity is nonpositive. For s>log⁡dBs>\sqrt{\frac{\log d}{B}}, first observe that for any x,yx,y, we have e−(x−y)2/e−(x+y)2=e4xye^{-(x-y)^{2}}/e^{-(x+y)^{2}}=e^{4xy}. As such, we compute

For dd large enough, this ratio is greater than

which converges to 11 as dd tends to infinity. This concludes the proof of the claim. ∎

Appendix E Combinatorial lemmas

Let x1,x2,…x_{1},x_{2},\ldots be a sequence of real numbers. Denote the partial sums by st=∑i=1txis_{t}=\sum_{i=1}^{t}x_{i}. For any 0<ρ<10<\rho<1, define the sequence w1,w2,…w_{1},w_{2},\ldots via the recursive formula

If there is some positive integer MM and some C>0C>0 such that ∣st∣≤C\lvert s_{t}\rvert\leq C for all t≤Mt\leq M, then we also have ∣wt∣≤2C+∣w0∣\lvert w_{t}\rvert\leq 2C+\lvert w_{0}\rvert for all t≤Mt\leq M.

We first prove by induction that the following representation for wtw_{t} holds:

First assume that the formula holds for some tt. Then starting from the definition, we use the inductive hypothesis to write

Applying the triangle inequality to (E.1), we get

Let x1,x2,…x_{1},x_{2},\ldots be a sequence of non-negative real numbers. Denote the partial sums by st=∑i=1txis_{t}=\sum_{i=1}^{t}x_{i}. Let ρ>0\rho>0 and ξ>0\xi>0 be such that we have the recursive inequality

Appendix F Concentration inequalities

This essentially follows from Theorem 1 in . For completeness and clarity, however, we give a direct proof of this here using the same technique.

First fix ϵ>0\epsilon>0. Let λ≤12K\lambda\leq\frac{1}{2K}, and for each positive integer tt, we write

so that the sequence Lt(λ)L_{t}(\lambda) forms a supermartingale, and remains so if we extend this to time 0 by setting L0(λ)=1L_{0}(\lambda)=1. We may now apply the supermartingale inequality. We have

It remains to choose λ\lambda appropriately in order to get the tail bound we want. If ϵ≤MK\epsilon\leq MK, then we set λ=ϵ/2K2M\lambda=\epsilon/2K^{2}M, observing that for this choice of λ\lambda,

and Lt(λ)L_{t}(\lambda) is indeed a supermartingale, and (F) holds. Plugging our choice of λ\lambda into the left hand and right hand sides, this yields the bound

On the other hand, if ϵ>MK\epsilon>MK, then we pick λ=1/2K\lambda=1/2K. Once again plugging this into (F), we get

Putting these two bounds together gives the upper tail in (3.4), and considering the negative of the sequence gives the lower tail. ∎

Let X1,X2,…,XMX_{1},X_{2},\ldots,X_{M} be a sequence of Bernoulli random variables adapted to a filtration {Gt}\{\mathcal{G}_{t}\}, and let θ\theta be such that for 1≤t≤M1\leq t\leq M, we have

Denote St≔∑i=1tXiS_{t}\coloneqq\sum_{i=1}^{t}X_{i}, observe that for any λ>0\lambda>0,

The rest of the proof is exactly the same as that of the regular Chernoff’s inequality (see ). ∎

Let X′X^{\prime} be an independent copy of XX. By Jensen’s inequality, followed by applying the contraction inequality pointwise, we have

Here, the inequality follows from Cauchy-Schwarz.

For the second claim, assume b>0b>0 and write

The first equality follows from the oddness of ρ\rho and the symmetry of XX, while the second follows from contraction. If b<0b<0, the statement may be proved similarly. ∎