On the Power of Differentiable Learning versus PAC and SQ Learning

Emmanuel Abbe, Pritish Kamath, Eran Malach, Colin Sandon, Nathan Srebro

Introduction

A leading paradigm that has become the predominant approach to learning is that of differentiable learning, namely using a parametric function class fw(x)f_{{\bm{w}}}(x), and learning by performing mini-batch Stochastic Gradient Descent (bSGD{\mathsf{bSGD}}) updates (using gradients of the loss on a mini-batch of bb independent samples per iteration) or full-batch Gradient Descent (fbGD{\mathsf{fbGD}}) updates (full gradient descent on the empirical loss, using the same mm samples in all iterations). Feed-forward neural networks are a particularly popular choice for the parametric model fwf_{{\bm{w}}}. One approach to understanding differentiable learning is to think of it as a method for minimizing the empirical error of fwf_{{\bm{w}}} with respect to w{\bm{w}}, i.e. as an empirical risk minimization (ERM\mathsf{ERM}). ERM\mathsf{ERM} is indeed well understood, and is in a sense a universal learning rule, in that any hypothesis class that is PAC{\mathsf{PAC}} learnable is also learnable using ERM\mathsf{ERM}. Furthermore, since poly-sized feed-forward neural networks can represent any poly-time computable function, we can conclude that ERM\mathsf{ERM} on neural networks can efficiently learn any tractable problem. But this view of differentiable learning ignores two things.

Firstly, many ERM\mathsf{ERM} problems, including ERM\mathsf{ERM} on any non-trivial neural network, are highly non-convex and (stochastic) gradient descent might not find the global minimizer. As we just discussed, pretending that bSGD{\mathsf{bSGD}} or fbGD{\mathsf{fbGD}} do find the global minimizer would mean we can learn all poly-time computable functions, which is known to be impossibleSubject to mild cryptographic assumptions, such as the existence of one-way functions, i.e. the existence of cryptography itself, and where we are referring to bSGD/fbGD{\mathsf{bSGD}}/{\mathsf{fbGD}} running for polynomially many steps. (e.g. Kearns and Valiant, 1994; Klivans and Sherstov, 2009). In fact, even minimizing the empirical risk on a neural net with two hidden units, and even if we assume data is exactly labeled by such a network, is already NP-hard, and cannot be done efficiently with bSGD{\mathsf{bSGD}}/fbGD{\mathsf{fbGD}} (Blum and Rivest, 1992). We already see that bSGD{\mathsf{bSGD}} and fbGD{\mathsf{fbGD}} are not the same as ERM\mathsf{ERM}, and asking “what can be learned by bSGD{\mathsf{bSGD}}/fbGD{\mathsf{fbGD}}” is quite different from asking “what can be learned by ERM\mathsf{ERM}”.

Furthermore, bSGD{\mathsf{bSGD}} and fbGD{\mathsf{fbGD}} might also be more powerful than ERM\mathsf{ERM}. Consider using a highly overparametrized function class fwf_{{\bm{w}}}, with many more parameters than the number of training examples, as is often the case in modern deep learning. In such a situation there would typically be many empirical risk minimizers (zero training error solutions), and most of them would generalize horribly, and so the ERM\mathsf{ERM} principal on its own is not sufficient for learning (Neyshabur et al., 2015). Yet we now understand how fbGD{\mathsf{fbGD}} and bSGD{\mathsf{bSGD}} incorporate intricate implicit bias that leads to particular empirical risk minimizers which might well ensure learning (Soudry et al., 2018; Nacson et al., 2019).

All this justifies studying differentiable learning as a different and distinct paradigm from empirical risk minimization. The question we ask is therefore:

What can be learned by performing mini-batch stochastic gradient descent, or full-batch gradient descent, on some parametric function class fwf_{{\bm{w}}}, and in particular on a feed-forward neural network?

Answering this question is important not only for understanding the limits of what we could possibly expect from differentiable learning, but even more so in guiding us as to how we should study differentiable learning, and ask questions, e.g., about its power relative to kernel methods (e.g. Yehudai and Shamir, 2019; Allen-Zhu and Li, 2019, 2020; Li et al., 2020; Daniely and Malach, 2020; Ghorbani et al., 2019, 2020; Malach et al., 2021), or the theoretical benefits of different architectural innovations (e.g. Malach and Shalev-Shwartz, 2020).

A significant, and perhaps surprising, advance toward answering this question was recently presented by Abbe and Sandon (2020), who showed that bSGD{\mathsf{bSGD}} with a single example per iteration (i.e. a minibatch of size 11) can simulate any poly-time learning algorithm, and hence is as powerful as PAC{\mathsf{PAC}} learning. On the other hand, they showed training using population gradients (infinite batch size, or even very large batch sizes) and low-precision (polynomial accuracy, i.e. a logarithmic number of bits of precision), is no more powerful than learning with Statistical Queries (SQ{\mathsf{SQ}}), which is known to be strictly less powerful than PAC{\mathsf{PAC}} learning (Kearns, 1998; Blum et al., 2003). This seems to suggest non-stochastic fbGD{\mathsf{fbGD}}, or even bSGD{\mathsf{bSGD}} with large batch sizes, is not universal in the same way as single-example bSGD{\mathsf{bSGD}}. As we will see below, it turns out this negative result depends crucially on the allowed precision, and its relationship to the batch, or sample, size.

In this paper, we take a more refined view of this dichotomy, and consider learning with bSGD{\mathsf{bSGD}} with larger mini-batch sizes b>1b>1, as is more typically done in practice, as well as with fbGD{\mathsf{fbGD}}. We ask whether the ability to simulate PAC{\mathsf{PAC}} learning is indeed preserved also with larger mini-batch sizes, or even full batch GD? Does this universality rest on using single examples, or perhaps very small mini-batches, or can we simulate any learning algorithm also without such extreme stochasticity? We discover that this depends on the relationship of the batch size bb and the precision ρ\rho of the gradient calculations. That is, to understand the power of differentiable learning, we need to also explicitly consider the numeric precision used in the gradient calculations, where ρ\rho is an arbitrary additive error we allow.

We first show that regardless of the mini-batch size, bSGD{\mathsf{bSGD}} is always able to simulate any SQ{\mathsf{SQ}} method, and so bSGD{\mathsf{bSGD}} is at least as powerful as SQ{\mathsf{SQ}} learning. When the mini-batch size bb is large relative to the precision, namely b=ω(log⁡(n)/ρ2)b=\omega(\log(n)/\rho^{2}), where nn is the input dimension and we assume the model size and number of iterations are polynomial in nn, bSGD{\mathsf{bSGD}} is not any more powerful than SQ{\mathsf{SQ}}. But when b<1/(8ρ)b<1/(8\rho), or in other words with fine enough precision ρ<1/(8b)\rho<1/(8b), bSGD{\mathsf{bSGD}} can again simulate any sample-based learning method, and is as powerful as PAC{\mathsf{PAC}} learning (the number of SGD iterations and size of the model used depend polynomially on the sample complexity of the method being simulated, but the mini-batch size bb and precision ρ\rho do not, and only need to satisfy bρ<1/8b\rho<1/8—see formal results in Section 3). We show a similar result for fbGD{\mathsf{fbGD}}, with a dependence on the sample size mm: with low precision (large ρ\rho) relative to the sample size mm, fbGD{\mathsf{fbGD}} is no more powerful than SQ{\mathsf{SQ}}. But with fine enough precision relative to the sample size, namely when ρ<1/(8m)\rho<1/(8m), fbGD{\mathsf{fbGD}} can again simulate any sample-based learning method based on mm samples (see formal results in Section 6).

We see then, that with fine enough precision, differentiable learning with any mini-batch size, or even using full-batch gradients (no stochasticity in the updates), is as powerful as any sample-based learning method. The required precision does depend on the mini-batch or sample size, but only linearly. That is, the number of bits of precision required is only logarithmic in the mini-batch or sample sizes. And with a linear (or even super-logarithmic) number of bits of precision (i.e. ρ=2−n\rho=2^{-n}), bSGD{\mathsf{bSGD}} and fbGD{\mathsf{fbGD}} can simulate arbitrary sample-based methods with any polynomial mini-batch size; this is also the case if we ignore issue of precision and assume exact computation (corresponding to ρ=0\rho=0).

On the other hand, with low precision (high ρ\rho, i.e. only a few bits of precision, which is frequently the case when training deep networks), the mini-batch size bb plays an important role, and simulating arbitrary sample based methods is provably not possible using fbGD{\mathsf{fbGD}}, or with bSGD{\mathsf{bSGD}} with a mini-batch size that is too large, namely b=ω(log⁡(n)/ρ2)b=\omega(\log(n)/\rho^{2}). Overall, except for an intermediate regime between 1/ρ1/\rho and log⁡(n)/ρ2\log(n)/\rho^{2}, we can precisely capture the power of bSGD{\mathsf{bSGD}}.

Computationally Bounded and Unbounded Learning.

Another difference versus the work of Abbe and Sandon is that we discuss both computationally tractable and intractable learning. We show that poly-time PAC{\mathsf{PAC}} and SQ{\mathsf{SQ}} are related, in the sense described above, to bSGD{\mathsf{bSGD}} and fbGD{\mathsf{fbGD}} on a poly-sized neural network, whereas computationally unbounded PAC{\mathsf{PAC}} and SQ{\mathsf{SQ}} (i.e. limited only by the number of samples or number of statistical queries, but not runtime) are similarly related to bSGD{\mathsf{bSGD}} and fbGD{\mathsf{fbGD}} on an arbitrary differentiable model fwf_{{\bm{w}}}. In fact, to simulate PAC{\mathsf{PAC}} and SQ{\mathsf{SQ}}, we first construct an arbitrary fwf_{{\bm{w}}}, and then observe that if the PAC{\mathsf{PAC}} or SQ{\mathsf{SQ}} method is poly-time computable, then the computations within it can be expressed as poly-size circuits, which can in turn be simulated as poly-size sub-networks, allowing us to implement fwf_{{\bm{w}}} as a neural net.

Answering 𝗦𝗤𝗦𝗤{\mathsf{SQ}}s using Samples.

Our analysis relies on introduction of a variant of SQ{\mathsf{SQ}} learning which we refer to as mini-batch Statistical Queries (bSQ{{\mathsf{bSQ}}}, and we similarly introduce a full-batch variant, fbSQ{{\mathsf{fbSQ}}}). In this variant, which is related to the Honest\mathsf{Honest}-SQ{\mathsf{SQ}} model (Yang, 2001, 2005), statistical queries are answered using a mini-batch of samples drawn from the source distribution, up to some precision. We first show that bSQ{{\mathsf{bSQ}}} methods can always be simulated by bSGD{\mathsf{bSGD}}, by constructing a differentiable model where at each step the derivatives with respect to some of the parameters contain the answers to the statistical queries. We then relate bSQ{{\mathsf{bSQ}}} to SQ{\mathsf{SQ}} and PAC{\mathsf{PAC}}, based on the relationship between the mini-batch size and precision. In order to simulate PAC{\mathsf{PAC}} using bSQ{{\mathsf{bSQ}}}, we develop a novel “sample extraction” that uses mini-batch statistical queries (on independently sampled mini-batches) to extract a single sample drawn from the source distribution. This procedure might be of independent interest, perhaps also in studying privacy, where such an extraction is not desirable. Our study of the relationship of bSQ{{\mathsf{bSQ}}} to SQ{\mathsf{SQ}} and PAC{\mathsf{PAC}}, summarized in Section 4, also sheds light on how well the SQ{\mathsf{SQ}} framework captures learning by answering queries using empirical averages on a sample, which is arguably one of the main motivations for the inverse-polynomial tolerance parameter in the SQ framework.

Learning Paradigms

where gtg_{t} is a ρ\rho-approximate rounding of the mini-batch (empirical) clipped gradient

Precision, Rounding and Clipping.

Learning with 𝗯𝗦𝗚𝗗𝗯𝗦𝗚𝗗{\mathsf{bSGD}}.

Full-batch Gradient Descent (𝗳𝗯𝗚𝗗𝗳𝗯𝗚𝗗{\mathsf{fbGD}}).

𝗣𝗔𝗖𝗣𝗔𝗖{\mathsf{PAC}} and 𝗦𝗤𝗦𝗤{\mathsf{SQ}} Learning.

Relating Classes of Methods.

That is, C′⪯δC\mathcal{C}^{\prime}\preceq_{\delta}\mathcal{C} means that C′\mathcal{C}^{\prime} is at least as powerful as C\mathcal{C}. Observe that for all classes of methods C1,C2,C3\mathcal{C}_{1},\mathcal{C}_{2},\mathcal{C}_{3}, if C1⪯δ1C2\mathcal{C}_{1}\preceq_{\delta_{1}}\mathcal{C}_{2} and C2⪯δ2C3\mathcal{C}_{2}\preceq_{\delta_{2}}\mathcal{C}_{3} then C1⪯δ1+δ2C3\mathcal{C}_{1}\preceq_{\delta_{1}+\delta_{2}}\mathcal{C}_{3}.

Main Results : 𝗯𝗦𝗚𝗗𝗯𝗦𝗚𝗗{\mathsf{bSGD}} versus 𝗣𝗔𝗖𝗣𝗔𝗖{\mathsf{PAC}} and 𝗦𝗤𝗦𝗤{\mathsf{SQ}}

Our main result, given below as a four-part Theorem, establishes the power of bSGD{\mathsf{bSGD}} learning relative to PAC{\mathsf{PAC}} (i.e. arbitrary sample based) and SQ{\mathsf{SQ}} learning. As previously discussed, the exact relation depends on the mini-batch size bb and gradient precision ρ\rho. First, we show that for any mini-batch size bb, with fine enough precision ρ\rho, bSGD{\mathsf{bSGD}} can simulate PAC{\mathsf{PAC}}.

For all bb and ρ<1/(8b)\rho<1/(8b), and for all m,r,δm,r,\delta, it holds that

To establish equivalence (when Theorem 1a holds), we also note that PAC{\mathsf{PAC}} is always at least as powerful as bSGD{\mathsf{bSGD}} (since bSGD{\mathsf{bSGD}} can be implemented using samples):

For all bb, ρ\rho and T,p,rT,p,r, it holds that

Furthermore, for all poly-time computable activations σ\sigma, it holds that

On the other hand, if the mini-batch size is large relative to the precision, bSGD{\mathsf{bSGD}} cannot go beyond SQ{\mathsf{SQ}}:

There exists a constant CC such that for all δ>0\delta>0, for all TT, ρ\rho, bb, pp, rr, such that bρ2>Clog⁡(Tp/δ)b\rho^{2}>C\log(Tp/\delta), it holds that

Furthermore, for all poly-time computable activations σ\sigma, it holds that

To complete the picture, we also show that regardless of the mini-batch size, i.e. even when bSGD{\mathsf{bSGD}} cannot simulate PAC{\mathsf{PAC}}, bSGD{\mathsf{bSGD}} can always, at the very least, simulate any SQ{\mathsf{SQ}} method. This also establishes equivalence to SQ{\mathsf{SQ}} when Theorem 1c holds:

There exists a constant CC such that for all δ>0\delta>0, for all bb and all k,τ,rk,\tau,r, it holds that

In the above Theorems, the reductions hold with parameters on the left-hand-side (the parameters of the model being reduced to) that are polynomially related to the dimension and the parameters on the right-hand-side. But the mini-batch size bb and gradient precision ρ\rho play an important role. In Theorem 1a, we may choose bb and ρ\rho as we wish, as long as they satisfy bρ<1/8b\rho<1/8—they do not need to be chosen based on the parameters of the PAC{\mathsf{PAC}} method, and we can always simulate PAC{\mathsf{PAC}} with any bb and ρ\rho satisfying bρ<1/8b\rho<1/8. Similarly, in Theorem 1d, we may chose b≥1b\geq 1 arbitrarily, and can always simulate SQ{\mathsf{SQ}}, although ρ\rho does have to be chosen according to τ\tau. The reverse reduction of Theorem 1c, establishing the limit of when bSGD{\mathsf{bSGD}} cannot go beyond SQ{\mathsf{SQ}}, is valid when bρ2=ω(log⁡n)b\rho^{2}=\omega(\log n), if the size of the model pp and number of SGD iterations TT are restricted to be polynomial in nn.

Focusing on the mini-batch size bb, and allowing all other parameters to be chosen to be polynomially related, Theorems 1a, 1b, 1c and 1d can be informally summarized as:

where the left relationship is tight if b<1/(8ρ)b<1/(8\rho) and the right relationship is tight if b>ω((log⁡n)/ρ2)b>\omega((\log n)/\rho^{2}). Equivalently, focusing on the precision ρ\rho and how it depends on the mini-batch size bb, and allowing all other parameters to be polynomially related, Theorems 1a, 1b, 1c and 1d can be informally summarized as:

For any (poly bounded, possible constant) dependence b(n,ρ)b(n,\rho), and for the activation function σ\sigma in Figure 1, it holds that

Moreover, if ∀n,ρ b(n,ρ)<1/(8ρ)\forall_{n,\rho}\ b(n,\rho)<1/(8\rho), then inclusions (2)(2) and (4)(4) are tight, and if b(n,ρ)≥ω(log⁡n)/ρ2b(n,\rho)\geq\omega(\log n)/\rho^{2}, then inclusions (1)(1) and (3)(3) are tight.

For any (poly bounded, possibly constant) b(n),ρ(n)b(n),\rho(n), and σ\sigma from Figure 1:

In Corollaries 1 and 2, for the sake of simplicity, we focused on realizable learning problems, where the minimal loss inf⁡fLDn(f)=0\inf_{f}\mathcal{L}_{\mathcal{D}_{n}}(f)=0 for each Dn∈Pn\mathcal{D}_{n}\in\mathcal{P}_{n}. However, we note that Theorems 1a, 1b, 1c and 1d are more general, as they preserve the performance of learning methods (up to an additive δ\delta) on all source distributions D\mathcal{D}. So, a result similar to Corollary 1 could be stated for other forms of learning, such as agnostic learning, weak learning etc.

Activation Functions and Fixed Weights.

The neural net simulations in Theorems 1a and 1d use a specific “stage-wise ramp” piecewise linear activation function with five linear pieces, depicted in Figure 1. A convenient property of this activation function is that it is has a central flat piece, making it easier for us to deal with weight drift due to rounding errors during training, and in particular drift of weights we would rather not change at all. Since any piecewise linear activation function can be simulated with ReLU activation, we could instead use a more familiar ReLU activation. However, the simulation “gadget” would involve weights that we would need fixed during training. That is, if we allow neural networks where some of the weights are fixed while others are trainable, we could use ReLU activation to simulate sample-based methods in Theorem 1a and SQ methods in Theorem 1d. As stated, we restrict ourselves only to neural nets where all edges have trainable weights, for which it is easier to prove the theorems with the specific activation function of Figure 1.

The Mini-Batch Statistical Query Model

En route to proving Theorems 1a, 1b, 1c and 1d, we introduce the model of mini-batch Statistical Queries (bSQ{{\mathsf{bSQ}}}). In this model, similar to the standard Statistical Query (SQ{\mathsf{SQ}}) learning model, learning is performed through statistical queries. But in bSQ{{\mathsf{bSQ}}}, these queries are answered based on an empirical average over a mini-batch of bb i.i.d. samples from the source distribution. That is, each query Φt:X×Y→p\Phi_{t}:\mathcal{X}\times\mathcal{Y}\to^{p} is answered with a response vtv_{t} s,t,

Note that we allow pp-dimensional vector “queries”, that is, pp concurrent scalar queries are answered based on the same mini-batch StS_{t}, drawn independently for each vector query.

The first step of our simulation of PAC{\mathsf{PAC}} and SQ{\mathsf{SQ}} with bSGD{\mathsf{bSGD}} is to simulate (a variant of) bSQ{{\mathsf{bSQ}}} using bSGD{\mathsf{bSGD}}. But beyond its use as an intermediate model in studying differentiable learning, bSQ{{\mathsf{bSQ}}} can also be though of as a realistic way of answering statistical queries. In fact, one of the main justifications for allowing errors in the SQ{\mathsf{SQ}} model, and the demand that the error tolerance τ\tau be polynomial, is that it is possible to answer statistical queries about the population with tolerance τ\tau by calculating empirical averages on samples of size O(1/τ2)O(1/\tau^{2}). In the bSQ{{\mathsf{bSQ}}} model we make this explicit, and indeed answer the queries using such samples. We do also allow additional arbitrary error beyond the sampling error, which we might think of as “precision”. The bSQ{{\mathsf{bSQ}}} model can thus be thought of as decomposing the SQ{\mathsf{SQ}} tolerance to a sampling error O(1/b)O(1/\sqrt{b}) and an additional arbitrary error τ\tau. If the arbitrary error τ\tau indeed captures “precision”, it is reasonable to take it to be exponentially small (corresponding to polynomially many bits of precision), while the sampling error would still be polynomial in a poly-time poly-sample method. Studying the bSQ{{\mathsf{bSQ}}} model can reveal to us how well the standard SQ{\mathsf{SQ}} model captures what can be done when most of the error in answering statistical queries is due to the sampling error.

Our bSQ{{\mathsf{bSQ}}} model is similar to the honest-SQ model studied by Yang (2001, 2005), who also asked whether answering queries based on empirical averages changes the power of the model. But the two models have some significant differences, which lead to different conclusions, namely: Honest\mathsf{Honest}-SQ{\mathsf{SQ}} does not allow for an additional arbitrary error (i.e. it uses τ=0\tau=0 in our notation), but an independent mini-batch is used for each single-bit query Φt:X×Y→{0,1}\Phi_{t}:\mathcal{X}\times\mathcal{Y}\to\{0,1\}, whereas bSQ{{\mathsf{bSQ}}} allows for pp concurrent real-valued scalar queries on the same mini-batch. Yang showed that, with a single bit query per mini-batch, and even if τ=0\tau=0, it is not possible to simulate arbitrary sample-based methods, and honest-SQ is strictly weaker than PAC{\mathsf{PAC}}. But we show that once multiple bits can be queried concurrentlyWe do so with polynomially many binary-valued queries, i.e. Φt(x)∈{0,1}p\Phi_{t}(x)\in\left\{0,1\right\}^{p}, and pp polynomial. It is also possible to encode this into a single real-valued query with polynomially many bits of precision. Once we use real-valued queries, if we do not limit the precision at all, i.e. τ=0\tau=0, and do not worry about processing time, its easy to extract the entire minibatch StS_{t} using exponentially many bits of precision. Theorem 2a shows that polynomially many bits are sufficient for extracting a sample and simulating PAC{\mathsf{PAC}}. the situation is quite different.

In fact, we show that when the arbitrary error τ\tau is small relative to the sample size bb (and thus the sampling error), bSQ{{\mathsf{bSQ}}} can actually go well beyond SQ{\mathsf{SQ}} learning, and can in fact simulate any sample-based method. That is, SQ{\mathsf{SQ}} does not capture learning using statistical queries answered (to within reasonable precision) using empirical averages:

(PAC{\mathsf{PAC}} to bSQ{{\mathsf{bSQ}}}) For all δ>0\delta>0, for all bb, and τ<1/(2b)\tau<1/(2b), and for all m,rm,r, it holds for k′=10m(n+1)/δk^{\prime}=10m(n+1)/\delta, p′=n+1p^{\prime}=n+1, r′=r+klog⁡2br^{\prime}=r+k\log_{2}b that

The main ingredient is a bSQ{{\mathsf{bSQ}}} method Sample-Extract (Algorithm 1) that extracts a sample (x,y)∼D(x,y)\sim\mathcal{D} by performing mini-batch statistical queries over independently sampled mini-batches. For ease of notation, let Z:=X×Y\mathcal{Z}:=\mathcal{X}\times\mathcal{Y} identifying it with {0,1}n+1\left\{0,1\right\}^{n+1} and denote z∈Zz\in\mathcal{Z} as (z1,…,zn+1)=(y,x1,…,xn)(z_{1},\ldots,z_{n+1})=(y,x_{1},\ldots,x_{n}). Sample-Extract operates by sampling the bits of zz one by one, drawing z^i\widehat{z}_{i} from the conditional distribution {zi∣z1,…,i−1}D\left\{z_{i}\mid z_{1,\ldots,i-1}\right\}_{\mathcal{D}}.

To complement the Theorem, we also note that sample-based learning is always at least as powerful as bSQ{{\mathsf{bSQ}}}, since bSQ{{\mathsf{bSQ}}} is specified based on a sample of size kbkb (see Section A.1 for a complete proof):

(bSQ{{\mathsf{bSQ}}} to PAC{\mathsf{PAC}}) For all bb, τ\tau and k,p,rk,p,r, it holds that

On the other hand, when the the mini-batch size bb is large relative to the precision (i.e. the arbitrary error τ\tau is large relative to the sampling error 1/b1/\sqrt{b}), bSQ{{\mathsf{bSQ}}} is no more powerful than standard SQ{\mathsf{SQ}}:

(bSQ{{\mathsf{bSQ}}} to SQ{\mathsf{SQ}}) There exists a constant C≥0C\geq 0 such that for all δ>0\delta>0, for all kk, τ\tau, bb, pp, rr, such that bτ2>Clog⁡(kp/δ)b\tau^{2}>C\log(kp/\delta), it holds that

When b≫1/τ2b\gg 1/\tau^{2}, the differences between the empirical and population averages become (with high probability) much smaller than the tolerance τ\tau, the population statistical query answers are valid responses to queries on the mini-batch, and we can thus simulate bSQ{{\mathsf{bSQ}}} using SQ{\mathsf{SQ}}. We do need to make sure this holds uniformly for the pp parallel scalar queries, and across all kk rounds—see Section A.3 for a complete proof. ∎

Finally, we show that with any mini-batch size, and enough rounds of querying, we can always simulate any SQ{\mathsf{SQ}} method using bSQ{{\mathsf{bSQ}}}:

(SQ{\mathsf{SQ}} to bSQ{{\mathsf{bSQ}}}) There exists a constant C≥0C\geq 0 such that for all δ>0\delta>0, for all bb and all k,τ,rk,\tau,r, it holds that

To obtain an answer to a statistical query on the population, even if the sample-size bb per query is small, we can average the responses for the same query over multiple mini-batches (i.e. over multiple rounds). This allows us to reduce the sampling error arbitrarily, and leaves us with only the arbitrary error τ′\tau^{\prime} (the arbitrary errors also get averaged, and since each element in the average is no larger than τ′\tau^{\prime}, the magnitude of this average is also no large than τ′\tau^{\prime}). See full proofs in Section A.3. ∎

Simulating Mini-Batch Statistical Queries with Differentiable Learning

Instead of working with, and simulating, any bSQ{{\mathsf{bSQ}}} method, we consider only alternating methods, denoted bSQ0/1{{\mathsf{bSQ^{0/1}}}}, where in each round, only one of the two possible labels is involved in the query. Formally, we say that a (mini-batch) statistical query Φ:X×Y→p\Phi:\mathcal{X}\times\mathcal{Y}\to^{p} is a y‾\overline{y}-query for y‾∈Y\overline{y}\in\mathcal{Y} if Φ(x,y)=0\Phi(x,y)=0 for all y≠y‾y\neq\overline{y}, or equivalently Φ(x,y)=\mathds1{y=y‾}⋅ΦX(x)\Phi(x,y)=\mathds{1}_{\left\{y=\overline{y}\right\}}\cdot\Phi_{\mathcal{X}}(x) for some ΦX:X→p\Phi_{\mathcal{X}}:\mathcal{X}\to^{p}. A bSQ0/1{{\mathsf{bSQ^{0/1}}}} (analogously, bSQTM0/1{{\mathsf{bSQ_{TM}^{0/1}}}}) method is a bSQ{{\mathsf{bSQ}}} (analogously, bSQTM{{\mathsf{bSQ_{TM}}}}) method such that for all odd rounds tt, Φt\Phi_{t} is a 11-query, and at all even rounds tt, Φt\Phi_{t} is a -query. As minor extensions of Theorems 2a and 2d (simulation PAC{\mathsf{PAC}} and SQ{\mathsf{SQ}} using bSQ{{\mathsf{bSQ}}} methods), we show that these simulations can in-fact be done using alternating queries, thus relating PAC{\mathsf{PAC}} and SQ{\mathsf{SQ}} to bSQ0/1{{\mathsf{bSQ^{0/1}}}}. We present the full details in Sections A.1 and A.3 respectively.

(PAC{\mathsf{PAC}} to bSQ0/1{{\mathsf{bSQ^{0/1}}}}) For all δ>0\delta>0, for all bb, and τ<1/(2b)\tau<1/(2b), then and for all m,rm,r, it holds for k′=20m(n+1)/δk^{\prime}=20m(n+1)/\delta, p′=n+1p^{\prime}=n+1, r′=r+klog⁡2br^{\prime}=r+k\log_{2}b that

(SQ{\mathsf{SQ}} to bSQ0/1{{\mathsf{bSQ^{0/1}}}}) There exists a constant C>0C>0, such that for all δ>0\delta>0, for all bb and all k,τ,rk,\tau,r, it holds that

We now show how to to simulate a bSQ0/1{{\mathsf{bSQ^{0/1}}}} method with bSGD{\mathsf{bSGD}}, with corresponding mini-batch and precision:

We first show how a single y‾\overline{y}-query can be simulated using a single step of bSGD{\mathsf{bSGD}} on a specific differentiable model. Given a -query Φ:X×Y→p\Phi:\mathcal{X}\times\mathcal{Y}\to^{p}, consider the following model:

With w(0)=0{\bm{w}}^{(0)}=0, the model fw(0)f_{{\bm{w}}^{(0)}} “guesses” the label to be 11 for all examples, and therefore suffers a loss only for examples with the label . Using a simple gradient calculations we get:

Combining Lemma 3a with Lemmas 1 and 2 establishes the first statements (about computationally unbounded learning) of Theorems 1a and 1d; full details in Appendix D.

Implementing the differentiable model as a Neural Network.

The simulation in Lemma 3a, as described above, uses some arbitrary differentiable model fw(x)f_{\bm{w}}(x), that is defined in terms of the mappings from responses to queries in the bSQ0/1{{\mathsf{bSQ^{0/1}}}} method. If the bSQ0/1{{\mathsf{bSQ^{0/1}}}} method is computationally bounded, the simulation can also be done using a neural network:

Combining Lemma 3b with Lemmas 1 and 2 establishes the second statements (about computationally bounded learning) of Theorems 1a and 1d; full details in Appendix D.

In order to establish Theorems 1b and 1c we rely on Theorems 2b and 2c and for that purpose note that bSGD{\mathsf{bSGD}} can be directly implemented using bSQ{{\mathsf{bSQ}}} (proof in Appendix D):

(bSGD{\mathsf{bSGD}} to bSQ{{\mathsf{bSQ}}}) For all T,ρ,b,p,rT,\rho,b,p,r, it holds that

Furthermore, for every poly-time computable activation σ\sigma, it holds that

So far we considered learning with mini-batch stochastic gradient descent (bSGD{\mathsf{bSGD}}), where an independent mini-batch of examples is used at each step. But this stochasticity, and the use of independent fresh samples for each gradient step, is not crucial for simulating PAC{\mathsf{PAC}} and SQ{\mathsf{SQ}}, provided enough samples overall, and a correspondingly fine enough precision. We show that analogous results hold for learning with full-batch Gradient Descent (fbGD{\mathsf{fbGD}}), i.e. gradient descent on the (fixed) empirical loss.

For all mm and ρ<1/(8m)\rho<1/(8m) and for all rr, it holds that

For all mm, ρ\rho and T,p,rT,p,r, it holds that

Furthermore, for all poly-time computable activations σ\sigma, it holds that

There exists a constant CC such that for all δ>0\delta>0, for all TT, ρ\rho, mm, pp, rr, such that mρ2>C(Tplog⁡(1/ρ)+log⁡(1/δ))m\rho^{2}>C(Tp\log(1/\rho)+\log(1/\delta)), it holds that

Furthermore, for all poly-time computable activations σ\sigma, it holds that

There exists a constant CC such that for all δ>0\delta>0, for all k,τ,rk,\tau,r, it holds for m,ρm,\rho such that ρ=τ/16\rho=\tau/16, mρ2>C(klog⁡(1/ρ)+log⁡(1/δ))m\rho^{2}>C(k\log(1/\rho)+\log(1/\delta)) that

The above theorems are analogous to Theorems 1a, 1b, 1c and 1d. They are proved in an analogous manner in Appendix E, by going through the intermediate model of fixed-batch statistical query fbSQ{{\mathsf{fbSQ}}}, in place of bSQ{{\mathsf{bSQ}}}. An fbSQ(k,τ,m,p,r){{\mathsf{fbSQ}}}(k,\tau,m,p,r) method is described identically to an bSQ(k,τ,b=m,p,r){{\mathsf{bSQ}}}(k,\tau,b=m,p,r) method, except that the responses for all queries are obtained using the same batch of samples in all rounds (i.e. St=SS_{t}=S for all tt, where S∼DmS\sim\mathcal{D}^{m} in Equation 8). Simulating fbSQ{{\mathsf{fbSQ}}} methods using fbGD{\mathsf{fbGD}}, or fbGDNN{\mathsf{fbGD_{NN}}} can be done using the exact same constructions as in Lemmas 3a and 3b. To establish Theorem 3a, we use an algorithm similar to (and simpler than) Algorithm 1 to extract all the samples batch of samples (see details in Lemma 5a in Section A.2). Relating fbSQ{{\mathsf{fbSQ}}} to SQ{\mathsf{SQ}} and establishing Theorems 3c and 3d requires more care, because of the adaptive nature of fbGD{\mathsf{fbGD}} on the full-batch. Instead, we consider all possible queries the method might make, based on previous responses. Since we have at most Tplog⁡(1/ρ)Tp\log(1/\rho) or kplog⁡(1/τ)kp\log(1/\tau) bits of response to choose a new query based on, we need to take a union bound over a number of queries exponential in this quantity, which results in the sample sized required to ensure validity scaling linear in kplog⁡(1/τ)kp\log(1/\tau). See Section A.4 for complete proofs and details.

fbGD=PAC{\mathsf{fbGD}}={\mathsf{PAC}}\quad and fbGDNN=PACTM\quad{\mathsf{fbGD_{NN}}}={\mathsf{PAC_{TM}}}.

But a significant difference versus bSGD{\mathsf{bSGD}} is that with fbGD{\mathsf{fbGD}} the precision depends (even if only polynomially) on the total number of samples used by the method. This is in contrast to bSGD{\mathsf{bSGD}}, where the precision only has to be related to the mini-batch size used, and with constant precision (and constant mini-batch size), we could simulate any sample based based method, regardless of the number of samples used by the method (only the number TT of SGD iterations and the size pp of the model increase with the number of samples used). Viewed differently, consider what can be done with some fixed precision ρ\rho (that is not allowed to depend on the problem size nn or sample size mm): methods that use up to 1/(8ρ)1/(8\rho) samples can be simulated even with fbGD{\mathsf{fbGD}}. But bSGD{\mathsf{bSGD}} allows us to simulate methods that use even more samples, by keeping the mini-batch size below 1/(8ρ)1/(8\rho).

Perhaps the most realistic differentiable learning approach is to use a fixed training set SS, and then at each iteration calculate a gradient estimate based on a mini-batch St⊂SS_{t}\subset S chosen at random, with replacement, from within the training set SS (as opposed to using fresh samples from the population distribution, as in bSGD{\mathsf{bSGD}}). Analogs of Theorems 1a and 1d and Theorem 3c should hold also for this hybrid class, but we do not provide details here.

Summary and Discussion

We provided an almost tight characterization of the learning power of mini-batch SGD, relating it to the well-studied learning paradigms PAC{\mathsf{PAC}} and SQ{\mathsf{SQ}}, and thus (nearly) settling the question of “what can be learned using mini-batch SGD?”. That single-sample SGD is able to simulate PAC{\mathsf{PAC}} learning was previously known, but we extended this result considerably, studied its limit, and showed that even outside this limit, bSGD{\mathsf{bSGD}} can still always simulate SQ{\mathsf{SQ}}. A gap still remains, when the mini-batch size is between 1/ρ1/\rho and log⁡(n)/ρ2\log(n)/\rho^{2}, where we do not know where bSGD{\mathsf{bSGD}} sits between SQ{\mathsf{SQ}} and PAC{\mathsf{PAC}}. We furthermore showed that with sufficient (polynomial) precision, even full Gradient Descent on an empirical loss can simulate PAC{\mathsf{PAC}} learning.

It is tempting to view our results, which show the theoretical power of differentiable learning, as explaining the success of this paradigm. But we do not think that modern deep learning behaves similar to the constructions in our work. While we show how any SQ{\mathsf{SQ}} or PAC{\mathsf{PAC}} algorithm can be simulated, this requires a very carefully constructed network, with an extremely particular initialization, which doesn’t look anything like deep learning in current practice. Our result certainly does not imply that SGD on a particular neural net can learn anything learnable by PAC{\mathsf{PAC}} or SQ{\mathsf{SQ}}, as this would imply that such network can learn any computationally tractable functionObserve that for any tractable function ff, there exists a trivial learning algorithm that returns ff regardless of its input, which means that the class {f}\{f\} is PAC{\mathsf{PAC}} learnable., which is known to be impossible (subject to mild cryptographic assumptions).

Rather, we view our work as guiding us as to what questions we should ask toward understanding how actual deep learning works. We see that understanding differentiable learning in such a broad generality as we did here is probably too strong, as it results in answers involving unrealistic initialization, and no restriction, and thus no insight, as to what makes learning problems learnable using deep learning. Can we define a class of neural networks, or initializations, which is broad enough to capture the power of deep learning, yet disallows such crazy initialization and does provide insight as to when deep learning is appropriate? Perhaps even mild restrictions on the initialization can already severely restrict the power of differentiable learning. E.g., Malach et al. (2021) recently showed that even just requiring that the output of the network at initialization is close to zero can significantly change the power of differentiable learning, Abbe et al. (2021) showed that imposing certain additional regularity assumptions on the architecture/initialization of neural networks restricts the learning power of (S)GD to function classes with a certain hierarchical property. An interesting direction for future work is understanding the power of differentiable learning under these, or other, restrictions. Does this lead to a different class of learnable problems, distinct from SQ{\mathsf{SQ}} and PAC{\mathsf{PAC}}, which is perhaps more related to deep learning in practice?

This work was done as part of the NSF-Simons Sponsored Collaboration on the Theoretical Foundations of Deep Learning. Part of this work was done while PK was at TTIC, and while NS was visiting EPFL. PK and NS were supported by NSF BIGDATA award 1546500 and NSF CCF/IIS award 1764032.

References

In this section, we prove Theorems 2a, 2b, 2c and 2d. Additionally, we state and prove analogous statements relating fbSQ{{\mathsf{fbSQ}}} to PAC{\mathsf{PAC}} and SQ{\mathsf{SQ}}.

For all bb, τ\tau satisfying bτ<1/2b\tau<1/2, we first design a bSQ(k=10(n+1),τ,b,p=n+1,r){{\mathsf{bSQ}}}(k=10(n+1),\tau,b,p=n+1,r) algorithm Sample-Extract (Algorithm 1) that generates a single sample (x,y)∼D(x,y)\sim\mathcal{D}; technically, this algorithm runs in at most 10(n+1)10(n+1) expected number of steps, but as we will see this is sufficient to complete the proof. For ease of notation, we let Z:=X×Y\mathcal{Z}:=\mathcal{X}\times\mathcal{Y} and we identify Z\mathcal{Z} with {0,1}n+1\left\{0,1\right\}^{n+1} and denote z∈Zz\in\mathcal{Z} as (z1,…,zn+1)=(y,x1,…,xn)(z_{1},\ldots,z_{n+1})=(y,x_{1},\ldots,x_{n}). Sample-Extract operates by sampling the bits of zz one by one, drawing z^i\widehat{z}_{i} from the conditional distribution {zi∣z1,…,i−1}D\left\{z_{i}\mid z_{1,\ldots,i-1}\right\}_{\mathcal{D}}. We show two things: (i) Once the algorithm is at a prefix ss, the algorithm indeed returns a sample z^\widehat{z} drawn from the conditional distribution {z∣z1,…,k=s}D\left\{z\mid z_{1,\ldots,k}=s\right\}_{\mathcal{D}}, and (ii) The algorithm returns a sample in T≤10(n+1)T\leq 10(n+1) expected number of steps.

Thus, the expected number of steps before we get at least one sample starting with ss is at most 2/bps2/bp_{s}. Also, Pr⁡[bs>0]=1−(1−ps)b≤bps\Pr[b_{s}>0]=1-(1-p_{s})^{b}\leq bp_{s}. So, for p0:=ps∘0p_{0}:=p_{s\circ 0} and p1:=ps∘1p_{1}:=p_{s\circ 1} (that is, the probabilities that a sample from D\mathcal{D} starts with s∘0s\circ 0 and s∘1s\circ 1 respectively), the expected number of steps remaining in the algorithm is at most

as claimed. Next, we show that if the algorithm currently has a prefix ss and bps>1/5bp_{s}>1/5, then the expected number of steps remaining is at most 10(n+1−∣s∣)10(n+1-|s|). We again prove this with a reverse induction on ∣s∣|s|. The base case of ∣s∣=n+1|s|=n+1 is trivial. With bps>1/5bp_{s}>1/5, the probability that a batch of bb samples has at least one sample starting with prefix ss is at least 1/81/8. So, the expected remaining number of steps before the algorithm terminates is at most

as desired. That completes the induction argument. The algorithm starts with s=ϵs=\epsilon, which every sample will start with, so the expected number of steps before the algorithm terminates is at most 10(n+1)10(n+1).

This is immediate, since a PAC(m=kb,r){\mathsf{PAC}}(m=kb,r) method can generate valid bSQ{{\mathsf{bSQ}}} responses using bb samples for each of the kk rounds by simply computing the empirical averages per batch. The number of random bits used remains unchanged. ∎

Finally we show that with a slight modification, Sample-Extract can be implemented as a bSQ0/1{{\mathsf{bSQ^{0/1}}}} algorithm, thereby proving Lemma 1, restated below for convenience.

Since z1=yz_{1}=y, all queries of Sample-Extract are already of the form \mathds1{y=y‾}∧ΦX(x)\mathds{1}\left\{y=\overline{y}\right\}\wedge\Phi_{\mathcal{X}}(x) for y‾∈{0,1}\overline{y}\in\{0,1\}. It can also be implemented as a bSQ0/1{{\mathsf{bSQ^{0/1}}}} algorithm as follows: After the first query we fix y=s1∈{0,1}y=s_{1}\in\{0,1\}. If s1=0s_{1}=0, we use only even rounds to perform the queries as done by Sample-Extract and when s1=1s_{1}=1, we use only odd rounds. This increases the total number of rounds by a factor of 22. ∎

A.2 𝗳𝗯𝗦𝗤𝗳𝗯𝗦𝗤{{\mathsf{fbSQ}}} versus 𝗣𝗔𝗖𝗣𝗔𝗖{\mathsf{PAC}}

We show the analogs of Theorems 2a and 2b for fbSQ{{\mathsf{fbSQ}}}.

(PAC{\mathsf{PAC}} to fbSQ{{\mathsf{fbSQ}}}) For all mm, and τ<1/(2m)\tau<1/(2m) and for all rr, it holds that

Starting with the prefix s=ϵs=\epsilon (empty string) and knowledge of ∣Sϵ∣=m|S_{\epsilon}|=m, we can recover all samples using at most m(n+1)m(n+1) fbSQ{{\mathsf{fbSQ}}}s, after which we can simply simulate the PAC(m,r){\mathsf{PAC}}(m,r) method. Note that unlike the reduction to bSQ{{\mathsf{bSQ}}}, here the algorithm always succeeds in extracting mm samples in m(n+1)m(n+1) steps. Hence there is no loss in the error ensured. ∎

(fbSQ{{\mathsf{fbSQ}}} to PAC{\mathsf{PAC}}) For all mm, τ\tau and k,p,rk,p,r, it holds that

This is immediate, since a PAC(m,r){\mathsf{PAC}}(m,r) method can generate valid fbSQ{{\mathsf{fbSQ}}} responses using mm samples for each of the kk rounds by simply computing the empirical averages over the entire batch of samples. The number of random bits used remains unchanged. ∎

Finally, we show that with a slight modification, in the same regime of Lemma 5a, any PAC{\mathsf{PAC}} method can be simulated by a fbSQ0/1{{\mathsf{fbSQ^{0/1}}}} method, analogous to Lemma 1.

(PAC{\mathsf{PAC}} to fbSQ0/1{{\mathsf{fbSQ^{0/1}}}}) For all mm, and τ<1/(2m)\tau<1/(2m), and for all rr, it holds that

By modifying proof of Lemma 5a, analogous to the modification to proof of Theorem 2a to get Lemma 1. ∎

A.3 𝗯𝗦𝗤𝗯𝗦𝗤{{\mathsf{bSQ}}} versus 𝗦𝗤𝗦𝗤{\mathsf{SQ}}

Fix a bSQ(k,τ,b,p,r){{\mathsf{bSQ}}}(k,\tau,b,p,r) method A\mathcal{A}, and consider any bSQ{{\mathsf{bSQ}}} query Φ:X×Y→p\Phi:\mathcal{X}\times\mathcal{Y}\to^{p}. Using Chernoff-Hoeffding’s bound and a union bound over pp entries, we have that

Finally we show that any SQ{\mathsf{SQ}} method can be simulated by a bSQ0/1{{\mathsf{bSQ^{0/1}}}} method thereby proving Lemma 2, restated below for convenience. The proof goes via an intermediate SQ0/1{\mathsf{SQ^{0/1}}} method (defined analogous to bSQ0/1{{\mathsf{bSQ^{0/1}}}}).

Finally, we use essentially the same argument as in Theorem 2d to obtain a bSQ0/1{{\mathsf{bSQ^{0/1}}}} method from A′\mathcal{A}^{\prime}. The only change needed to preserve the alternating nature of A′\mathcal{A}^{\prime} is that we perform all the 11-queries in qq odd rounds, interleaved with -queries in qq even rounds. This completes the proof. ∎

A.4 𝗳𝗯𝗦𝗤𝗳𝗯𝗦𝗤{{\mathsf{fbSQ}}} versus 𝗦𝗤𝗦𝗤{\mathsf{SQ}}

We show the analogs of Theorems 2c and 2d for fbSQ{{\mathsf{fbSQ}}}. The number of samples required depends linearly on the number of queries, instead of logarithmically. The reason for this is that unlike in the proof of Theorems 2c and 2d, a naive union bound does not suffice, since the queries can be adaptive.

(fbSQ{{\mathsf{fbSQ}}} to SQ{\mathsf{SQ}}) There exist a constant C≥0C\geq 0 such that for all δ>0\delta>0, for all k,τ,m,p,rk,{\bm{\tau}},{\bm{m}},p,r such that mτ2>C(kplog⁡(1/τ)+log⁡(1/δ))m\tau^{2}>C(kp\log(1/\tau)+\log(1/\delta)), it holds that

Using Chernoff-Hoeffding’s bound (Equation 10) and a union bound over all these (1/α+1)kp(1/\alpha+1)^{kp} possible queries we have that with probability at least 1−2p(1+1/α)kp⋅e−η2m/21-2p(1+1/\alpha)^{kp}\cdot e^{-\eta^{2}m/2} over sampling S∼DmS\sim\mathcal{D}^{m}, it holds for each such query Φ:X×Y→p\Phi:\mathcal{X}\times\mathcal{Y}\to^{p} that

(SQ{\mathsf{SQ}} to fbSQ{{\mathsf{fbSQ}}}) There exist a constant C≥0C\geq 0 such that for all δ>0\delta>0, k,τ,rk,\tau,r, it holds for τ′=τ2\tau^{\prime}=\frac{\tau}{2} and all mm such that mτ2>C(klog⁡(1/τ)+log⁡(1/δ))m\tau^{2}>C(k\log(1/\tau)+\log(1/\delta)) that

Finally, we show that with a slight modification, in the same regime of Lemma 5d, any SQ{\mathsf{SQ}} method can be simulated by a fbSQ0/1{{\mathsf{fbSQ^{0/1}}}} method, analogous to Lemma 2.

(SQ{\mathsf{SQ}} to fbSQ0/1{{\mathsf{fbSQ^{0/1}}}}) There exist a constant C≥0C\geq 0 such that for all δ>0\delta>0, k,τ,rk,\tau,r, it holds for τ′=τ4\tau^{\prime}=\frac{\tau}{4} and all mm such that mτ2>C(klog⁡(1/τ)+log⁡(1/δ))m\tau^{2}>C(k\log(1/\tau)+\log(1/\delta)) that

By modifying proof of Lemma 5d, analogous to the modification to proof of Theorem 2d to get Lemma 2. ∎

Appendix B Simulating 𝗯𝗦𝗤𝗯𝗦𝗤{{\mathsf{bSQ}}} with 𝗯𝗦𝗚𝗗𝗯𝗦𝗚𝗗{\mathsf{bSGD}} : Proof of Lemma 3a

In this section, we show how to simulate any bSQ0/1{{\mathsf{bSQ^{0/1}}}} method A\mathcal{A} as bSGD{\mathsf{bSGD}} on some differentiable model constructed according to A\mathcal{A}, and thus prove Lemma 3a:

As a first step towards showing Lemma 3a, we show how a single y‾\overline{y}-query can be simulated using a single step of bSGD{\mathsf{bSGD}}. We consider parameterized queries, that can depend on some of parameters of the differentiable model. Namely, Φ:q×X×Y→p\Phi:^{q}\times\mathcal{X}\times\mathcal{Y}\to^{p} is a query that given some parameters θ∈q\theta\in^{q}, an input x∈Xx\in\mathcal{X} and a label y∈Yy\in\mathcal{Y}, returns some vector value Φ(θ,x,y)\Phi(\theta,x,y). We show the following:

∥θ(1)−1b∑(x,y)∈SΦ(θ^(0),x,y)∥∞≤ε+ρ\left\lVert\theta^{(1)}-\frac{1}{b}\sum_{(x,y)\in S}\Phi(\widehat{\theta}^{(0)},x,y)\right\rVert_{\infty}\leq\varepsilon+\rho,

Assume Φ\Phi is a -query, namely Φ(x,y)=(1−y)Φ(θ^,x,0)\Phi(x,y)=(1-y)\Phi(\widehat{\theta},x,0). Then, we define a differentiable model as follows:

Then, after performing one step of bSGD{\mathsf{bSGD}} we have:

And therefore, after one step of bSGD{\mathsf{bSGD}} we have κ(1)≥ε−ρ\kappa^{(1)}\geq\varepsilon-\rho.

Assume Φ\Phi is a 11-query, namely Φ(x,y)=yΦ(θ^,x,0)\Phi(x,y)=y\Phi(\widehat{\theta},x,0). Then, we define a differentiable model as follows:

Then, after performing one step of bSGD{\mathsf{bSGD}} we have:

And therefore, after one step of bSGD{\mathsf{bSGD}} we have κ(1)≥ε−ρ\kappa^{(1)}\geq\varepsilon-\rho.∎

Our goal is to use Lemma 8 to simulate a bSQ0/1{{\mathsf{bSQ^{0/1}}}} method. Any bSQ0/1{{\mathsf{bSQ^{0/1}}}} method A\mathcal{A} is completely described by a sequence of (potentially adaptive) queries Φ1,…,ΦT\Phi_{1},\dots,\Phi_{T}, and a predictor hh which depends on the answer to previous queries, namely:

Φt\Phi_{t} depends on rr random bits, and on the answers of the previous t−1t-1 queries, namely:

For every compact set KK and open set UU such that K⊆U⊆qK\subseteq U\subseteq^{q}, there exists a smooth function Ψ:q→\Psi:^{q}\to such that Ψ(x)=1\Psi(x)=1 for every x∈Kx\in K and Ψ(x)=0\Psi(x)=0 for every x∉Ux\notin U.

Our differentiable model will use the following parameter:

TT “clock” parameters κ1,…κT\kappa_{1},\dots\kappa_{T}, that indicate which query should be issued next. We initialize κ1,…,κT=0\kappa_{1},\dots,\kappa_{T}=0.

We denote by θ(t,i),κt(i)\theta^{(t,i)},\kappa_{t}^{(i)} the value of the tt-th set of parameters in the ii-th iteration of SGD. The differentiable model is defined as follows:

Now, denote v0:=θ(0,0)v_{0}:=\theta^{(0,0)}, and for every t>0t>0 denote vt:=θ(t,t)v_{t}:=\theta^{(t,t)}. We have the following claim:

Claim: for every iteration ii of bSGD{\mathsf{bSGD}},

For every t>it>i we have θ(t,i)=0\theta^{(t,i)}=0 and κt(i)=0\kappa^{(i)}_{t}=0.

For t<it<i we have θ(t,i)=θ(t,t)\theta^{(t,i)}=\theta^{(t,t)}.

For t≤it\leq i we have κt(i)≥ρ\kappa_{t}^{(i)}\geq\rho.

For i=1i=1, notice that by the initialization, κt(0)=0\kappa_{t}^{(0)}=0 for every tt. Fix some t>1=it>1=i, and note that c(κt−1,κt,α)=0c(\kappa_{t-1},\kappa_{t},\alpha)=0, so the gradient w.r.t θ(t,i)\theta^{(t,i)}, κt(i)\kappa_{t}^{(i)} is zero, and so condition 1 hold (the initialization is zero, and the gradient is zero). For t=1=it=1=i, notice that since c(κ0,κ1(0),α)=αc(\kappa_{0},\kappa_{1}^{(0)},\alpha)=\alpha, we have:

and by applying Lemma 8 with ε=2ρ\varepsilon=2\rho we get:

and so condition 4 follows. Finally, condition 3 is vacuously true.

Fix some i>0i>0, and assume the claim holds for ii. We will prove the claim for i+1i+1. By the assumption, we have κt(i)=0\kappa_{t}^{(i)}=0 for every t>it>i and κt(i)≥ρ\kappa_{t}^{(i)}\geq\rho for every t≤it\leq i. Therefore, by definition of cc, we have c(κt−1(i),κt(i),α)=\mathds1t=i+1αc(\kappa_{t-1}^{(i)},\kappa_{t}^{(i)},\alpha)=\mathds{1}_{t=i+1}\alpha. So,

Therefore, conditions 1 and 3 follow from the fact that the gradient with respect to θ(t,i)\theta^{(t,i)} and κt(i)\kappa_{t}^{(i)}, for every t≠i+1t\neq i+1, is zero. Now, using Lemma 8 with ε=2ρ\varepsilon=2\rho, condition 2 follows, and we also have κi+1(i+1)≥ρ\kappa_{i+1}^{(i+1)}\geq\rho. Finally, for every t<i+1t<i+1, by the assumption we have κt(i)≥ρ\kappa_{t}^{(i)}\geq\rho, and since the gradient with respect to κt(i)\kappa_{t}^{(i)} is zero, we also have κt(i+1)≥ρ\kappa_{t}^{(i+1)}\geq\rho. Therefore, condition 4 follows.

Finally, to prove the Theorem, observe that by the previous claim, for every tt:

So, vtv_{t} is a valid response for the tt-th query.

By the previous claim, we have κt(T)≥ρ\kappa_{t}^{(T)}\geq\rho for every 1≤t≤T1\leq t\leq T. Therefore, we have:

and, using the fact that v0,…,vTv_{0},\dots,v_{T} are valid responses to the method’s queries, we get the required. ∎

B.1 From Arbitrary Differentiable Models to Neural Networks.

In this section we proved the key lemma for our main results, showing that alternating batch-SQ methods can be simulated by gradient descent over arbitrary differentiable models. We would furthermore like to show that if the alternative batch-SQ method is computationally bounded, the differentiable model we defined can be implemented as a neural network of bounded size.

Indeed, observe that when the method can be implemented using a Turing-machine, each query (denoted by Φ\Phi in the proof) can be simulated by a Boolean circuit [see Arora and Barak, 2009], and hence by a neural network with some fixed weights. Therefore, one can show with little extra effort that the differentiable model introduced in the proof of Lemma 3a can be written as a neural network, with some of the weights being fixed. To show that the same behavior is guaranteed even when all the weights are trained, it is enough to show that all the relevant weights (e.g., θ(0),…,θ(T)\theta^{(0)},\dots,\theta^{(T)}) have zero gradient, unless they are correctly updated. This can be achieved using the “clock” mechanism (the function c(α1,α2,α3)c(\alpha_{1},\alpha_{2},\alpha_{3}) in the construction), which in turn can be implemented by a neural network that is robust to small perturbations of its weights, and hence does not suffer from unwanted updates of gradient descent. One possible way to implement the clock mechanism using a neural network that is robust to small perturbations is to rely either on large weight magnitudes and small step-sizes, or on the clipping of large gradients.

We do not include these details, Instead, in the next Section, we provide complete details and a rigorous proof of an alternate, more direct, construction of a neural network defining a differentiable model that simulates a given bSQTM0/1{{\mathsf{bSQ_{TM}^{0/1}}}} method. This direct neural-network construction is based on the same ideas, but is different in implementation from the construction shown in this section, involving some technical details to ensure that the network is well behaved under the gradient descent updates.

In this section, we show a direct construction of a neural network such that gradient descent on the neural net simulates a given bSQTM0/1{{\mathsf{bSQ_{TM}^{0/1}}}} method, thus proving Lemma 3b:

Given a bSQ{{\mathsf{bSQ}}} algorithm with a specified bounded runtime, we will design a neural network such that the mini-batch gradients at each step correspond to responses to queries of the bSQ{{\mathsf{bSQ}}} algorithm. Our proof of this is based on the fact that any efficient bSQTM0/1{{\mathsf{bSQ_{TM}^{0/1}}}} algorithm must decide what query to make next and what to output based on some efficient computation performed on random bits and the results of previous queries. Any efficient algorithm can be performed by a neural net, and it is possible to encode any circuit as a neural net of comparable size in which every vertex is always at a flat part of the activation function. Doing that would ensure that none of the edge weights ever change, and thus that the net would continue computing the desired function indefinitely. So, that allows us to give our net a subgraph that performs arbitrary efficient computations on the net’s inputs and on activations of other vertices.

Also, we can rewrite any bSQTM0/1{{\mathsf{bSQ_{TM}^{0/1}}}} algorithm to only perform binary queries by taking all of the queries it was going to perform, and querying the iith bit of their binary representation for all sufficiently small ii instead. For each of the resulting binary queries, we will have a corresponding vertex with an edge going to it from the constant vertex and no other edges going to it. So, the computation subgraph of the net will be able to determine the current weights of the edges leading to the query vertices by checking their activations. Also, each query vertex will have paths from it to the output vertex with intermediate vertices that will either get inputs in the flat parts of their activation function or not depending on the output of some vertices in the computation subgraph. The net effect of this will be to allow the computation subgraph to either make the value encoded by the query edge stay the same or make it increase if the net’s output differs from the sample output based on any efficiently computable function of the inputs and other query vertices’ activations. This allows us to encode an arbitrary bSQfgt{{\mathsf{bSQ^{\text{fgt}}}}} algorithm as a neural net.

In order to prove the capabilities of a neural net trained by batch stochastic gradient descent, we will start by proving that any algorithm in bSQ{{\mathsf{bSQ}}} can be emulated by a neural net trained by batch stochastic gradient descent under appropriate parameters. In this section our net will use an activation function σ\sigma, as defined in Figure 1, namely

Any bSQ{{\mathsf{bSQ}}} algorithm repeatedly makes a query and then computes what query to perform next from the results of the previous queries. So, our neural net will have a component designed so that we can make it update targeted edge weights by an amount proportional to the value of an appropriate query on the current batch and a component designed to allow us to perform computations on these edge weights. We will start by proving that we can make the latter component work correctly. More formally, we assert the following.

Let h:{0,1}m→{0,1}m′h:\{0,1\}^{m}\rightarrow\{0,1\}^{m^{\prime}} be a function that can be computed by a circuit made of AND, OR, and NOT gates with a total of bb gates. Also, consider a neural net with mm inputNote that these will not be the nn data input of the general neural net that is being built; these input vertices take both the data inputs and some inputs from the memory component. vertices v1′,...,vm′v^{\prime}_{1},...,v^{\prime}_{m}, and choose real numbers yi(0)<yi(1)y_{i}^{(0)}<y_{i}^{(1)} for each 1≤i≤m1\leq i\leq m. It is possible to add a set of at most bb new vertices to the net, including output vertices v1′′,...,vm′′′v^{\prime\prime}_{1},...,v^{\prime\prime}_{m^{\prime}}, along with edges leading to them such that for any possible addition of edges leading from the new vertices to old vertices, if the net is trained by bSGD, the output of vi′v^{\prime}_{i} is either less than yi(0)y_{i}^{(0)} or more than yi(1)y_{i}^{(1)} for every ii in every timestep, then the following hold:

The derivative of the loss function with respect to the weight of each edge leading to a new vertex is in every timestep, and no paths through the new vertices contribute to the derivative of the loss function with respect to edges leading to the vi′v^{\prime}_{i}.

In any given time step, if the output of vi′v^{\prime}_{i} encodes xix_{i} with values less than yi(0)y_{i}^{(0)} and values greater than yi(1)y_{i}^{(1)} representing and 11 respectively for each ii, then the output of vj′′v^{\prime\prime}_{j} encodes hj(x1,...,xm)h_{j}(x_{1},...,x_{m}) for each jj with −2-2 and 22 encoding and 11 respectively.

In order to do this, we will add one new vertex for each gate and each input in a circuit that computes hh. When the new vertices are used to compute hh, we want each vertex to output 22 if the corresponding gate or input outputs a 11 and −2-2 if the corresponding gate or input outputs a , and we want the derivative of its activation with respect to its input to be . In order to do that, we need the vertex to receive an input of more than 22 if the corresponding gate outputs a 11 and an input of less than −3-3 if the corresponding gate outputs a .

In order to make one new vertex compute the NOT of another new vertex, it suffices to have an edge of weight −2-2 to the vertex computing the NOT and no other edges to that vertex. We can compute an AND of two new vertices by having a vertex with two edges of weight 22 from these vertices and an edge of weight −4-4 from the constant vertex. Similarly, we can compute an OR of two new vertices by having a vertex with two edges of weight 22 from these vertices and an edge of weight 44 from the constant vertex. For each ii, in order to make a new vertex corresponding to the iith input, we add a vertex and give it an edge of weight 8/(y(1)−y(0))8/(y^{(1)}-y^{(0)}) from the associated vi′v^{\prime}_{i} and an edge of weight −(4y(1)+4y(0))/(y(1)−y(0))-(4y^{(1)}+4y^{(0)})/(y^{(1)}-y^{(0)}) from the constant vertex. These provide an overall input of at least 44 to the new vertex if vi′v^{\prime}_{i} has an output greater than y(1)y^{(1)} and an input of at most −4-4 if vi′v^{\prime}_{i} has an output less than y(0)y^{(0)}.

This ensures that if the outputs of the vi′v^{\prime}_{i} encode binary values x1,...,xmx_{1},...,x_{m} appropriately, then each of the new vertices will output the value corresponding to the output of the appropriate gate or input. So, these vertices compute h(x1,...,xm)h(x_{1},...,x_{m}) correctly. Furthermore, since the input to each of these vertices is outside of $,thederivativesoftheiractivationfunctionswithrespecttotheirinputsareall.Assuch,thederivativeofthelossfunctionwithrespecttoanyoftheedgesleadingtothemisalways,andpathsthroughthemdonotcontributetochangesintheweightsofedgesleadingtothe, the derivatives of their activation functions with respect to their inputs are all . As such, the derivative of the loss function with respect to any of the edges leading to them is always , and paths through them do not contribute to changes in the weights of edges leading to thev^{\prime}_{i}$. ∎

Our next order of business is to show that we can perform queries successfully. So, we define the query subgraph as follows:

Given τ>0\tau>0, let QQ be the weighted directed graph with vertices v0v_{0}, v1v_{1}, v2v_{2}, v2′v^{\prime}_{2}, v3v_{3}, v4v_{4}, vcv_{c}, and virv^{r}_{i} for 0≤i<log⁡2(1/τ)0\leq i<\log_{2}(1/\tau) and the following edges:

An edge of weight 1/121/12 from v0v_{0} to v1v_{1}

Edges of weight 11 from v1v_{1} to v2v_{2} and v2′v^{\prime}_{2}, and from v2v_{2} and v2′v^{\prime}_{2} to v3v_{3}.

An edge of weight 1/41/4 from v3v_{3} to v4v_{4}.

An edge of weight 1010 from vcv_{c} to v2v_{2}.

An edge of weight −10-10 from vcv_{c} to v2′v^{\prime}_{2}.

An edge of weight −1/4-1/4 from vcv_{c} to v3v_{3}.

An edge of weight 12/τ12/\tau from v1v_{1} to virv^{r}_{i} for each ii.

An edge of weight −1/τ+6−6⋅2⌈log⁡2(1/τ)⌉-1/\tau+6-6\cdot 2^{\lceil\log_{2}(1/\tau)\rceil} from v0v_{0} to virv^{r}_{i} for each ii.

An edge of weight −3⋅2i-3\cdot 2^{i} from virv^{r}_{i} to vjrv^{r}_{j} for each i>ji>j.

Also, let Q′Q^{\prime} be the graph that is exactly like QQ except that in it the edge from v3v_{3} to v4v_{4} has a weight of −1-1.

v0<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><msub><mi>v</mi><mn>1</mn></msub></mrow><annotationencoding="application/x−tex">v1</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.5806em;vertical−align:−0.15em;"></span><spanclass="mord"><spanclass="mordmathnormal"style="margin−right:0.0359em;">v</span><spanclass="msupsub"><spanclass="vlist−tvlist−t2"><spanclass="vlist−r"><spanclass="vlist"style="height:0.3011em;"><spanstyle="top:−2.55em;margin−left:−0.0359em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmtight">1</span></span></span></span></span><spanclass="vlist−s">​</span></span><spanclass="vlist−r"><spanclass="vlist"style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>v2<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><msubsup><mi>v</mi><mn>2</mn><momathvariant="normal"lspace="0em"rspace="0em">′</mo></msubsup></mrow><annotationencoding="application/x−tex">v2′</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:1.0489em;vertical−align:−0.247em;"></span><spanclass="mord"><spanclass="mordmathnormal"style="margin−right:0.0359em;">v</span><spanclass="msupsub"><spanclass="vlist−tvlist−t2"><spanclass="vlist−r"><spanclass="vlist"style="height:0.8019em;"><spanstyle="top:−2.453em;margin−left:−0.0359em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmtight">2</span></span></span></span><spanstyle="top:−3.113em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmtight">′</span></span></span></span></span><spanclass="vlist−s">​</span></span><spanclass="vlist−r"><spanclass="vlist"style="height:0.247em;"><span></span></span></span></span></span></span></span></span></span></span>v3<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><msub><mi>v</mi><mn>4</mn></msub></mrow><annotationencoding="application/x−tex">v4</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.5806em;vertical−align:−0.15em;"></span><spanclass="mord"><spanclass="mordmathnormal"style="margin−right:0.0359em;">v</span><spanclass="msupsub"><spanclass="vlist−tvlist−t2"><spanclass="vlist−r"><spanclass="vlist"style="height:0.3011em;"><spanstyle="top:−2.55em;margin−left:−0.0359em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmtight">4</span></span></span></span></span><spanclass="vlist−s">​</span></span><spanclass="vlist−r"><spanclass="vlist"style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>vc<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><msubsup><mi>v</mi><mn>1</mn><mi>r</mi></msubsup></mrow><annotationencoding="application/x−tex">v1r</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.9614em;vertical−align:−0.247em;"></span><spanclass="mord"><spanclass="mordmathnormal"style="margin−right:0.0359em;">v</span><spanclass="msupsub"><spanclass="vlist−tvlist−t2"><spanclass="vlist−r"><spanclass="vlist"style="height:0.7144em;"><spanstyle="top:−2.453em;margin−left:−0.0359em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmtight">1</span></span></span></span><spanstyle="top:−3.113em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmathnormalmtight"style="margin−right:0.0278em;">r</span></span></span></span></span><spanclass="vlist−s">​</span></span><spanclass="vlist−r"><spanclass="vlist"style="height:0.247em;"><span></span></span></span></span></span></span></span></span></span></span>v2r<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><msubsup><mi>v</mi><mn>3</mn><mi>r</mi></msubsup></mrow><annotationencoding="application/x−tex">v3r</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.9614em;vertical−align:−0.247em;"></span><spanclass="mord"><spanclass="mordmathnormal"style="margin−right:0.0359em;">v</span><spanclass="msupsub"><spanclass="vlist−tvlist−t2"><spanclass="vlist−r"><spanclass="vlist"style="height:0.7144em;"><spanstyle="top:−2.453em;margin−left:−0.0359em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmtight">3</span></span></span></span><spanstyle="top:−3.113em;margin−right:0.05em;"><spanclass="pstrut"style="height:2.7em;"></span><spanclass="sizingreset−size6size3mtight"><spanclass="mordmtight"><spanclass="mordmathnormalmtight"style="margin−right:0.0278em;">r</span></span></span></span></span><spanclass="vlist−s">​</span></span><spanclass="vlist−r"><spanclass="vlist"style="height:0.247em;"><span></span></span></span></span></span></span></span></span></span></span>v4rv_{0}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi>v</mi><mn>1</mn></msub></mrow><annotation encoding="application/x-tex">v_{1}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.5806em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0359em;">v</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:-0.0359em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">1</span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>v_{2}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msubsup><mi>v</mi><mn>2</mn><mo mathvariant="normal" lspace="0em" rspace="0em">′</mo></msubsup></mrow><annotation encoding="application/x-tex">v_{2}^{\prime}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:1.0489em;vertical-align:-0.247em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0359em;">v</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.8019em;"><span style="top:-2.453em;margin-left:-0.0359em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">2</span></span></span></span><span style="top:-3.113em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">′</span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.247em;"><span></span></span></span></span></span></span></span></span></span></span>v_{3}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msub><mi>v</mi><mn>4</mn></msub></mrow><annotation encoding="application/x-tex">v_{4}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.5806em;vertical-align:-0.15em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0359em;">v</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.3011em;"><span style="top:-2.55em;margin-left:-0.0359em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">4</span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.15em;"><span></span></span></span></span></span></span></span></span></span></span>v_{c}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msubsup><mi>v</mi><mn>1</mn><mi>r</mi></msubsup></mrow><annotation encoding="application/x-tex">v^{r}_{1}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.9614em;vertical-align:-0.247em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0359em;">v</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.7144em;"><span style="top:-2.453em;margin-left:-0.0359em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">1</span></span></span></span><span style="top:-3.113em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mathnormal mtight" style="margin-right:0.0278em;">r</span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.247em;"><span></span></span></span></span></span></span></span></span></span></span>v^{r}_{2}<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><msubsup><mi>v</mi><mn>3</mn><mi>r</mi></msubsup></mrow><annotation encoding="application/x-tex">v^{r}_{3}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.9614em;vertical-align:-0.247em;"></span><span class="mord"><span class="mord mathnormal" style="margin-right:0.0359em;">v</span><span class="msupsub"><span class="vlist-t vlist-t2"><span class="vlist-r"><span class="vlist" style="height:0.7144em;"><span style="top:-2.453em;margin-left:-0.0359em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mtight">3</span></span></span></span><span style="top:-3.113em;margin-right:0.05em;"><span class="pstrut" style="height:2.7em;"></span><span class="sizing reset-size6 size3 mtight"><span class="mord mtight"><span class="mord mathnormal mtight" style="margin-right:0.0278em;">r</span></span></span></span></span><span class="vlist-s">​</span></span><span class="vlist-r"><span class="vlist" style="height:0.247em;"><span></span></span></span></span></span></span></span></span></span></span>v^{r}_{4} The idea behind this construction is as follows. v0v_{0} will be the constant vertex, and v4v_{4} will be the output vertex. If vcv_{c} has activation 22 then v2v_{2} will have activation 22 and v2′v^{\prime}_{2} will have activation −2-2, leaving v3v_{3} with activation . So, this subgraph will have no effect on the output and the derivative of the loss with respect to any of its edge weights will be . However, if vcv_{c} has activation then this subgraph will contribute to the net’s output, and the weights of the edges will change. So, when we do not want to use this subgraph to perform a query we simply set vcv_{c} to 22. In a timestep where we do want to use it to perform a query, we set vcv_{c} to or 22 based on the query’s value on the current input so that the weight of the edge from v0v_{0} to v1v_{1} will change based on the value of the query. The activations of the vrv^{r} will always give a binary representation of the activation of v1v_{1}. So, we can read off the current weight of the edge from v0v_{0} to v1v_{1}. This construction works in the following sense.

Let bb be an integer greater than 11, and 0<τ<1/120<\tau<1/12. Next, let (f,G)(f,G) be a neural net such that GG contains QQ or Q′Q^{\prime} as a subgraph with v0v_{0} as the constant vertex and v4v_{4} as GG’s output vertex, and there are no edges from vertices outside this subgraph to vertices in the subgraph other than vcv_{c} and v4v_{4}. Now, assume that this neural net is trained using bSGD with learning rate 22 and loss function LL for TT time steps, and the following hold:

vcv_{c} outputs or 22 on every sample in every time step.

There is at most one timestep in which the net has a sample on which vcv_{c} outputs . On any such sample, the output of the net is −2τ-2\tau if the subgraph is QQ and 1+2τ1+2\tau if the subgraph is Q′Q^{\prime}.

The derivatives of the loss function with respect to the weights of all edges leaving this subgraph are always .

The edge from v3v_{3} to v4v_{4} makes no contribution to v4v_{4} on any sample where vcv_{c} is 22, a contribution of exactly 1/241/24 on any sample where vcv_{c} is if the subgraph is QQ, and a contribution of exactly −1/24-1/24 on any sample where vcv_{c} is if the subgraph is −Q-Q. Also, if we regard an output of 22 as representing the digit 11 and an output of −2-2 as representing the digit then the binary string formed by concatenating the outputs of v⌈log⁡2(1/τ)⌉−1rv^{r}_{\lceil\log_{2}(1/\tau)\rceil-1},…,v0rv^{r}_{0} is within 3/23/2 of the number of samples in previous steps where vcv_{c} output and the net’s output did not match the sample output divided by bτb\tau plus double the total number of samples in previous steps where vcv_{c} output divided by bb.

First, let let tt be the timestep on which there is a sample with vcv_{c} outputting if any, and T+1T+1 otherwise. We claim that none of the weights of edges in this subgraph change on any timestep except step tt, and prove it by induction on tt. First, observe that by the definition of tt, given any sample the net receives on a timestep before tt, vcv_{c} has an output of 22. So, on any such sample v2v_{2} and v2′v^{\prime}_{2} will output 22 and −2-2 respectively, which results in v3v_{3} having an input of −1/2-1/2 and output of , and thus in both the derivative of v4v_{4} with respect to the weight of the edge from v3v_{3} to v4v_{4} and the derivative of v4v_{4} with respect to the input of v3v_{3} being . So, the derivative of the loss with respect to any of the edge weights in this component is on all samples received before step tt.

During step tt, for any sample on which vcv_{c} has an output of , v1v_{1}, v2v_{2}, and v2′v^{\prime}_{2} all have output 1/121/12. That results in v3v_{3} having an output of 1/61/6. So, the derivative of the loss with respect to the weight of the edge from v0v_{0} to v1v_{1} is τ\tau if the net’s output agrees with the sample output and (1+2τ)/2(1+2\tau)/2 otherwise. That means that the weight of the edge from v0v_{0} to v1v_{1} increases by 1/b1/b times the number of samples in this step for which vcv_{c} had output and the net’s output disagreed with the sample output plus 2τ/b2\tau/b times the total number of samples in this step for which vcv_{c} had an output of plus an error term of size at most 3τ/23\tau/2. Meanwhile, all of the other edges in the subgraph change by at most 1+3τ1+3\tau, and the weights of the edges from v2v_{2} and v2′v^{\prime}_{2} to v3v_{3} are left with weights within 3τ3\tau of each other because the derivatives of the gradients with respect to their weights are the same by symmetry, so only the error term differentiates them. The weights of the edges from vcv_{c} do not change because the samples for which this subgraph’s gradient is nonzero all have vcv_{c} outputting and thus making their weights irrelevant.

On any step after step tt, vcv_{c} gives an output of 22 on every sample. The weights of the edges from v1v_{1} to v2v_{2} and v2′v^{\prime}_{2} have absolute value of at most 33, so the edges from vcv_{c} still provide enough input to them to ensure that they output 22 and −2-2 respectively. That means that the input to v3v_{3} is within 6τ6\tau of −1/2-1/2, and thus that v3v_{3} outputs . That in turn means that the derivative of the loss with respect to any of the edge weights in the subgraph are so their weights do not change. All of this combined shows that the edge weights only ever change in step tt as desired.

Now, define even mm such that the weight of the edge from v0v_{0} to v1v_{1} increased by mτm\tau in step tt. We know that mτm\tau is within 3τ/23\tau/2 of the fraction of samples in step tt for which vcv_{c} output and the net’s output did not match the sample output plus τ\tau times the overall fraction of samples in step tt for which vcv_{c} output . The output of v1v_{1} is 1/121/12 until step tt and mτ+1/12m\tau+1/12 after step tt. Also, let k=⌈log⁡2(1/τ)⌉k=\lceil\log_{2}(1/\tau)\rceil. Now, pick some t′>tt^{\prime}>t and let rir_{i} be the output of virv^{r}_{i} on step t′t^{\prime}. In order to prove that the rir_{i} will be a binary encoding of mm we induct on k−ik-i. So, assume that vi+1rv^{r}_{i+1},…,vk−1rv^{r}_{k-1} encode the correct binary digits. Then the input to virv^{r}_{i} is

If the 2i2^{i} digit of mm’s binary representation is 11 this will be at least 66 while if it is it will be at most −6-6, so virv^{r}_{i} will output 22 if the digit is 11 and −2-2 if it is as desired. Showing that they all output −2-2 before step tt is just the m=0m=0 case of this.

Finally, recall that on every sample received after step tt, vcv_{c} will output 22, v3v_{3} will output , and thus the edge from v3v_{3} to v4v_{4} will provide no input to v4v_{4}. During or before step tt, the edges in this subgraph will all still have their original weights. So, if vcv_{c} outputs 22 then v3v_{3} outputs and makes no contribution to v4v_{4}, while if vcv_{c} outputs then v3v_{3} outputs 1/61/6 and makes a contribution of magnitude 1/241/24 and the appropriate sign to v4v_{4}, as desired. ∎

At this point, we claim that we can build a neural net that emulates any bSQTM0/1{{\mathsf{bSQ_{TM}^{0/1}}}} algorithm using the same value of bb and an error of τ/4\tau/4, provided τ<1/3\tau<1/3. In order to do that, we will structure our net as follows. First of all, we will have (plog⁡2(1/τ)+3)T(p\log_{2}(1/\tau)+3)T copies of QQ and Q′Q^{\prime}. Then, we build a computation component that takes input from all the copies of the vrv^{r} and computes from them what to output in the next step, what to query next, and what values those queries take on the current input. This component is built never to change as explained in Lemma 10. We will use a loss function of LL and learning rate of 22 when we train this net.

The net will also have TT primary output control vertices, (plog⁡2(1/τ)+3)T(p\log_{2}(1/\tau)+3)T secondary output control vertices, and 11 final output control vertex. Each of these will have edges of weight 11 from two different outputs of the computation component and an edge of weight −1/2-1/2 from the constant vertex. That way, the computation component will be able to control whether each of these vertices outputs −2-2, , or 22. Each primary output control vertex will have an edge of weight (1+2ρ)/2(1+2\rho)/2 to the output, each secondary output control vertex will have an edge of weight 1/481/48 to the output, and the final output control vertex will have an edge of weight 1/21/2 to the output. Our plan is to set a new group of output control vertices to nonzero values in each timestep and to use it to control the net’s output.

Each copy of vcv_{c} will have an edge of weight 1/21/2 from an output of the computation component and an edge of weight 11 from the constant vertex so the computation component will be able to control if it outputs or 22. That allows the computation component to query an arbitrary function to {0,1}\{0,1\} that is whenever the sample output takes on the wrong value by setting vcv_{c} to on every input for which the function is potentially nonzero and 11 on every input on which it is regardless of the sample output. Then the computation component can use a primary output control vertex to provide a value of ±(1+2ρ)\pm(1+2\rho) to the output and use a set of secondary output control vertices to cancel out the effects of the copies of QQ and Q′Q^{\prime} on the output.

A little more precisely, we will have one primary output control vertex, (plog⁡2(1/τ)+3)(p\log_{2}(1/\tau)+3) secondary output control vertices, and (plog⁡2(1/τ)+3)(p\log_{2}(1/\tau)+3) copies of QQ and Q′Q^{\prime} associated with each time step. Then, in that step it will determine what queries the bSQTM0/1{{\mathsf{bSQ_{TM}^{0/1}}}} algorithm it is emulating would have made, and use the copies of QQ and Q′Q^{\prime} to query the first ⌊log⁡2(1/τ)⌋+2\lfloor\log_{2}(1/\tau)\rfloor+2 digits of each of them and the constant function 11. In order to determine the details of this, it will consider the current time step as being the first step for which the component querying the constant function still gives an output of and consider each prior query performed by the bSQTM0/1{{\mathsf{bSQ_{TM}^{0/1}}}} algorithm as having given an output equal to the sum over 1≤i≤log⁡2(1/τ)+21\leq i\leq\log_{2}(1/\tau)+2 of ρ/2i\rho/2^{i} times the value given by the copy of QQ or Q′Q^{\prime} used to query its iith bit. For any given input/output pair the actual value of the query will be within ρ/2\rho/2 of the value given by its first ⌊log⁡2(1/τ)⌋+2\lfloor\log_{2}(1/\tau)\rfloor+2 binary digits. Also, if the conditions of Lemma 11 are satisfied then the value derived from the outputs of the QQ and Q′Q^{\prime} will be within ∑i(7/2)ρ⋅2−i≤(7/2)ρ\sum_{i}(7/2)\rho\cdot 2^{-i}\leq(7/2)\rho of the average of the values of the first ⌊log⁡2(1/τ)⌋+2\lfloor\log_{2}(1/\tau)\rfloor+2 bits of the queried function on the batch. So, the values used by the computation component will be within τ\tau of the average values of the queried functions on the appropriate batches, as desired. Also, any component that was ever used to query the function 11 will always return a nonzero value, so the computation component will be able to track the current timestep correctly.

We claim that in every timestep the net will output −2ρ-2\rho or 1+2ρ1+2\rho as determined by the computation component, the query subgraphs will update so that the computation component can read the results of the desired queries from them, and none of the edges will change in weight except the appropriate edges in copies of QQ or Q′Q^{\prime} and possibly weights of edges from the output control vertices used in this timestep to the output, and we can prove this by induction on the timesteps.

So, assume that this has held so far. In the current timestep all of the output control vertices except those designated for this timestep are set to , and the ones designated for this timestep are set to the value chosen by the computation component. The weights of the edges from the current output control vertices still have their original values because the only way they could not is if they had been set to nonzero values before. The edges from the other output control vertices to the output have no effect on its value so the primary output control vertices as a whole make a contribution of ±(1+2ρ)\pm(1+2\rho) as chosen by the computation component to the output. Meanwhile, there at most (plog⁡2(1/τ)+3)(p\log_{2}(1/\tau)+3) copies of QQ or Q′Q^{\prime} that have vcv_{c} set to for any sample, so the computation component can compute the contribution they make to the output and use the secondary output control vertices for the timestep to cancel it out. So, the input to the output vertex will be exactly what we wanted it to be, which means that the copies of QQ and Q′Q^{\prime} will update in the manner given by Lemma 11. That in turn means that the query subgraphs will update in the desired manner and the computation component will be able to determine valid values for the queries the bSQTM0/1{{\mathsf{bSQ_{TM}^{0/1}}}} algorithm would have made. None of the edges in or to the computation component will change by Lemma 10, and none of the edges from the computation component to any of the copies of vcv_{c} or any output control vertices will change because they are always at flat parts of their activation functions. So, the net behaves as described. That means that the net continues to be able to make queries, perform arbitrary efficient computations on the results of those queries, and output the result of an arbitrary efficient computation.

Once it is done training the compuation component can use the final output control vertex to make the net output or 11 based on the current input and the valuesof the previous queries. This process takes kk steps to run, and uses a polynomial number of query subgraphs per step. Any Turing machine can be converted into a circuit with size polynomial in the number of steps it runs for, so this can all be done by a net of size polynomial in the parameters. ∎

Before proving Theorems 1a, 1b, 1c and 1d, we first prove Lemma 4, restated below for convenience.

Using the tools developed in Sections 4 and 5 we now prove Theorems 1a, 1b, 1c and 1d. We only show the proofs for the computationally unbounded case, as a near identical derivation achieves the required results for the computationally bounded case, since all relevant Theorems/Lemmas have both versions (except Lemmas 3a and 3b which are stated as different lemmas).

From Lemma 1, we have that for all bb and τ<1/(2b)\tau<1/(2b), and k=O(mn/δ)k=O(mn/\delta), p=n+1p=n+1 and r′=r+klog⁡2br^{\prime}=r+k\log_{2}b, it holds that

From Lemma 3a (correspondingly Lemma 3b for the computationally bounded case), we then have that for T=k=O(mn/δ)T=k=O(mn/\delta), ρ=τ/4<1/(8b)\rho=\tau/4<1/(8b), p′=r′+(p+1)k=r+O((n+log⁡b)mn/δ)p^{\prime}=r^{\prime}+(p+1)k=r+O((n+\log b)mn/\delta), it holds that

The proof is complete by combining the above. ∎

Any bSGD(T,ρ,b,p,r){\mathsf{bSGD}}(T,\rho,b,p,r) algorithm can be simulated by a PAC(m,r){\mathsf{PAC}}(m,r) algorithm with m=Tbm=Tb samples, since bb samples are required to perform one bSGD{\mathsf{bSGD}} iteration. Moreover, there is no loss in the error ensured by this simulation. ∎

From Lemma 4 it holds for all T,ρ,b,p,rT,\rho,b,p,r and k=Tk=T, τ=ρ/4\tau=\rho/4 that

And by Theorem 2c, there exists a constant CC such that for k′=kp=Tpk^{\prime}=kp=Tp and τ′=τ2=ρ8\tau^{\prime}=\frac{\tau}{2}=\frac{\rho}{8} it holds for all bb such that bτ2>Clog⁡(kp/δ)b\tau^{2}>C\log(kp/\delta) that

The proof is complete by combining the above. ∎

From Lemma 2, it holds for all bb, τ′=τ4\tau^{\prime}=\frac{\tau}{4}, k′=k⋅⌈Clog⁡(k/δ)bτ2⌉k^{\prime}=k\cdot\left\lceil\frac{C\log(k/\delta)}{b\tau^{2}}\right\rceil and p=1p=1 that

Finally, from Lemma 3a (correspondingly Lemma 3b for the computationally bounded case) we have that for T=kT=k, ρ=τ′4\rho=\frac{\tau^{\prime}}{4} and p′=r+(p+1)k′=r+2k′p^{\prime}=r+(p+1)k^{\prime}=r+2k^{\prime} that

The proof is complete by combining the above. ∎

Before proving Theorems 3a, 3b, 3c and 3d, we state the analogs of Lemmas 3a, 3b and 4 relating fbGD{\mathsf{fbGD}} and fbSQ{{\mathsf{fbSQ}}} (fbSQ0/1{{\mathsf{fbSQ^{0/1}}}}) (the proofs follow in an identical manner, so we skip it).

(fbGD{\mathsf{fbGD}} to fbSQ{{\mathsf{fbSQ}}}) For all T,ρ,m,p,rT,\rho,m,p,r, it holds that

Finally, we put together all the tools developed in Sections 4 and 5 along with the above Lemmas to prove Theorems 3a, 3b, 3c and 3d. As in Appendix D, we only show the proofs for the computationally unbounded case, as a near identical derivation achieves the required results for the computationally bounded case.

From Lemma 6, we have that for all mm, τ<1/(2m)\tau<1/(2m) and rr, it holds for k=2m(n+1)k=2m(n+1), p=1p=1, r′=rr^{\prime}=r that

From Lemma 12a (correspondingly Lemma 12b for the computationally bounded case), we then have that for T=k=O(mn)T=k=O(mn), ρ=τ/4<1/(8m)\rho=\tau/4<1/(8m), p′=r′+(p+1)k=r+O(mn)p^{\prime}=r^{\prime}+(p+1)k=r+O(mn), it holds that

The proof is complete by combining the above. ∎

Any fbGD(T,ρ,m,p,r){\mathsf{fbGD}}(T,\rho,m,p,r) algorithm can be simulated by a PAC(m,r){\mathsf{PAC}}(m,r) algorithm with mm samples. Moreover, there is no loss in the error ensured by this simulation. ∎

From Lemma 13 it holds for all T,ρ,m,p,rT,\rho,m,p,r and k=Tk=T, τ=ρ/4\tau=\rho/4 that

And by Lemma 5c, there exists a constant CC such that for k′=kp=Tpk^{\prime}=kp=Tp and τ′=τ2=ρ8\tau^{\prime}=\frac{\tau}{2}=\frac{\rho}{8} it holds for all mm such that mτ2>C(kplog⁡(1/τ)+log⁡(1/δ))m\tau^{2}>C(kp\log(1/\tau)+\log(1/\delta)) that

The proof is complete by combining the above. ∎

From Lemma 7, there exists a constant CC such that for all mm, τ′=τ4\tau^{\prime}=\frac{\tau}{4} satisfying mτ2≥C(klog⁡(1/τ)+log⁡(1/δ))m\tau^{2}\geq C(k\log(1/\tau)+\log(1/\delta)) it holds for k′=2kk^{\prime}=2k and p=1p=1 that

Finally, from Lemma 12a (correspondingly Lemma 12b for the computationally bounded case) we have that for T=k′T=k^{\prime}, ρ=τ′4\rho=\frac{\tau^{\prime}}{4} and p′=r+(p+1)k′=r+2k′p^{\prime}=r+(p+1)k^{\prime}=r+2k^{\prime} that

The proof is complete by combining the above. ∎