The Neural Covariance SDE: Shaped Infinite Depth-and-Width Networks at Initialization

Mufan Bill Li, Mihai Nica, Daniel M. Roy

Introduction

Of the many milestones in deep learning theory, the precise characterization of the infinite-width limit of neural networks at initialization as a Gaussian process with a non-random covariance matrix was a turning point. The so-called Neural Network Gaussian process (NNGP) theory laid the mathematical foundation to study various limiting training dynamics under gradient descent . The Neural Tangent Kernel (NTK) limit formed the foundation for a rush of theoretical work, including advances in our understanding of generalization for wide networks . Besides the NTK limit, the infinite-width mean-field limit was developed , where the different parameterization demonstrates benefits for feature learning and hyperparameter tuning .

Fundamentally, the infinite-width paradigm derives results from the assumption that the depth of the network is held fixed while the widths of all layers grow to infinity. Unfortunately, this assumption can be problematic for modeling real-world networks, as the microscopic fluctuations from layer to layer are neglected in this limit (see Figure 1). In particular, infinite-width predictions are shown to be poor approximations of real networks unless the depth is much less than the width .

Impressive achievements of deep networks with billions of parameters crystallize the importance of understanding extremely large, deep neural networks (DNNs). An alternative to the infinite-width paradigm is the infinite-depth-and-width paradigm. In this setting, both the network depth dd and the width nn of each layer are simultaneously scaled to infinity, while their relative ratio d/nd/n remains fixed . Recent work also explores using d/nd/n as an effective perturbation parameter or to study concentration bounds in terms of d/nd/n . This limit has the distinct advantage of being incredibly accurate at predicting the output distribution for finite size networks at initialization — a significant improvement over the NNGP theory. Furthermore, it has also been shown that there is feature learning in this limit , in contrast to the linear regime of infinite-width limits . Considering the mathematical success of the NNGP techniques, the infinite-depth-and-width limit hints at the possibility of developing an accurate theory for training and generalization.

An immediate issue of the infinite-depth limit is that this limit predicts that network output becomes degenerate as depth increases: on initialization the network becomes a constant function sending all inputs to the same (random) output . While degenerate outputs are not necessarily an issue in theory, it poses a more serious problem in practice: degenerate correlations imply a “sharp” input–output Jacobian, and therefore exploding gradients . Intuitively, the output is not very sensitive to changes in the input, hence the gradient must be very large in the earlier layers.

A promising new attack on this problem is to modify the activation function (“shaping”) to reduce to the effect of degeneracy . In this prior work, extensive experiments show that shaping the activation significantly improves training speed without the need for normalization layers. This method has been proven effective for problems as large as standard ResNets on ImageNet data. The authors designed several criteria including reducing estimated output correlation, and numerically optimized the shape of activation functions for improved training results. However, their deterministic estimation of output correlation using the infinite-width limit leads to a poor approximation of real networks, as the additional randomness has both non-zero mean and heavy skew (see Figure 1 right column). Furthermore, numerically searching for the activation shape obscures the picture on how shaping should depend on the network depth and width.

In this paper, we address these problems by providing a precise theory of shaped infinite-depth-and-width networks, extending both the NNGP theories and the activation shaping techniques. In particular, we prescribe an exact scaling of the activation function shape as a function of network width nn that leads to a non-trivial nonlinear limit. By keeping track of microscopic O(n−1/2)O(n^{-1/2}) random fluctuations in each layer of the network, we show that the cumulative effect is described by a stochastic differential equation (SDE) in the limit. In contrast to existing infinite-width theory, we are able to characterize the random distribution of the output covariance, which matches closely to simulations of real networks. In a similar spirit to how the NNGP theory laid the foundation for studying training and generalization in the infinite-width limit, we also see this work as building the mathematical tools for an infinite-depth-and-width theory of training and generalization.

Similar to the NNGP approach, we use the fact that the output is Gaussian conditional on the penultimate layer. However, unlike in the infinite-width paradigm, the covariance matrix is no longer deterministic in the infinite-depth-and-width limit. Our focus in this paper is to study this random covariance matrix. Our main contributions are as follows:

We introduce the tool of stochastic n\sqrt{n}-expansions and convergence to SDEs for analyzing the distribution of covariances in DNNs.

For unshaped ReLU-like activations, we show that the norm of each layer evolves according to geometric Brownian motion and correlations evolve according to a discrete Markov process. See left column of Figure 1 and Section 2.

For both ReLU-like and a large class of smooth activation functions, we derive the Neural Covariance SDE characterizing the distribution of the shaped infinite-depth-and-width limit. See right column of Figure 1 and Section 3.

We show our prescribed shape scaling is exact, as other rates of scaling leads to either degenerate or linear network limits. See 3.4 and 3.10.

For smooth activations, we derive an if-and-only-if condition for exploding/vanishing norms based on properties of the activation function. See 3.7 and Section 4.

We provide simulations to verify theoretical predictions and help interpret properties of real DNNs. See Figures 1 and 4 and supplemental simulations in Appendix F.

Limits for Unshaped ReLU-Like Activations

In this section, we analyze ReLU-like activations by which we mean activations which are linear on the negative and positive numbers given respectively by two slopes s+s_{+} and s−s_{-}:

where the convergence is in the Skorohod topology (see Appendix A). When φ\varphi is the ReLU function (s+=1,s−=0s_{+}=1,s_{-}=0), we have c=2c=2 and σ2=5\sigma^{2}=5, which recovers known results in . We remark again this simple Markov chain example illustrates the main technique we use in later sections to establish SDE convergence for shaped networks in Section 3.

Neural Covariance SDEs: Shaped Infinite-Depth-and-Width Limit

We will show that with shaping of Definition 3.1, one gets non-trivial SDEs that describe the covariance (3.2) and correlations (3.3) of the network. The precise scaling is shown to be the critical scaling for a non-trivial limit in 3.4. All proofs for results in this section appear in Appendix C.

where ν(ρ)≔(c+−c−)22π(1−ρ2−ρarccos⁡ρ),ρtαβ≔VtαβVtααVtββ\nu(\rho)\coloneqq\frac{(c_{+}-c_{-})^{2}}{2\pi}\left(\sqrt{1-\rho^{2}}-\rho\arccos\rho\right),\rho^{\alpha\beta}_{t}\coloneqq\frac{V^{\alpha\beta}_{t}}{\sqrt{V^{\alpha\alpha}_{t}V^{\beta\beta}_{t}}}

Furthermore, the output distribution can be described conditional on VTV_{T} evaluated at final time TT

Here we remark that ν(1)=0\nu(1)=0, and therefore the drift component of diagonal entries (VtααV^{\alpha\alpha}_{t}) are zero, as they are geometric Brownian motion. However, we emphasize that the mm-point joint output distribution is not characterized by the marginal for each of the pairs, as the output zoutαz_{\text{out}}^{\alpha} is not Gaussian. In particular, we observe the diffusion matrix entry corresponding to Vtαβ,VtγδV^{\alpha\beta}_{t},V^{\gamma\delta}_{t} involves other processes Vtαγ,Vtβδ,Vtαδ,VtβγV^{\alpha\gamma}_{t},V^{\beta\delta}_{t},V^{\alpha\delta}_{t},V^{\beta\gamma}_{t}! This implies that the Neural Covariance SDE limit cannot be described by a kernel, unlike stacking random features or NNGP.

That being said, it is still instructive to study the marginal for a pair of data points. More specifically, it turns out in the generalized ReLU case, we can derive the marginal SDE for the correlation process.

To help interpret the SDE, we observe that μ\mu and σ\sigma are entirely independent of the activation function. In other words, these terms will be present in this limit even for linear networks. At the same time, ν\nu describes the influence of the shaped activation function in this limit. has derived a related ordinary differential equation (ODE) of dρt=ν(ρt) dtd\rho_{t}=\nu(\rho_{t})\,dt in the sequential limit of n→∞n\to\infty then d→∞d\to\infty, where the activation is shaped depending on depth. Here we also note that ν(ρ)\nu(\rho) is closely related to the J1J_{1} function derived in . See Section C.3 for the mm-point joint version of the correlation SDE, and Appendix F for an empirical measure of convergence in the Kolmogorov–Smirnov distance.

It is also possible to transform this SDE via Itô’s Lemma for potentially more interpretability, such as the angle form θtαβ=arccos⁡(ρtαβ)\theta^{\alpha\beta}_{t}=\arccos(\rho^{\alpha\beta}_{t})

where for θ≈0\theta\approx 0 we have that θcot⁡θ−1≈−θ23\theta\cot\theta-1\approx\frac{-\theta^{2}}{3} and sin⁡θ≈θ\sin\theta\approx\theta, which converges rapidly to .

One immediate consequence of the correlation SDE is that we can show the n−1/2n^{-1/2} scaling in Definition 3.1 is the only case where the limit is neither degenerate nor a linear network.

the degenerate limit: ρtαβ=1\rho^{\alpha\beta}_{t}=1 for all t>0t>0, if 0≤p<120\leq p<\frac{1}{2}, and c+≠c−c_{+}\neq c_{-},

the critical limit: the SDE from 3.3, if p=12p=\frac{1}{2},

the linear network limit: if p>12p>\frac{1}{2} , the following SDE, with μ,σ\mu,\sigma as defined in (3.5),

Here we remark that the unshaped network case (p=0p=0) is contained by the above in case (i). At the same time, we observe that case (iii) is equivalent to the correlation SDE in 3.3 except with ν=0\nu=0. In particular, we observe this limit is also reached when c+=c−c_{+}=c_{-}, which implies φs(x)=s+x\varphi_{s}(x)=s_{+}x is linear, which is the reason we call this the linear network limit. Furthermore, without much additional work, the same argument also implies the joint covariance SDE also loses the drift component, i.e., dVt=Σ(Vt)1/2 dBtdV_{t}=\Sigma(V_{t})^{1/2}\,dB_{t}.

2 Neural Covariance SDE for Shaped Smooth Activations

In this section, we consider smooth activation functions and derive a similar covariance SDE. All the proofs for results in this section can be found in Appendix D.

Following the ideas of , we consider the following shaping of a smooth activation function.

In this regime, we can similarly characterize the joint output distribution, however the limiting SDEs are not always well behaved. In particular, they can have finite time explosions as described by the Feller test for explosions [42, Theorem 5.5.29]. Here the SDE in 3.7 is exactly the VtααV^{\alpha\alpha}_{t} marginal of the Neural Covariance SDE, with the parameter bb determined by the activation function φ\varphi and controls whether or not finite time explosions happen (see Equation 4.1).

Technically speaking, the main culprit behind finite time explosions is the non-Lipschitzness of the drift coefficient. This issue requires us to weaken the sense of convergence in this section; the ordinary convergence in the Skorohod topology is in general not true when the diffusion has finite time explosions. A weakened type of convergence is the best we can hope for. To this goal, we introduce the following definition.

We say a sequence of processes XnX^{n} converge locally to XX in the Skorohod topology if for any r>0r>0, we define the following stopping times

and we have that Xt∧τnnX^{n}_{t\wedge\tau^{n}} converge to Xt∧τX_{t\wedge\tau} in the Skorohod topology.

This weakened sense of convergence essentially constrains the processes Xn,XX^{n},X in a bounded set by adding an absorbing boundary condition. Not only do these stopping times rule out explosions, the drift coefficient is now also Lipschitz on a compact set. With this notion of convergence, we can now state a precise Neural Covariance SDE result for general smooth activation functions.

where Σ(Vt)\Sigma(V_{t}) is the same as 3.2 and

Furthermore, if VTV_{T} is finite, then the output distribution can be described conditional on VTV_{T} as

and otherwise the distribution of [zoutα]α=1m[z_{\text{out}}^{\alpha}]_{\alpha=1}^{m} is undefined.

We also have a similar critical scaling result for general smooth activations.

the degenerate limit: if 0<p<120<p<\frac{1}{2}

for all t>0t>0 and 1≤α≤β≤m1\leq\alpha\leq\beta\leq m,

the critical limit: the solution of the SDE from 3.9, if p=12p=\frac{1}{2},

the linear network limit: the stopped solution to the SDE dVt=Σ(Vt) dBtdV_{t}=\Sigma(V_{t})\,dB_{t} with coefficient Σ\Sigma defined in 3.3, if p>12p>\frac{1}{2}.

Here we observe that in case (i) when 34φ′′(0)2+φ′′′(0)≤0\frac{3}{4}\varphi^{\prime\prime}(0)^{2}+\varphi^{\prime\prime\prime}(0)\leq 0, we also have a constant (in time) correlation ρtαβ\rho^{\alpha\beta}_{t} similar to the ReLU case in 3.4, however in this case ρtαβ\rho^{\alpha\beta}_{t} is not necessarily equal to 11. At the same time, the linear network limit in case (iii) also has the same covariance SDE as 3.4.

Consequences, Discussion, and Future Directions

So far, we have derived the Neural Covariance SDE. Analysis of this SDE reveals important behaviour of the network on initialization. Here we lay out one concrete example and provide some discussion and future directions.

Exploding and Vanishing Norms. Here we consider the behaviour of shaping smooth activation functions, as it is done in the experiments of . While the authors here avoided exploding and vanishing norms by numerically optimizing shaping parameters, we can actually describe the precise behaviour a priori with the Neural Covariance SDE. Recall the shaping parameter aa from Definition 3.6. Let VtV_{t} be the solution to the SDE in Equation 3.9. We can write down the marginal SDE for VtααV^{\alpha\alpha}_{t} as

which implies by 3.7 that VtV_{t} has a finite time explosion (with non-zero probability) if and only if 34φ′′(0)2+φ′′′(0)>0\frac{3}{4}\varphi^{\prime\prime}(0)^{2}+\varphi^{\prime\prime\prime}(0)>0. This criterion can be used to help choose how activation functions should be centered for shaping; below are two examples.

We start with the sigmoid activation σ(x)=11+e−x\sigma(x)=\frac{1}{1+e^{-x}}, then we can define φ(x)≔4σ(x)−2\varphi(x)\coloneqq 4\sigma(x)-2 to satisfy Assumption 3.5, which leads to φ′′(0)=0,φ′′′(0)=−12\varphi^{\prime\prime}(0)=0,\varphi^{\prime\prime\prime}(0)=-\frac{1}{2}, and therefore leads to a stable network. It turns out φ(x)≔tanh⁡(x)\varphi(x)\coloneqq\tanh(x) already satisfies Assumption 3.5, which leads to φ′′(0)=0,φ′′′(0)=−2\varphi^{\prime\prime}(0)=0,\varphi^{\prime\prime\prime}(0)=-2, and therefore is also stable.

More generally, if σ\sigma behaves like a cumulative distribution function for a symmetric unimodal density, we will have that φ′′(0)=0\varphi^{\prime\prime}(0)=0 and φ′′′(0)<0\varphi^{\prime\prime\prime}(0)<0 as desired.

Relationship to Edge of Chaos. The finite time explosion example above resembles the Edge of Chaos (EOC) analysis of gradient stability , where the weight and bias variance at initialization determines a stability criterion. However, we note that the EOC regime is sufficiently different that the results are not directly comparable. More precisely, the EOC analysis is in the sequential limit of infinite-width and then infinite-depth, which also leaves the activation function unchanged. Under very weak assumptions, the variance (diagonal of VtV_{t}) will not explode in this regime; instead, the gradient can explode due to the covariance (off diagonals). On the other hand, our finite explosion result is in the joint limit of depth and width, where the variance (diagonal of VtV_{t}) can explode instead.

Simulating SDEs. Both the Markov chains and SDEs predict neural networks at initialization very well (see Figure 1), but the SDE is significantly faster to simulate. In particular, we can view the Markov chain as an approximate Euler discretization of the SDE, but with a very small step size n−1n^{-1}. In contrast, to simulate the SDE we should only need a step size that is small on the scale of depth-to-width ratio T=d/nT=d/n, which is independent of width nn. Therefore, practitioners using the shaping techniques of can now simulate the covariance SDEs at a low computational cost to significantly improve estimates of the output correlation (see Figure 1 and additional simulations in Appendix F).

Analytical Tractability of SDEs. Besides numerical tractability, the SDEs are also far more tractable to analyze. For example, in the one input case, we arrive at geometric Brownian motion Equation 2.6, which is known to have a log-normal distribution at fixed times. Similarly, our finite time explosions hinge on the fact we identified an SDE limit. In the same way that NNGP theory played a major role in the infinite-width regime, the Neural Covariance SDEs and the techniques developed here also serve as a mathematical foundation for studying training and generalization.

Acknowledgement

We would like to thank Sinho Chewi, James Foster, Boris Hanin, Cameron Jakub, Jeffrey Negrea, Nuri Mert Vural, Guodong Zhang, Matthew S. Zhang, and Yuchong Zhang for helpful discussions and draft feedback. We would like to thank Sam Buchanan and Soufiane Hayou for pointing out a gap in the proof of B.8. ML is supported by Ontario Graduate Scholarship and the Vector Institute. MN is supported by an NSERC Discovery Grant. DMR is supported in part by Canada CIFAR AI Chair funding through the Vector Institute, an NSERC Discovery Grant, Ontario Early Researcher Award, a stipend provided by the Charles Simonyi Endowment, and a New Frontiers in Research Exploration Grant.

References

Appendix A Background on Markov Chain Convergence to SDEs

In this section we briefly review the background and technical results required to characterize the convergence of a Markov chain to an SDE. Majority of the content in this section are based on .

T\mathcal{T} induces the Skorohod convergence xn→sxx_{n}\xrightarrow{s}x,

T\mathcal{T} generates the Borel σ\sigma-field generated by the evaluation maps πt\pi_{t}, t≥0t\geq 0, where πt(x)=xt\pi_{t}(x)=x_{t}.

We also need to define Feller semi-groups. To start we let SS be a locally compact separable metric space and C0≔C0(S)C_{0}\coloneqq C_{0}(S) be the space of continuous functions that vanishes at infinity, and we equip C0C_{0} with the sup norm to make it a Banach space. T:C0→C0T:C_{0}\to C_{0} is a positive contraction operator if for all 0≤f≤10\leq f\leq 1 we have 0≤Tf≤10\leq Tf\leq 1. A semi-group of such operators (Tt)(T_{t}) on C0C_{0} is called a Feller semi-group if it additionally satisfies

Let D⊂C0\mathcal{D}\subset C_{0} and A:D→C0A:\mathcal{D}\to C_{0}, and we say that (A,D)(A,\mathcal{D}) is a generator of (Tt)(T_{t}) if D\mathcal{D} is the maximal set such that for all f∈Df\in\mathcal{D}, we have that

An operator AA with domain D\mathcal{D} on a Banach space BB is said to be closed, if its graph G={(f,Af)∣f∈D}G=\{(f,Af)|f\in\mathcal{D}\} is a closed subset of B×BB\times B. If the closure of GG is the graph of an operator Aˉ\bar{A}, we say Aˉ\bar{A} is the closure of AA. Finally, we will define a linear subspace D⊂DD\subset\mathcal{D} as a core of AA if the closure of A∣DA|_{D} is AA. If (A,D)(A,\mathcal{D}) is a generator of a Feller semigroup, every dense invariant subspace D⊂DD\subset{D} is a core of AA [48, Proposition 17.9]. In particular, we will work with the core C0∞C^{\infty}_{0} of smooth functions vanishing at infinity.

We will state a sufficient condition required for an semi-group to be Feller based on its generator.

generates a Feller semi-group on C0C_{0}.

We will next state a set of equivalent criterion for convergence of Feller processes.

Let X,X1,X2,X3,⋯X,X^{1},X^{2},X^{3},\cdots be Feller processes in SS with semi-groups (Tt),(Tn,t)(T_{t}),(T_{n,t}) and generators (A,D),(An,Dn)(A,\mathcal{D}),(A_{n},\mathcal{D}_{n}), respectively, and fix a core DD for AA. Then these conditions are equivalent:

for any f∈Df\in D, there exists some fn∈Dnf_{n}\in\mathcal{D}_{n} with fn→ff_{n}\to f and Anfn→AfA_{n}f_{n}\to Af,

Tn,t→TtT_{n,t}\to T_{t} strongly for each t>0t>0,

Tn,tf→TtfT_{n,t}f\to T_{t}f for every f∈C0f\in C_{0}, uniformly for bounded t>0t>0,

Once again, we note that it is common to choose the core D=C0∞D=C^{\infty}_{0}, and that checking condition (i) is sufficient for convergence in the Skorohod topology. This is translated to the Markov chain setting by the next theorem.

Let Y1,Y2,Y3,⋯Y^{1},Y^{2},Y^{3},\cdots be discrete time Markov chains in SS with transition operators U1,U2,U3,⋯U_{1},U_{2},U_{3},\cdots, and let XX be a Feller process with semi-group (Tt)(T_{t}) and generator AA. Fix a core DD for AA, and let 0<hn→00<h_{n}\to 0. Then conditions (i)−(iv)(i)-(iv) of A.3 remain equivalent for the operators and processes

It remains to check that the generators AnA_{n} converges to AA with respect to the core D=C0∞D=C^{\infty}_{0}, and we will use a criterion from . Here we will first let Πn(x,dy)\Pi_{n}(x,dy) be the Markov transition kernel of YnY^{n}, and define

The following two conditions are equivalent:

Finally, we summarize the above results in a user friendly form for our applications.

Suppose otherwise bn,σnb_{n},\sigma_{n} are only locally Lipschitz (but still uniform in nn), then XnX^{n} converges locally to XX in the same topology (see Definition 3.8). More precisely, for any fixed r>0r>0, we consider the stopping times

then the stopped process Xt∧τnnX^{n}_{t\wedge\tau^{n}} converges in distribution to the stopped solution Xt∧τX_{t\wedge\tau} of the above SDE in the same topology.

We start with Δnϵ(x)\Delta^{\epsilon}_{n}(x). Given that the randomness in the Markov chain have bounded moments (uniform in nn), then by a Markov inequality we have that for any q>0q>0

therefore choosing q>1q>1 we have sup⁡∣x∣≤RΔnϵ(x)=O(n−2p(q−1))→0\sup_{|x|\leq R}\Delta^{\epsilon}_{n}(x)=O(n^{-2p(q-1)})\to 0 for any fixed ϵ\epsilon.

where we note the drift’s randomness contributes the higher order n−2pn^{-2p} term and therefore also vanishes in the limit. This implies sup⁡x≤∣R∣∥an(x)−a(x)∥op→0\sup_{x\leq|R|}\|a_{n}(x)-a(x)\|_{op}\to 0, which gives us the desired result.

Appendix B Unshaped ReLU Markov Chain

In this section, we will derive the Markov chain update Equation 2.10 with explicit coefficients. For the rest of this section, we will adopt the following notation. Let φ(x)≔max⁡(x,0)\varphi(x)\coloneqq\max(x,0) be the ReLU activation function. Let f(x)=12πe−x2/2f(x)=\frac{1}{\sqrt{2\pi}}e^{-x^{2}/2} be the density of a standard Gaussian, and let F(x)=∫−∞xf(t) dtF(x)=\int_{-\infty}^{x}f(t)\,dt be the cumulative distribution function (CDF).

For g∼N(0,1)g\sim\mathcal{N}(0,1) and hh is weakly differentiable, we have that

where ff is the standard Gaussian density.

We start by writing the expectation as an integral

Here by observing that f′(x)=−xf(x)f^{\prime}(x)=-xf(x), we can use integration by parts for u=h(x),dv=xf(x) dxu=h(x),dv=xf(x)\,dx to get du=h′(x) dx,v=−f(x)du=h^{\prime}(x)\,dx,v=-f(x), and therefore

Finally we recover the desired result using symmetry of f(−a)=f(a)f(-a)=f(a).

We will note the special case of a=0a=0 to get

Let g∼N(0,1),ρ∈,q=1−ρ2g\sim\mathcal{N}(0,1),\rho\in,q=\sqrt{1-\rho^{2}}, then we have that

We will again write the expectation as an integral

at this point, we can complete the square to write

Finally, we can use the substitution y=x+aρq,dy=1qdxy=\frac{x+a\rho}{q},dy=\frac{1}{q}dx to get

We will start by calculating simpler quantities.

For the second and fourth moments, we simply observe that g2g^{2} is symmetric and φ\varphi is exactly half of of the integral. For the first integral we will use Gaussian integration-by-parts with h(g)=1h(g)=1 to get

We will also recall the following result from

Let ρ∈,q=1−ρ2\rho\in,q=\sqrt{1-\rho^{2}} and let ρ,w∼N(0,1)\rho,w\sim\mathcal{N}(0,1) be independent. Then we have that

We will need to compute the following quantity.

Let ρ∈,q=1−ρ2\rho\in,q=\sqrt{1-\rho^{2}} and let ρ,w∼N(0,1)\rho,w\sim\mathcal{N}(0,1) be independent. Then we have that

At this point we can use the substitution formula from B.2 to write

which is the desired result after simplifying.

where φ(x):=max⁡(x,0)\varphi(x):=\max(x,0) is the usual ReLU activation.

We will compute several basic moments first.

Furthermore, this implies the normalizing constant is c=2s+2+s−2c=\frac{2}{s_{+}^{2}+s_{-}^{2}} and

To start we first recall the Gaussian integration by parts calculation

then the first moment follows immediately from rewriting in terms of φ\varphi

For the second moment, we will also rewrite in terms of φ\varphi

where we used that φ(g)φ(−g)=0\varphi(g)\varphi(-g)=0 almost surely and g=d−gg\overset{d}{=}-g, and the desire result follows from Gaussian integration by parts

For the fourth moment, we will similarly observe that all mixed moments φ(g)pφ(−g)r=0\varphi(g)^{p}\varphi(-g)^{r}=0 almost surely whenever p,r>0p,r>0, which allows us to write

and the desire result follows from the Gaussian integration by parts calculation

where g,w∼N(0,1)g,w\sim N(0,1) and we define g^=ρg+qw\hat{g}=\rho g+qw with q=1−ρ2q=\sqrt{1-\rho^{2}}. We will also use the short hand notation to write Jˉp:=Jˉp,p,Kp:=Kp,p\bar{J}_{p}:=\bar{J}_{p,p},K_{p}:=K_{p,p}.

Let ρ∈\rho\in, q=1−ρ2q=\sqrt{1-\rho^{2}}, g,w∼N(0,1)g,w\sim\mathcal{N}(0,1), and g^=ρg+qw\hat{g}=\rho g+qw. Then we have the following formulas

Before we start, we will make several observations. Using the fact that (g,w)=d(±g,±w)(g,w)\overset{d}{=}(\pm g,\pm w), we have the following equality in distribution relations

In particular, we note that the two Gaussian random variable (g,−g^)(g,-\hat{g}) have correlation −ρ-\rho.

With K2K_{2}, we will additionally make use of the fact that φ(g)φ(−g)=0\varphi(g)\varphi(-g)=0 almost surely to write

K3,1K_{3,1} follows from a similar calculation

We will also define the bounded Lipschitz function norm as

which induces the bounded Lipschitz distance for probability measures

We will also note that O(n−1)O(n^{-1}) error in the result arise from replacing the O(n−1/2)O(n^{-1/2}) with a Gaussian due to Berry–Esseen, and the O(n−1)O(n^{-1}) term with its expectation, as these are the dominant error terms in the approximation.

where we recall the X=O(n−3/2)X=O(n^{-3/2}) notation denotes a random variable (the Taylor remainder term) where all moments of n3/2Xn^{3/2}X are bounded by a constant independent of nn.

This allows us to write (considering the well defined case)

To complete the proof, we will need to control these differences in terms of the bounded Lipschitz distance on the Markov transition kernels. To this goal, we let hh be such that ∥h∥BL≤1\|h\|_{BL}\leq 1, hence it must be both bounded by 11 and at worst 11-Lipschitz. We will first condition on EcE^{c} to write the Taylor expansion, and then “uncondition” to recover the original distribution, both at a cost of an O(2−n)O(2^{-n}) error term. More precisely, we will write

At this point we observe that we can now “uncondition” the Taylor expansion by essentially doing the same trick, or more precisely observe that

Since hh is 11-Lipschitz, we have that h(x+y)≤h(x)+∣y∣h(x+y)\leq h(x)+|y|, and therefore we can write

Finally since the above results do not depend on the choice of the test function hh, so we have that

Appendix C Proofs for ReLU Shaping Results

where φ(x)≔max⁡(x,0)\varphi(x)\coloneqq\max(x,0) is the usual ReLU activation.

where g,wg,w are iid N(0,1)\mathcal{N}(0,1) and we define g^=ρg+qw\hat{g}=\rho g+qw with q=1−ρ2q=\sqrt{1-\rho^{2}}. We will also use the short hand notation to write Jˉp:=Jˉp,p,Kp:=Kp,p\bar{J}_{p}:=\bar{J}_{p,p},K_{p}:=K_{p,p}.

We will also recall from B.6 the following moment calculations

In the shaped case, we will calculate a Taylor expansion for the function cK1(ρ)cK_{1}(\rho).

Let s±=1+c±ns_{\pm}=1+\frac{c_{\pm}}{\sqrt{n}}, then

where ν(ρ)=(c+−c−)22π(1−ρ2+ρarccos⁡ρ)\nu(\rho)=\frac{(c_{+}-c_{-})^{2}}{2\pi}\left(\sqrt{1-\rho^{2}}+\rho\arccos\rho\right).

We start by consider plugging in the formula from Equation C.4 to get

where we used the fact that arccos⁡(−ρ)=π−arccos⁡(ρ)\arccos(-\rho)=\pi-\arccos(\rho).

After substituting s±=1+c±ns_{\pm}=1+\frac{c_{\pm}}{\sqrt{n}}, we can use SymPy to Taylor expand with respect to the variable x=n−1/2x=n^{-1/2} about x0=0x_{0}=0 and get

where we used the simplify function on the coefficients to reduce the size of the expression.

We will also need an approximation result for fourth moments.

and similarly for other pairs of α,β,γ,δ\alpha,\beta,\gamma,\delta. Then

where the constant in the O(⋅)O(\cdot) notation is universal.

We will also calculate a useful covariance.

and similarly for other pairs of α,β,γ,δ\alpha,\beta,\gamma,\delta. If we also define

then we have the following covariance formula:

We first observe that since each entry of the sum in RαβR^{\alpha\beta} are iid and zero mean, it is sufficient to just compute the covariance a single term. In other words

Since c=1+O(n−1/2)c=1+O(n^{-1/2}) and K1(ρ)=ρ+O(n−1)K_{1}(\rho)=\rho+O(n^{-1}) from C.1, we can further write this as

and we can use the fourth moment approximation C.2 to get

where we denote ν(ρ)≔(c+−c−)22π(1−ρ2−ρarccos⁡ρ),ρtαβ≔VtαβVtααVtββ\nu(\rho)\coloneqq\frac{(c_{+}-c_{-})^{2}}{2\pi}\left(\sqrt{1-\rho^{2}}-\rho\arccos\rho\right),\rho^{\alpha\beta}_{t}\coloneqq\frac{V^{\alpha\beta}_{t}}{\sqrt{V^{\alpha\alpha}_{t}V^{\beta\beta}_{t}}} and write

Furthermore, the output distribution can be described conditional on VTV_{T} evaluated at final time TT

which essentially recovers the Markov chain form we want from A.6, where the drift is

It remains to simply compute the covariance conditioned on previous layer. To this end, we will use C.3 to write

C.2 Proof of Theorem 3.3 (Correlation SDE, ReLU)

Using the expansion of cK1(ρ)cK_{1}(\rho) from C.1, we can now write

Furthermore, we also have that by C.1 and C.2

Finally, we can recover the desired SDE via A.6.

C.3 Joint Correlation SDE

In this section, we will extend 3.3 to a general joint process over all the possible pairs of correlations.

It’s sufficient to just compute the covariance matrix Σ\Sigma for the random terms of the Markov chain Equation B.42, which reduces down to

Using C.1 and C.3, we can calculate this explicitly as

C.4 Proof for Proposition 3.4 (Critical Exponent, ReLU)

the degenerate limit: ρtαβ=1\rho^{\alpha\beta}_{t}=1 for all t>0t>0, if 0≤p<120\leq p<\frac{1}{2}, and c+≠c−c_{+}\neq c_{-},

the critical limit: the SDE from 3.3, if p=12p=\frac{1}{2},

the linear network limit: if p>12p>\frac{1}{2} , the following SDE, with μ,σ\mu,\sigma as defined in (3.5),

Case (ii) follows from 3.3, therefore it is sufficient to only consider cases (i) and (iii). In the case that p=0p=0, we can recover the following recursion in the limit as n→∞n\to\infty

Next we will recall the result of C.1 and observe that we can simply replace n\sqrt{n} with npn^{p} to recover the expansion

This gives us the following Markov chain from the proof of 3.3

In the case that 0<p<1/20<p<1/2, we can consider the time step size hn=n−2ph_{n}=n^{-2p} instead of n−1n^{-1} and apply A.6, where we recover the ODE

but on the time scale of ρ^sαβ,n=ρ^⌊sn2p⌋αβ\hat{\rho}^{\alpha\beta,n}_{s}=\hat{\rho}^{\alpha\beta}_{{\left\lfloor sn^{2p}\right\rfloor}}. Converting it back to the time scale of ρtαβ,n=ρ⌊tn⌋αβ\rho^{\alpha\beta,n}_{t}=\rho^{\alpha\beta}_{{\left\lfloor tn\right\rfloor}} implies that we have

And since ν(ρ)>0\nu(\rho)>0 for all ρ<1\rho<1 and that ν(ρ)=C(1−ρ)3/2+O((1−ρ)5/2)\nu(\rho)=C(1-\rho)^{3/2}+O((1-\rho)^{5/2}) as ρ→1\rho\to 1, we have that ρ^∞αβ=1\hat{\rho}^{\alpha\beta}_{\infty}=1 as desired.

In the case p>12p>\frac{1}{2}, we have that since ν\nu is deterministic, we observe the drift term used in A.6 in the limit as n→∞n\to\infty is

which would simply recover the desired SDE with drift μ\mu only.

Appendix D Proofs for Smooth Shaping Results

In this section, we consider smooth activation functions φ\varphi satisfying Assumption 3.5, that is φ∈C4,φ(0)=0,φ′(0)=1\varphi\in C^{4},\varphi(0)=0,\varphi^{\prime}(0)=1, and that ∣φ(4)(x)∣≤C(1+∣x∣p)|\varphi^{(4)}(x)|\leq C(1+|x|^{p}) for some C,p>0C,p>0. We recall the shaping we consider for activations of this type is via the following definition for s>0s>0

so that lim⁡s→∞φs(x)=x\lim_{s\to\infty}\varphi_{s}(x)=x.

Before we start, we will calculate the behaviour of the normalizing constant cc up an error order of s−3s^{-3}.

Let φs\varphi_{s} be defined as above with φ\varphi satisfying Assumption 3.5. Then if g∼N(0,1)g\sim N(0,1), we have that

We will first Taylor expand φs(g)\varphi_{s}(g) about g=0g=0

where we note by Assumption 3.5 the remainder term is at most polynomial in gg.

where O(s−3)O(s^{-3}) is bounded due to Gaussians have all bounded moments.

Therefore, for s>0s>0 sufficiently small, we have the following expansion

where b=34φ′′(0)2+φ′′′(0)b=\frac{3}{4}\varphi^{\prime\prime}(0)^{2}+\varphi^{\prime\prime\prime}(0), which is the desired result.

where Σ(Vt)\Sigma(V_{t}) is the same as 3.2 and

Furthermore, if VTV_{T} is finite, then the output distribution can be described conditional on VTV_{T} as

and otherwise the distribution of [zoutα]α=1m[z_{\text{out}}^{\alpha}]_{\alpha=1}^{m} is undefined.

where R3(⋅)R_{3}(\cdot) is the Taylor remainder term, which has polynomial growth by Assumption 3.5.

By using the fact that φ(0)=0,φ′(0)=1\varphi(0)=0,\varphi^{\prime}(0)=1 and observing that the derivatives of φs\varphi_{s} satisfies φs(k)(0)=φ(k)(0)sk−1\varphi^{(k)}_{s}(0)=\frac{\varphi^{(k)}(0)}{s^{k-1}}, we can further write

Then we can compute the inner product with the same expansion as

and we will proceed by analyzing the product terms separately. We start with the terms of order O(s0)O(s^{0}) first, which are

For the first order terms, i.e., terms of order O(s−1)O(s^{-1}), we have the terms

We then turn our attention to the second order terms, i.e., terms of order O(s−2)O(s^{-2})

Since this term is order s−2=1a2ns^{-2}=\frac{1}{a^{2}n}, it can only contribute to the drift term, and in view of A.6, we only need to compute its mean. To this goal, we will simply invoke Isserlis’ Theorem and calculate

At this point, we have fully recovered the drift term, and we observe the covariance structure is the same as C.3 in the limit as n→∞n\to\infty. Therefore we can invoke A.6 to recover the desired SDE.

D.2 Proof of Proposition 3.10 (Critical Exponent, Smooth)

We will restate and prove the proposition.

the degenerate limit: if 0<p<120<p<\frac{1}{2}

for all t>0t>0 and 1≤α≤β≤m1\leq\alpha\leq\beta\leq m,

the critical limit: the solution of the SDE from 3.9, if p=12p=\frac{1}{2},

the linear network limit: the stopped solution to the SDE dVt=Σ(Vt) dBtdV_{t}=\Sigma(V_{t})\,dB_{t} with coefficient Σ\Sigma defined in 3.3, if p>12p>\frac{1}{2}.

Similar to the proof of 3.9, we will borrow the same notation and write down the Markov chain update and consider the time scale depending on the value of pp. In case (i) where 0<p<120<p<\frac{1}{2}, we will consider the time scale hn=1s2=1a2n2ph_{n}=\frac{1}{s^{2}}=\frac{1}{a^{2}n^{2p}} and observe that based on the Taylor expansion of φs\varphi_{s} about , we can write

In view of the time scale s−2s^{-2} for A.6, it is then only important to keep track of the expected value of the s−2s^{-2} terms and the covariance of the s−1s^{-1} terms. However, since there is no terms on the order of s−1s^{-1}, we essentially have

where we used the fact that c=1−bs2+O(s−3)c=1-\frac{b}{s^{2}}+O(s^{-3}) for b=φ′′′(0)+34φ′′(0)2b=\varphi^{\prime\prime\prime}(0)+\frac{3}{4}\varphi^{\prime\prime}(0)^{2} from D.1.

Hence, we have that Utαα,n≔V⌊ts2⌋ααU^{\alpha\alpha,n}_{t}\coloneqq V^{\alpha\alpha}_{{\left\lfloor ts^{2}\right\rfloor}} converging to the ODE via A.6

where we observe if b>0b>0 this ODE is “mean avoiding” as it will drift towards or ∞\infty. And since the VtV_{t} time scale is on the order of 1n\frac{1}{n}, for all t>0t>0 we have that

therefore if b>0b>0 we have that Vtαα=0V^{\alpha\alpha}_{t}=0 or ∞\infty as desired in the first case of (i). When b=0b=0 we observe that Vtαα=V0ααV^{\alpha\alpha}_{t}=V^{\alpha\alpha}_{0} since the time derivative is zero. Furthermore if b<0b<0 we also have that Vtαα=1V^{\alpha\alpha}_{t}=1 in the second case of (i).

When b≤0b\leq 0, we can also write down the ODE for UtαβU^{\alpha\beta}_{t} using a similar argument and keeping only the s−2s^{-2} terms. More precisely, we can modify Equation D.18 to get

Since Utαα,UtββU^{\alpha\alpha}_{t},U^{\beta\beta}_{t} converge to constants as t→∞t\to\infty, ∣Utαβ∣≤UtααUtββ|U^{\alpha\beta}_{t}|\leq\sqrt{U^{\alpha\alpha}_{t}U^{\beta\beta}_{t}} by definition and Cauchy–Schwarz inequality, and that UtαβU^{\alpha\beta}_{t} satisfies a first order ODE (so it cannot have a periodic solution), we must also have that lim⁡t→∞Utαβ=const.\lim_{t\to\infty}U^{\alpha\beta}_{t}=\text{const.} This completes the proof for case (i).

Case (ii) follows directly from 3.9, therefore we can then consider case (iii) with the same Taylor expansion, however this time on the time scale of n−1n^{-1} instead. We will again follow A.6 to only track the mean of the order n−1n^{-1} term and the variance of the n−1/2n^{-1/2} term. Since p>12p>\frac{1}{2}, the only term that remains is the diffusion on the order of n−1/2n^{-1/2}

which gives us the desired SDE from calculating the covariance from 3.9.

D.3 Proof of Proposition 3.7 (Finite Time Explosion Criterion)

We will start by recalling several definitions from [42, Section 5.5]. Firstly, we consider the one dimensional Itô diffusion on I≔(0,∞)I\coloneqq(0,\infty)

where the drift and diffusion coefficients satisfy the following conditions

We will also define the following functions for some fixed c∈Ic\in I

We will also define the following sequence of stopping times for M>0M>0

and let τ∗≔sup⁡M>0τM\tau^{\ast}\coloneqq\sup_{M>0}\tau_{M}. Now we will state the main results we need for finite time explosions.

We will begin our derivations for the SDE Equation D.27.

Let XtX_{t} be a solution to the following SDE

then we have that τ∗=∞\tau^{\ast}=\infty a.s.

By Feller’s test for explosions D.5, we have the desired result.

Suppose XtX_{t} is a solution of the following equation

In particular, when b≤−1b\leq-1, we have that lim⁡x→0v(x)=lim⁡x→∞v(x)=∞\lim_{x\to 0}v(x)=\lim_{x\to\infty}v(x)=\infty.

Then we can also calculate the integral via a substitution of y=bxy=bx to get the desired result.

where γ\gamma is the lower incomplete gamma function, and therefore finite for all values of xx including the limits x→0,∞x\to 0,\infty.

The b=0b=0 case follows from D.6. Finally when b<0b<0 we can write

which clearly diverges to ∞\infty as x→∞x\to\infty.

On the other hand, we can observe as that as x→0x\to 0, we have that y∈y\in and therefore 1≤e∣b∣y≤e∣b∣1\leq e^{|b|y}\leq e^{|b|}. This implies we only need to consider the integral −∫x1y−∣b∣ dy-\int_{x}^{1}y^{-|b|}\,dy, which diverges to −∞-\infty if and only if ∣b∣≥1|b|\geq 1. In other words we have

Suppose XtX_{t} is a solution of the following equation

We will start by calculating the following integral using the exponential series expansion

We first consider the case when x→0x\to 0, in which case we have e−∣b∣≤e−by≤e∣b∣e^{-|b|}\leq e^{-by}\leq e^{|b|} and therefore will not affect convergence or divergence, so we can safely ignore the factor e−bye^{-by} and write (for k>0,x→0k>0,x\to 0)

Since the exponential series ∑k>0bkk!=eb−1\sum_{k>0}\frac{b^{k}}{k!}=e^{b}-1 converges, and we have terms strictly smaller than the exponential series, we have convergence of these terms when k>0k>0. We now return to handle a couple of edge case terms, firstly when k=0k=0

which is a desired behaviour. Secondly we consider when k=b+1k=b+1

from which we can conclude lim⁡x→0v(x)=∞\lim_{x\to 0}v(x)=\infty.

Next we consider the case when x→∞x\to\infty. Firstly, since we already have that p(x)→∞p(x)\to\infty when b≤0b\leq 0, therefore D.4 implies v(x)→∞v(x)\to\infty. Therefore we only need to consider when b>0b>0.

Since b>0b>0 we will have that e−bxe^{-bx} will dominate, and therefore we can safely ignore all the edge case terms and consider the series

Observe that as x→∞x\to\infty we actually recover the gamma integral in the terms i.e.

where we observe the second term is independent of kk, and therefore the series converges due to comparison with the exponential Taylor series. This implies we only need to focus on the first term, which is

where the series converges since it’s a sum of k−2k^{-2} type. This allows us to conclude that lim⁡x→∞v(x)<∞\lim_{x\to\infty}v(x)<\infty as desired.

We can now prove the desired result of 3.7, which we restate below.

Putting the results of D.7 and D.8 together, we have the following table

We will extend the above Lemma to a slightly modified update as well.

We will start the induction proof at n=1n=1

Then we assume the inequality holds for xnx_{n}, we will similarly write

and plugging in the inequality for xnx_{n} we get

To complete the proof it’s sufficient to show

Since 2n3+1≥2(n+1)3\frac{2n}{3}+1\geq\frac{2(n+1)}{3}, we only need to compare the first coefficient, which is

and this is equivalent to n≥1n\geq 1, and therefore satisfied by the induction. This completes the proof.

At the same time, we also conjecture the following bound.

Suppose we want to establish the approximation of

Then for the initial induction n=1n=1 step, we only need

Using WolframAlpha (probably through the quartic formula), we find the desired solution for b∈(0,1/2)b\in(0,1/2) is

This function b(x0)b(x_{0}) is a strictly decreasing function on $,anditsatisfies, and it satisfiesb(0)=\frac{1}{2},b(1)=\sqrt{2}-1.Thisimpliesthatwhenever. This implies that wheneverx_{0}issmall,wecanchooseis small, we can choosebclosertocloser to\frac{1}{2}inthein then=1$ step of the induction.

Similarly, for the induction step, it’s sufficient to show

Again, since we are always choosing b≤1/2b\leq 1/2, therefore we have 1≥2b1\geq 2b, and we will only need to focus on the first coefficient. To this end we rewrite the first term as

This implies we require n≥b1−2bn\geq\frac{b}{1-2b}, which increases as we choose bb closer to 1/21/2. However, if the induction starts the step ⌈b1−2b⌉\lceil\frac{b}{1-2b}\rceil, then this is not a problem, which leads to our conjecture.

This allows us to consider the upper bound

Similarly, the conjecture leads to the following approximation

Appendix F Additional Simulations and Discussions

In this section, we have additional simulations plotting the densities of ρdαβ\rho^{\alpha\beta}_{d} and VdαβV^{\alpha\beta}_{d} for shaped ReLU-like, sigmoid, and softplus networks. In particular, the density of VdαβV^{\alpha\beta}_{d} for ReLU-like networks can be found in Figure 6, the densities for sigmoid in Figure 7, and the densities for softplus in Figure 8.

From Figure 9, we can show that our results (3.3) converges at a rate of n−1/2n^{-1/2} in terms of the KS-distance.

F.2 Tuning Shape and Depth-to-Width Ratio

Since the existing shaping methods estimates the output correlation based on the infinite-width limit, we can easily improve the shape tuning based on the covariance SDEs. In particular, we consider the example of ReLU-like activations with correlation described by the SDE Equation 3.4. By simulating both the SDE and the infinite-width limit ODE, we arrive at the results in Figure 10.

We observe that simply by increasing c−c_{-} towards zero does not automatically reduce effects on the correlation when time tt (the depth-to-width ratio) is large. In other words, even a linear network will observe an increase in correlation when depth is large enough. Therefore shaping the activation alone is insufficient, but we also need to account for the depth-to-width ratio.

We also remark that Figure 10 only plotted the median for simplicity, but if we recall the density plots from Figure 1, correlation is heavily skewed and concentrated near 11. More precisely, while the median correlation is approximately 0.550.55, roughly 20%20\% of the samples are larger than 0.90.9. In other words, one in five random initializations will lead to a correlation worse than 0.90.9! As a consequence, practitioners implementing the shaping methods of should consider simulating the correlation SDE to account for the heavy skew.