Gradient flow dynamics of shallow ReLU networks for square loss and orthogonal inputs

Etienne Boursier, Loucas Pillaud-Vivien, Nicolas Flammarion

Introduction

Artificial neural networks are nowadays trained successfully to solve a large variety of learning tasks. However, a large number of fundamental questions surround their impressive success. Among them, the convergence to global minima of their non-convex training dynamics and their ability to generalise well despite fitting perfectly the dataset have challenged traditional machine learning belief. While a complete theory is still lacking, the machine learning community has recently come up with key steps that allow to tame the complexity of the problem: proving the convergence of gradient flow to zero loss (Mei et al., 2018; Chizat and Bach, 2018; Sirignano and Spiliopoulos, 2020; Rotskoff and Vanden-Eijnden, 2022), investigating the algorithmic selection of a specific global minimum, often referred as the implicit bias of an algorithm (Neyshabur et al., 2014; Zhang et al., 2021); while paying attention to the importance of the initialisation (Woodworth et al., 2020; Chizat et al., 2019). The aim of this article is to analyse precisely these three points for regression problems. This is done in a specific setting: for orthogonal inputs, we provide a complete characterisation of the gradient flows dynamics of training one-hidden layer ReLU neural networks with the square loss at small initialisation. We show that this non-convex optimisation dynamics captures most of the complexity mentioned above and thus could be a first step towards analysing more general setups.

Showing convergence of the gradient flow to a global minimum is an open and important question. Beyond the lazy regime (see next paragraph), only a few results were proven in the regression setting. The most promising route might be the link with Wassertein gradient flows for infinite neural networks. In that case, global convergence happens under mild conditions (Chizat and Bach, 2018; Wojtowytsch, 2020). Other works focus on local convergence (Zhou et al., 2021; Safran et al., 2021), or general criteria that eventually fail to encompass practical setups (Chatterjee, 2022; Chen et al., 2022). These latter works rest on Polyak-Łojasiewicz inequalities that in fact cannot be satisfied through the whole process if the dynamics travels near saddle points (Liu et al., 2022), as empirically observed (Dauphin et al., 2014). On the contrary, the present paper proves global convergence without resorting to large overparameterisation, dealing carefully with saddles.

The scale of initialisation plays an essential role in the behavior of the training dynamics. Indeed, an important example is that, at large initialisation, known as the lazy regime (Chizat et al., 2019), the neurons move relatively slightly implying that the dynamics is nearly convex and described by an effective kernel method with respect to the Neural Tangent Kernel (Jacot et al., 2018; Allen-Zhu et al., 2019; Arora et al., 2019). Instead, we are interested in another regime where the initialisation scale is small. This regime is known to be richer as it performs feature learning (Yang and Hu, 2021) but is also more challenging to analyse as it follows a truly non-convex dynamics (see details in Section 4).

In the regression case, the starting point governs where the flow converges. This observation suggests that a complete analysis of the trajectory may be required when one wants to understand the implicit bias in this case. Such descriptions have been undertaken by Maennel et al. (2018), who describe the initial alignment phase at small initialisation, and Li et al. (2020); Jacot et al. (2021) who conjecture that the dynamics travels from saddle to saddle. These papers provide intuitive content that we prove rigorously in the orthogonal setup. Finally, closest to our work are the following results on the classification of orthogonally separable data (Phuong and Lampert, 2020; Wang and Pilanci, 2021) and linearly separable, symmetric data (Lyu et al., 2021). The classification setup provides easier tools to analyse the problem: indeed, after the initial alignment phase, the network has already perfectly classified the data points in these settings. From there, it is known that the training loss converges to zero and that the parameters direction is biased towards KKT points of the max-margin problem (Lyu and Li, 2019; Ji and Telgarsky, 2020). Such tools cannot be applied after the alignment phase for regression, and we resort to a refined analysis of the trajectory to show both global convergence and implicit bias. On the other hand, Lyu et al. (2021) require a precise description of the dynamics to ensure convergence towards specific KKT points of the max-margin problem. Yet, the analysis of the dynamics is simplified by their symmetry assumption: the trajectory does not go through intermediate saddles and all the labels are simultaneously fitted. On the contrary, the dynamics we describe travels near an intermediate saddle point which separates two distinct fitting phases. This behaviour largely complicates the analysis, besides being more representative of the saddle to saddle dynamics observed in general settings.

1 Main contributions

We prove the convergence of the gradient flow towards a global minimum of the non-convex training loss for small enough initialisation and finite width.

As important as the convergence result, the dynamics is portrayed in Section 4: we quantitatively detail its different phases (alignment and fitting) and show it follows a saddle to saddle dynamics.

2 Notations

Setup and preliminaries

We introduce here the main assumptions on the data inputs.

The input points form an orthonormal family, i.e. ∀k,k′∈⟦n⟧\forall k,k^{\prime}\in\llbracket n\rrbracket, ⟨xk,xk′⟩=\mathds1k=k′\langle x_{k},x_{k^{\prime}}\rangle=\mathds{1}_{k=k^{\prime}}.

The data are assumed to be normalized only for convenience—the real limitation being that they are pairwise orthogonal. This assumption is exhaustively discussed in Section 3.2.

For all k∈⟦n⟧k\in\llbracket n\rrbracket, yk≠0y_{k}\neq 0 and ∑k∣yk>0yk2≠∑k∣yk<0yk2\sum_{k\mid y_{k}>0}y_{k}^{2}\neq\sum_{k\mid y_{k}<0}y_{k}^{2}.

This assumption on the data output is mild, e.g. has zero Lebesgue measure, and only permits to exclude degenerate situations.

As the limiting dynamics of the (stochastic) gradient descent with infinitesimal step-sizes (Li et al., 2019), we study the following gradient flow

initialised at θ0:=(a0,W0)\theta^{0}:=(a^{0},W^{0}). Since the ReLU is not differentiable at , the dynamics should be defined as a subgradient inclusion flow (Bolte et al., 2010). However, we show in Appendix D that the only ReLU subgradient that guarantees the existence of a global solution is σ′(x)=\mathds1x>0\sigma^{\prime}(x)=\mathds{1}_{x>0}. Hence, we stick with this choice throughout the paper. Another important difficulty of this non-differentiability is that Cauchy-Lipschitz theorem does not apply and uniqueness is not ensured. There have been attempts to define the solution of this Ordinary Diffential Equation (ODE) unequivocally (Eberle et al., 2021) as well as ways to circumvent this difficulty by resorting to smooth activations or additional data assumptions (Wojtowytsch, 2020; Chizat and Bach, 2020). Yet, we do not follow this line and demonstrate our results for all the gradient flows satisfying Equation 2.

2 Preliminary properties and initialisation

Let us derive here some preliminary properties of the gradient flows. If we rewrite explicitly the dynamics of Equation 2 on each layer separately, we have straightforwardly that for all j∈⟦m⟧j\in\llbracket m\rrbracket,

Importantly, by Lemma 1, the study of Equation 3 reduces to the hidden layer WW solely. We consider the following balanced initialisation:

As already stated, we are interested in the regime where the initialisation scale λ>0\lambda>0 is small. We also introduce the following sets of neurons that are crucial in the fitting process

The sets S+,1S_{+,1} and S−,1S_{-,1} are both non-empty.

Assumption 3 states that there are some neurons in two given cones at initialisation. It holds with probability 11 when the support of initialisation covers all directions and the width mm of the network goes to infinity. This is thus a weaker condition than the omni-directionality of neurons at initialisation (Wojtowytsch, 2020), which is instrumental to show convergence in the mean field regime (Chizat and Bach, 2018). On the other hand, it is stronger than the alignment condition of Abbe et al. (2022), which is known to be necessary for weak learning but might not lead to the implicit bias described in the next section.

Convergence and implicit bias characterisation

1 below states our main result on the convergence and implicit bias of one-hidden layer ReLU networks for regression tasks with orthogonal data.

Under Assumptions 1, 2 and 3, there exists λ∗>0\lambda^{*}>0 depending only on the data and the width such that, if λ≤λ∗\lambda\leq\lambda^{*}, the gradient flow initialised according to Equation 4 converges almost surely to some θλ∞\theta_{\lambda}^{\infty} of zero training loss, i.e. L(θλ∞)=0L(\theta_{\lambda}^{\infty})=0. Furthermore, there exists θ∗\displaystyle\theta^{*} such that

The significance of this result is thoroughly discussed in Section 3.2. Note that a quantitative and non-asymptotic version of 1, both in time and λ\lambda, is stated in Lemma 12 (Appendix B). Roughly, it states that the dynamics has already nearly converged after a time of order −ln⁡(λ)-\ln(\lambda) and then that the convergence happens at exponential speed. Note also that the neural network need not be overparametrised for the result to hold: the only sufficient and necessary requirement on the width mm stems from Assumption 3.

The proof of 1 rests on a precise description of the training dynamics, which is divided into four different phases. We here only sketch it at a very high level and a more thorough description, with quantitative intermediate lemmas, is given in Section 4.

During the first phase, hidden neurons align to a few representative directions, while remaining close to in norm. In particular, all hidden neurons in S+,1S_{+,1} (resp. S−,1S_{-,1}) align with some key vector D+D_{+} (resp. −D−-D_{-}) defined in Section 4.1. During the second phase, the neurons aligned with D+D_{+} grow in norm, while staying aligned with D+D_{+}, until fitting all the positive labels of the dataset (up to some error scaling with λ\lambda). Meanwhile, all the other neurons stay idle. Then similarly, the neurons aligned with −D−-D_{-} grow in norm during the third phase, until nearly fitting all the negative labels. Meanwhile, these neurons remain aligned with −D−-D_{-} and all other neurons remain idle. The precise description of these three phases is obtained by analysing the solutions of the limit ODEs when λ=0\lambda=0. The approximation errors that occur from dealing with non-zero λ\lambda are then carefully handled via Grönwall comparison arguments. Due to the large time scales (of order −ln⁡(λ)-\ln(\lambda)), the error can propagate on such large time spans. Handling these error terms is the main challenge of our proof and remains intricate despite the orthogonality assumption.

2 Discussion

Even if the orthogonal setting we consider is quite restrictive, it carries several characteristics that may be generic, either because they have been observed empirically, shown in related contexts or simply conjectured. We discuss these important points below.

1 states that the gradient flow converges to zero loss. Such a result is simple to show when the loss satisfies a Polyak-Łojasiewicz (PL) inequality (Bolte et al., 2007): ∥∇L∥2≥cL\|\nabla L\|^{2}\geq cL for c>0c>0. However, here, as the dynamics travels near saddles, this inequality is not verified through all the process. Circumventing this global argument, it is yet possible to formulate a refined analysis and show convergence if the dynamics arrives in a region where a local PL stands with a large enough constant. This refined analysis, inspired by the recent work of Chatterjee (2022)Note that the argument is certainly not new, but the cited article has the benefit of clearly presenting it., allows to characterise properly the last phase of the dynamics. We believe that this approach may help in showing convergence in other non-convex gradient flow/descent.

This regularisation is implicit, meaning that this effect does not result from any explicit regularisation (e.g. weight decay) performed during training (Shevchenko et al., 2021; Parhi and Nowak, 2022). This is only a consequence of the inner structure of the gradient flow and the scale of initialisation.

Following Equation 8, note the striking parallel between the inductive bias of infinitesimally small initialisation for regression and that of the classification problem with the logistic loss as a max-margin problem with respect to the F1\mathcal{F}_{1}-norm (Chizat and Bach, 2020). As already observed in the linear case (Woodworth et al., 2020), in contrast with classification, infinitesimally small initialisation is instrumental in regression to be biased towards small F1\mathcal{F}_{1}-norm functions. The role of initialisation is illustrated empirically in Appendix A.

Finally, let us stress that we did not address the question of what functions solve Equation 8, nor the question of the generalisation implied by such a bias. Related works on the first point come from a functional description of norms related to F1\mathcal{F}_{1} (Savarese et al., 2019; Ongie et al., 2019; Debarre et al., 2022). For the generalisation properties of small F1\mathcal{F}_{1} norm functions, we refer to Kurková and Sanguineti (2001); Bach (2017). Importantly, we recall that the question of how well low F1\mathcal{F}_{1}-norm functions generalise depends heavily on the a priori we have on the ground-truth (Petrini et al., 2022).

An important characteristics of the loss landscape is that the origin is a saddle point. Hence, as the dynamics is initialised at small scale λ\lambda, the radial movement is slow and neurons move out of the saddle after time scale −ln⁡λ-\ln\lambda. Meanwhile, the tangential movement of the neurons rules the dynamics and aligns their directions towards specific vectors. This has been first explained by Maennel et al. (2018) and referred as the quantisation phenomenon, because neural networks weights collapse to a small finite number of directions. We emphasise that this phase happens generically when initialisation is near the origin and that this part of our analysis can be directly extended to the general (i.e. non-orthognal) case. Phuong and Lampert (2020); Lyu et al. (2021) analysed a similar early alignment for classification with specific data structures.

When initialising the dynamics of a gradient flow near a saddle point of the loss, it is expected (but hard to prove generically) that the dynamics will alternate slow movements near saddles and rapid junctions between them. Such a behavior has been conjectured for linear neural networks (Li et al., 2020; Jacot et al., 2021) initialised near the origin. We precisely prove that such a phenomenon occurs: after initialisation, the dynamics visits one strict saddle. See Section 4, 1 for more details.

As its main limitation, 1 assumes orthogonal data points xkx_{k}. The orthogonality assumption disentangles the analysis as the different phases, where either the neurons align towards some direction or grow in norm, are well separated in that case. More precisely, the neurons do not change in direction once they have a non-zero norm in the case of orthogonal data. This separation between alignment and norm growth does not hold in the general case, as observed empirically in Appendix A. Extending our result to more general data thus remains a major challenge and requires additional theoretical tools. Nonetheless, as it can be the case in high dimension, our analysis can easily be extended to nearly orthogonal data where ∣⟨xk,xk′⟩∣≤δ|\langle x_{k},x_{k^{\prime}}\rangle|\leq\delta, with δ\delta of order λ\lambda. If however δ\delta is much larger than the initialisation scale, the dynamics is drastically different and becomes as hard as the general case to analyse. In Appendix A, we observe similar dynamics for high dimensional data, where the loss converges towards , goes through an intermediate saddle point and the final solution is close to a 22 neurons network.

A minor assumption is the balanced initialisation, i.e. ∥wj0∥=∣aj0∣\|w_{j}^{0}\|=|a_{j}^{0}|. If instead we initialise a0a^{0} as a Gaussian scaling with λ\lambda, the initialisation would be nearly balanced for small λ\lambda. This assumption is thus mostly used for simplicity and our analysis can be extended to unbalanced initialisations.

The exact value of λ∗\lambda^{*} is omitted for exposition’s clarity. Roughly, it can be inferred from the analysis that λ∗\lambda^{*} scales as Θ(1)me−Θ(n)\frac{\Theta(1)}{\sqrt{m}}e^{-\Theta(n)}. Interestingly, the 1m\frac{1}{\sqrt{m}} term is reminiscent of the mean field regime, which is known to induce implicit bias (Chizat and Bach, 2020; Lyu et al., 2021). On the other hand, the exponential dependency in nn is common in the implicit bias literature (Woodworth et al., 2020). For larger values of λ\lambda (but still in the mean field regime), the parameters empirically seem to also converge towards a minimal norm interpolator. The analysis yet becomes more intricate and we do not observe any separation between Phase 2 and Phase 3, i.e. there is no intermediate saddle in the trajectory.

Fine dynamics description: alignments and saddles

This section describes thoroughly the training dynamics of the gradient flow. It presents and discusses quantitative lemmas on the state of the neural network at the end of each different phase. In particular, mathematical formulations of the early alignment and saddle to saddle phenomena are provided.

First, we need to introduce additional notations for this section. We define vectors D+D_{+} and D−D_{-} that are the two directions towards which the neurons align

We also need to define c≔max⁡j∈⟦m⟧∥wj0∥/λc\coloneqq\max_{j\in\llbracket m\rrbracket}\|w_{j}^{0}\|/\lambda and r≔∥D+∥/∥D−∥r\coloneqq\|D_{+}\|/\|D_{-}\|. Assumption 2 implies that r≠1r\neq 1, and by symmetry we can assume r>1r>1 without any loss of generality. We additionally fix constants λ∗,ε>0\lambda_{*},\varepsilon>0, small enough and depending only on the dataset and the width mm.

2 Training dynamics

This section precisely describes the phases of the dynamics, summarised in Figure 1.

During the first phase, all the neurons remain small in norm, while moving tangentially (i.e. in directions). The neurons align according to several key directions: an initial clustering of neurons’ directions happens in this early phase, as observed by Maennel et al. (2018). As the neurons have small norm, hθt≈0h_{\theta^{t}}\approx 0 for this phase and Equation 9 approximates

This ODE corresponds to the descent/ascent gradient flow (depending on the sign of sj\mathsf{s}_{j}) on the sphere with objective ⟨Dj0,wj⟩\langle D^{0}_{j},\mathsf{w}_{j}\rangle. All neurons end up minimizing or maximizing their scalar product with Dj0D^{0}_{j}, which only depends on the activation of wj\mathsf{w}_{j}. As a consequence, neurons with similar activations align towards the same vector, leading to some quantisation of the neurons’ directions. This alignment happens in a relatively short time, so that the neurons cannot largely grow in norm. Lemma 2 below quantifies this effect for neurons in S+1,S_{+1,} and S−,1S_{-,1}, which are crucial to the training dynamics. Since all other neurons remain small in norm during the whole process, we do not focus on their direction.

For λ≤λ∗\lambda\leq\lambda^{*}, we have the following inequalities for t1=−εln⁡(λ)∥D−∥t_{1}=\frac{-\varepsilon\ln(\lambda)}{\|D_{-}\|}:

neurons in S+,1S_{+,1} are aligned with D+D_{+}: +∀j∈S+,1,⟨wjt1,D+⟩≥(1−2λε)∥D+∥\forall j\in S_{+,1},\langle\mathsf{w}_{j}^{t_{1}},D_{+}\rangle\geq(1-2\lambda^{\varepsilon})\|D_{+}\|,

neurons in S−,1S_{-,1} are aligned with −D−-D_{-}: ∀j∈S−,1,⟨wjt1,−D−⟩≥(1−2λε)∥D−∥\forall j\in S_{-,1},\langle\mathsf{w}_{j}^{t_{1}},-D_{-}\rangle\geq(1-2\lambda^{\varepsilon})\|D_{-}\|,

all neurons have small norm: ∀j∈⟦m⟧,∥wjt1∥≤2cλ1−rε\forall j\in\llbracket m\rrbracket,\|w_{j}^{t_{1}}\|\leq 2c\lambda^{1-r\varepsilon}.

During the second phase, the norm of the neurons in S+,1S_{+,1} (which are aligned with D+D_{+}) grows until fitting all positive labels. Meanwhile, all the other neurons do not move significantly. The key approximate ODE of this phase is given for u+(t):=∑j∈S+,1∥wjt∥2u_{+}(t):=\sum_{j\in S_{+,1}}\|w_{j}^{t}\|^{2} by

This equation implies that u+(t)u_{+}(t), the sum of the squared norms of neurons in S+,1S_{+,1}, eventually converges to n∥D+∥n\|D_{+}\| within a time −ln⁡(λ)/∥D+∥{-\ln(\lambda)}/{\|D_{+}\|} . Meanwhile, it needs to be shown that these neurons remain aligned with D+D_{+} and that the other neurons remain small in norm. This fine control is technical and relies on the orthogonality assumption. If data were not orthogonal, neurons could indeed realign while growing in norm as illustrated by Figure 4(d) in Appendix A. Lemma 3 below describes the state of the network at the end of the second phase.

If λ≤λ∗\lambda\leq\lambda^{*}, then for some time t2≤−1+3ε∥D+∥ln⁡(λ)t_{2}\leq-\frac{1+3\varepsilon}{\|D_{+}\|}\ln(\lambda):

neurons in S+,1S_{+,1} are aligned with D+D_{+}: ∀j∈S+,1,⟨wjt2,D+⟩≥∥D+∥−λε2\forall j\in S_{+,1},\langle\mathsf{w}_{j}^{t_{2}},D_{+}\rangle\geq\|D_{+}\|-\lambda^{\frac{\varepsilon}{2}},

neurons in S+,1S_{+,1} have a large norm: ∑j∈S+,1∥wjt2∥2=n∥D+∥−λε5\sum_{j\in S_{+,1}}\|w_{j}^{t_{2}}\|^{2}=n\|D_{+}\|-\lambda^{\frac{\varepsilon}{5}},

other neurons have small norm: ∀j∈⟦m⟧∖S+,1,∥wjt2∥≤2cλε\forall j\in\llbracket m\rrbracket\setminus S_{+,1},\|w_{j}^{t_{2}}\|\leq 2c\lambda^{\varepsilon}.

These three points directly imply that the loss is of order λε5\lambda^{\frac{\varepsilon}{5}} on the positive labels at time t2t_{2}.

As explained above, the positive labels are almost fitted by the action of the neurons belonging to S+,1S_{+,1} at the end of the second phase, whereas the other neurons still have infinitesimally small norm. At this point, the dynamics has reached the vicinity of a strict saddle point and requires a long time to escape it. The analysis actually leads to the following fact:

There exists a (strict) saddle point θS≠0\theta_{S}\neq 0 of LL such that if λ≤λ∗\lambda\leq\lambda^{*}:

The training trajectory thus starts at the saddle point and passes through a second non-trivial saddle point at the end of the second phase. This lemma illustrates the phenomenon of saddle to saddle dynamics discussed in Section 3.2 and conjectured for linear models by Li et al. (2020); Jacot et al. (2021). This intermediate saddle point is escaped when the norms of the neurons in S−,1S_{-,1} have significantly grown (i.e. become non-zero), which happens during a third phase described below.

The norm of the neurons in S−,1S_{-,1} (which are aligned with −D−-D_{-}) grows until fitting all negative labels during the third phase. Meanwhile, all other neurons do not move significantly. The additional difficulty in the analysis of this phase compared to the second one is that of controlling the possible movements of neurons in S+,1S_{+,1}. Their norm is indeed large during the whole phase, but they do not change consequently, because the positive labels are nearly perfectly fitted.

If λ≤λ∗\lambda\leq\lambda^{*}, then for some time t3≤−1+3rε∥D−∥ln⁡(λ)t_{3}\leq-\frac{1+3r\varepsilon}{\|D_{-}\|}\ln(\lambda):

neurons in S−,1S_{-,1} are aligned with −D−-D_{-}: ∀j∈S−,1,⟨wjt3,−D−⟩≥∥D−∥−λε14\forall j\in S_{-,1},\langle\mathsf{w}_{j}^{t_{3}},-D_{-}\rangle\geq\|D_{-}\|-\lambda^{\frac{\varepsilon}{14}},

neurons in S−,1S_{-,1} have a large norm: ∑j∈S−,1∥wjt3∥2=n∥D−∥−λε29\sum_{j\in S_{-,1}}\|w_{j}^{t_{3}}\|^{2}=n\|D_{-}\|-\lambda^{\frac{\varepsilon}{29}},

neurons in S+,1S_{+,1} did not move since phase 22: ∀j∈S+,1,∥wjt2−wjt3∥≤λε15\forall j\in S_{+,1},\|w_{j}^{t_{2}}-w_{j}^{t_{3}}\|\leq\lambda^{\frac{\varepsilon}{15}},

other neurons have small norm: ∀j∈⟦m⟧∖(S+,1∪S−,1),∥wjt3∥≤3cλε\forall j\in\llbracket m\rrbracket\setminus\left(S_{+,1}\cup S_{-,1}\right),\|w_{j}^{t_{3}}\|\leq 3c\lambda^{\varepsilon}.

To show this final convergence, we use a local PL condition given by Lemma 5.

For λ≤λ∗\lambda\leq\lambda^{*}, we have the following lower bound on the PL constant

where Θ\Theta is the set of parameters verifying the balancedness property.

Adapting arguments from the recent work by Chatterjee (2022), this implies that the training trajectory converges to an interpolator and stays in the aforementioned ball. It thus converges exponentially at a rate ∥D−∥\|D_{-}\| to a point close to a minimal norm interpolator, and the distance to this point goes to when λ\lambda goes to , hence implying 1. This exponential rate is only asymptotical: the dynamics still require a large time −ln⁡(λ)/∥D−∥-{\ln(\lambda)}/{\|D_{-}\|} to escape the two first saddles.

Experiments

This section confirms empirically the dynamics described in Section 4 on an orthogonal toy example. The code and animated versions of the figures are available in github.com/eboursier/GFdynamics. Additional experiments can be found in Appendix A; they illustrate the necessity of small initialisation for implicit bias and present similar experiments on non-orthogonal toy data. For the latter, we observe some similar training phenomena, but major differences appearing in the dynamics highlight the difficulty of dealing with non-orthogonality.

We consider the following two-point dataset: x1,y1=(−0.5,1),−1x_{1},y_{1}=(-0.5,1),-1 and x2,y2=(2,1),1x_{2},y_{2}=(2,1),1. It corresponds to unidimensional data with a second 11 coordinate for the bias term. We choose unidimensional data for a simpler visualisation. However, it restricts the number of observations to n=2n=2 to maintain orthogonality. Also, the inputs’ norms are not 11 here, but we recall that our analysis is not specific to this case. The width of the neural network is m=60m=60. We choose a balanced initialisation at scale λ=10−6/m\lambda=10^{-6}/\sqrt{m}. We then run gradient descent with a step size 10−310^{-3} to approximate the gradient flow trajectory.

Figure 2 shows the training dynamics on this example. In particular, the state of the network is shown at different steps. In Figure 2(a), all the neurons are close to at initialisation. Figure 2(b) shows the end of the first phase, where the neurons are aligned towards two key directions. After the second phase, shown in Figure 2(c), all the neurons aligned with D+D_{+} have grown in norm and the positive label is perfectly fitted. Similarly at the end of third phase in Figure 2(d), all neurons aligned with −D−-D_{-} have grown in norm and the negative label is fitted.

At the end of training, the loss is and the estimated function is simple. In particular, it only has two kinks, which illustrates the sparsity induced by the implicit bias. Also, the final estimated function might be counter-intuitive. Previous works on implicit bias indeed conjectured that the learned estimator is linear if the data can be linearly fitted (Kalimeris et al., 2019; Lyu et al., 2021). However, the learned function in Figure 2(d) has a smaller F1\mathcal{F}_{1}-norm than the linear interpolator.

Figure 3 shows the evolution of the loss during training. The saddle to saddle dynamics is well observed here: the parameters vector starts from the saddle point at initialisation and needs 50005000 iterations to leave this first saddle. A second saddle is then encountered at the end of the second phase and the trajectory only leaves this saddle around iteration 1100011000, once the norm of the neurons in S−,1S_{-,1} start being significant during the third phase. All these different experiments confirm 1 and the precise dynamics described in Section 4. Moreover, such training phenomena are not specific to the orthogonal data case, as observed in Appendix A.

Conclusion and perspectives

References

Appendix

This section presents additional experiments in the general setting where data are non-orthogonal. To be able to visualize the results, similarly to Section 5, we consider unidimensional data with a bias term, but with 55 data points. Here again, we consider a neural network of width m=60m=60 and run gradient descent with step size 10−310^{-3}.

Similarly to Figure 2, Figure 4 presents the training dynamics with a small initialisation (λ=10−4m\lambda=\frac{10^{-4}}{\sqrt{m}}). As in the orthogonal case, we first observe an early alignment phase in Figure 4(b). Afterwards, two groups of neurons (against a single group in the orthogonal case) grow in norms until reaching an intermediate saddle point in Figure 4(c). Note that during this norm-growth phase, the group of neurons does not remain aligned with a fixed direction, but changed in direction. A similar phenomenon happens when leaving the intermediate saddle point in Figure 4(d): the norm of these neurons still grow but they also change their direction. This behavior is what makes the general case fundamentally harder to analyse than the orthogonal one, where such a behavior does not happen. We believe that controlling this type of behavior is the key towards dealing with the general set up.

In Figure 4(e), we have a zero training loss and as in the orthogonal case, the estimated function is sparse in its number of kinks. Figure 4(f) shows that a saddle to saddle dynamics also happens in this general setup.

Figure 5 on the other hand studies the impact of the initialisation scale. In particular, it shows the training dynamics for a large scale of initialisation (λ=10m\lambda=\frac{10}{\sqrt{m}}). By comparing Figures 5(a) and 5(b), we indeed observe the lazy regime [Chizat et al., 2019]: the neuron weights do not significantly move between the initialisation and the end of training.

The final estimated function is not as simple as for small initialisation: it is approximately the interpolator coming from the Neural Tangent Kernel at initialisation, whose associated RKHS is described in Bietti and Mairal . The strong biased induced by initialisation and the poor generalising properties of this kernel interpolator illustrate the benefits of the rich regime obtained for small initialisation. However, this large initialisation has the advantage of converging towards loss much faster: the trajectory indeed does not go through saddle points that might significantly slow down the learning process. Similar results are observed for large initialisations when either considering unbalanced initialisation or orthogonal data.

A.2 High dimensional data

This section presents additional experiments for high dimensional, nearly orthogonal data. We generated n=75n=75 data points xix_{i} drawn independently at random according to a standard Gaussian distribution of dimension d=150d=150. It is then known that such points are almost orthogonal with large probability. The labels yiy_{i} are then given by a 66-neurons teacher network, whose weights were drawn at random following a Gaussian distribution.

From there, we trained a neural network with width m=200m=200 (without bias terms), an initialisation scale λ=10−20md\lambda=\frac{10^{-20}}{\sqrt{md}} and a gradient step size of 10−310^{-3}. Figure 6 illustrates the training dynamics of the parameters when projected onto the 22 dimensional space of the two principal components of the 200×150200\times 150 matrix associated to the hidden layer of the network.

The training loss profile is given by Figure 7(a). Figure 7(b) finally shows the explained variance ratio of the principal components of the PCA used in Figure 6.

Behaviors close to the orthogonal case can easily be observed here. First, we can see in Figure 7(a) that the training loss converges to and that the dynamics goes through an intermediate saddle point around iteration 7500075000. Also, thanks to Figure 6, we see the training dynamics (at least when projected onto the 22 dimensional space of the two principal components) follows similar phases. At first, we observe an early alignment phase where the neurons align towards two key directions. During a second phase, a first cluster of neurons grows in norm while keeping a fixed direction. During a third phase, the same happens for the other cluster of neurons. Only after these three phases, the neurons will slightly move from these two key directions and reach the final state of Figure 6(d), for which we see that the neurons are not exactly aligned with two directions.

This last state is what differs from the exactly orthogonal case, where the final solution consists of a 22 neurons network. This is obviously not the case here, as can be seen from Figure 6(d), but also from the explained variance ratios given in Figure 7(b). Indeed, we can there see that 80%80\% of the final state neurons are explained by the two principal components. This implies that the estimated interpolator is close to a 22 neurons network, but still far from being only represented by 22 direction (around 20%20\% of its variance is explained by the remaining directions).

Note that we here chose a very small initialisation scale (of order 10−2010^{-20}). As explained in Section 3.2, this confirms the exponential dependence of λ\lambda in the number of data points and is merely needed to observe a clear saddle point in the dynamics, but similar final states of training are observed for much larger values initialisation scales.

Appendix B Main proofs

In this section, we prove the main theorem. Sections B.1 and B.2 provide additional notations and recall the assumption on the initialisation. Then, each subsection corresponds to the study of the different phases of the dynamics: Sections B.3, B.4, B.5 and B.6 prove respectively the alignment phase, the positive, then negative label fitting phases and finally the convergence phase.

First, we introduce the following additional notations. We need to define the two dynamical vectors that encode respectively the fitting of positive and negative labels

Note that D+∈S+D_{+}\in S_{+} and −D−∈S−-D_{-}\in S_{-} and that neurons defined in Equations 5 and 6 correspond to

B.2 Initialisation and assumptions on the dataset

We recall here, for the sake of completeness, the setup at initialisation. We initialise the dynamics on a balanced fashion: aj0=sj∥wj0∥a_{j}^{0}=\mathsf{s}_{j}\|w_{j}^{0}\|. We take λ>0\lambda>0 and assume that wj0=λgjw_{j}^{0}=\lambda g_{j}, where each gj∼N(0,Id)g_{j}\sim\mathcal{N}(0,I_{d}) is an independent standard Gaussian. We assume that both S+,1S_{+,1} and S−,1S_{-,1} are non-empty and that for all kk, yk≠0y_{k}\neq 0. We also assume without loss of generality that ∥D+∥>∥D−∥\|D_{+}\|>\|D_{-}\|. This makes r=∥D+∥/∥D−∥>1r=\|D_{+}\|/\|D_{-}\|>1. We fix ε>0\varepsilon>0 small enough so that

We finally introduce an arbitrarily small λ∗>0\lambda_{*}>0 that only depends on the training set (xk,yk)k(x_{k},y_{k})_{k} and set c=max⁡j∈⟦m⟧∥wj0∥/λ=max⁡j∈⟦m⟧∥gj∥c=\max_{j\in\llbracket m\rrbracket}\|w_{j}^{0}\|/\lambda=\max_{j\in\llbracket m\rrbracket}\|g_{j}\|.

Because of the non-differentiability of the ReLU activation, the gradient flow is not uniquely defined. In the orthogonal case, we have in particular the following ODE

For any t′≥tt^{\prime}\geq t, j∈⟦m⟧j\in\llbracket m\rrbracket and k∈[n]k\in[n]:

This is a direct consequence of Equation 11. ∎

B.3 Phase 1: proof of Lemma 2

During the first phase, the neurons remain small in norm and they move tangentially to align with the vectors D+D_{+} and D−D_{-}. The duration of this movement is typically sub-logarithmic in λ\lambda, the initialisation scale. As ∥D+∥>∥D−∥\|D_{+}\|>\|D_{-}\|, neurons in S+,1S_{+,1} move slightly faster to align with D+D_{+} than the ones in S−,1S_{-,1} that align with −D−-D_{-}. We begin by describing the dynamics of neurons in S+,1S_{+,1} in Lemma 7 and derive similar results for the one of S−,1S_{-,1} in Lemma 8. These two Lemmas constitute together Lemma 2 of the main text.

First, we define the following ending time of the phase +1+1:

1+1). If λ≤λ∗\lambda\leq\lambda_{*}, then we have the following inequalities,

∀j∈⟦m⟧,∥wjt+1∥≤2cλ1−ε\forall j\in\llbracket m\rrbracket,\|w_{j}^{t_{+1}}\|\leq 2c\lambda^{1-\varepsilon}

∀j∈S+,1,⟨wjt+1,D+⟩≥(1−2λε)∥D+∥\forall j\in S_{+,1},\langle\mathsf{w}_{j}^{t_{+1}},D_{+}\rangle\geq(1-2\lambda^{\varepsilon})\|D_{+}\|.

For all j∈S+,1j\in S_{+,1}, let kk such that yk<0y_{k}<0, then ⟨wjt+1,xk⟩=0\langle w_{j}^{t_{+1}},x_{k}\rangle=0.

The condition (i) of the lemma means that the neurons do not grow so much in this first phase, whereas (ii) states that they have aligned with vectors D+D_{+}. Finally (iii) shows that neurons in S+,1S_{+,1} deactivate along negative labels during this first phase.

We divide the proof in three steps. First one is to control the growth of the neurons norm until t+1t_{+1}. The second step shows that the tangential movement is faster, while the third one shows that neurons in S+,1S_{+,1} “unalign” with the directions of negative labels.

First step: we show (i), i.e. that for t≤t+1t\leq t_{+1}, we have ∥wjt∥≤2cλ1−ε\|w_{j}^{t}\|\leq 2c\lambda^{1-\varepsilon}. Note first that, thanks to the balancedness, ajt=sj∥wjt∥a_{j}^{t}=\mathsf{s}_{j}\|w_{j}^{t}\|. Let τλ:=inf⁡{t≥0∣∃j∈⟦m⟧,∥wjt∥>2cλ1−ε}\tau_{\lambda}:=\inf\{t\geq 0\mid\exists j\in\llbracket m\rrbracket,\|w_{j}^{t}\|>2c\lambda^{1-\varepsilon}\}, then for all t≤τλt\leq\tau_{\lambda}, we have ∥wjt∥≤2cλ1−ε\|w_{j}^{t}\|\leq 2c\lambda^{1-\varepsilon} and hence,

Then, for λ∗\lambda_{*} such that λ∗2(1−ε)≤min⁡k∣yk∣4mc2\lambda_{*}^{2(1-\varepsilon)}\leq\frac{\min_{k}|y_{k}|}{4mc^{2}}, yk−hθt(xk)y_{k}-h_{\theta^{t}}(x_{k}) and yky_{k} have the same sign for any t≤τλt\leq\tau_{\lambda}. As a consequence, for jj such that sj=1\mathsf{s}_{j}=1, Equation 3 yields

Now, denote D+,jt:=1n∑k∣⟨xk,wjt⟩>0(yk−hθt(xk))xk\mathds1yk>0\displaystyle D^{t}_{+,j}:=\frac{1}{n}\sum_{\begin{subarray}{c}k\mid\langle x_{k},w_{j}^{t}\rangle>0\end{subarray}}\left(y_{k}-h_{\theta^{t}}(x_{k})\right)x_{k}\mathds{1}_{y_{k}>0}, we have

The same series of inequalities holds in the case sj=−1\mathsf{s}_{j}=-1 for ∥D−t∥\|D_{-}^{t}\|. Overall,

By Grönwall’s lemma, this gives for any t≤τλt\leq\tau_{\lambda}, ∥wjt∥≤∥wj0∥e(∥D+∥+4mc2λ2(1−ε))t\|w_{j}^{t}\|\leq\|w_{j}^{0}\|e^{\left(\|D_{+}\|+4mc^{2}\lambda^{2(1-\varepsilon)}\right)t}. This shows that for t≤min⁡(τλ,t+1)t\leq\min(\tau_{\lambda},t_{+1}),

where λ∗\lambda^{*} has been taken small enough. Hence τλ≥t+1\tau_{\lambda}\geq t_{+1}.

Second step: we show condition (ii). Indeed, let us now choose any j∈S+,1j\in S_{+,1}. We have the following decomposition:

so that as all vectors in the D−,jtD_{-,j}^{t} are different from the one of D+D_{+}, by orthogonality we have,

Now, let us analyse the tangential movement for these jj’s: for all t≤t+1t\leq t_{+1}, Equation 9 leads to the following growth comparison

and as we have, −⟨Dj,−t,wjt⟩⟨wjt,D+⟩≥0-\langle D^{t}_{j,-},\mathsf{w}_{j}^{t}\rangle\langle\mathsf{w}_{j}^{t},D_{+}\rangle\geq 0 and the third term lower bounded by −4mc2∥D+∥λ2(1−ε)-4mc^{2}\|D_{+}\|\lambda^{2(1-\varepsilon)}, it yields

Solutions of the ODE f′(t)=a2−f2(t)f^{\prime}(t)=a^{2}-f^{2}(t) with value in (−a,a)(-a,a) are of the form f(t)=atanh⁡(a(t+t0))f(t)=a\tanh(a(t+t_{0})) for some t0t_{0}. Note in the remaining of the proof a2=∥D+∥2−4mc2∥D+∥λ2(1−ε)a^{2}=\|D_{+}\|^{2}-4mc^{2}\|D_{+}\|\lambda^{2(1-\varepsilon)}. In our case, define t0t_{0} such that ⟨wj0,D+⟩=atanh⁡(at0)\langle\mathsf{w}_{j}^{0},D_{+}\rangle=a\tanh(at_{0}), then, by Grönwall comparison, it yields:

Note that ⟨wj0,D+⟩≥0\langle\mathsf{w}_{j}^{0},D_{+}\rangle\geq 0, so that t0≥0t_{0}\geq 0. As a consequence, we simply have

Now, using the inequality tanh⁡(x)≥1−2e−2x\tanh(x)\geq 1-2e^{-2x}, we have

We have the following inequalities on aa when choosing λ∗\lambda_{*} small enough:

This leads for any j∈S+,1j\in S_{+,1} and λ∗\lambda^{*} small enough to

Third step: we show (iii). Consider j∈S+,1j\in S_{+,1} and kk such that yk<0y_{k}<0. If ⟨wj0,xk⟩<0\langle w_{j}^{0},x_{k}\rangle<0, then by Lemma 6, there is nothing to prove. Otherwise, as for t≤t+1t\leq t_{+1} we have ∥wjt∥≤2cλ1−ε\|w_{j}^{t}\|\leq 2c\lambda^{1-\varepsilon}, it yields (as long as ⟨wjt,xk⟩>0\langle\mathsf{w}_{j}^{t},x_{k}\rangle>0)

Hence, if for some time τ≤t+1\tau\leq t_{+1}, we have ⟨wjτ,xk⟩=0\langle\mathsf{w}_{j}^{\tau},x_{k}\rangle=0, then for all times t∈[τ,t+1]t\in[\tau,t_{+1}], this quantity remains . We now continue to upperbound the derivative of ⟨wjτ,xk⟩\langle\mathsf{w}_{j}^{\tau},x_{k}\rangle:

From this, we sum all the yk<0y_{k}<0, and noting f(t):=−1n ⁣ ⁣∑k∣⟨wjt,xk⟩>0 ⁣ ⁣ ⁣ ⁣yk⟨wjt,xk⟩\mathds1yk<0\displaystyle f(t):=-\frac{1}{n}\!\!\sum_{k\mid\langle\mathsf{w}_{j}^{t},x_{k}\rangle>0}\!\!\!\!y_{k}\langle\mathsf{w}_{j}^{t},x_{k}\rangle\mathds{1}_{y_{k}<0}, we have

If we let a2:=1n2∑k∣yk<0yk2+5mc2λ2(1−ε)1n∑k∣yk<0yk>0a^{2}:=\frac{1}{n^{2}}\sum_{k\mid y_{k}<0}y^{2}_{k}+5mc^{2}\lambda^{2(1-\varepsilon)}\frac{1}{n}\sum_{k\mid y_{k}<0}y_{k}>0, we have f′(t)≤−a2+f(t)2f^{\prime}(t)\leq-a^{2}+f(t)^{2}, and as by Cauchy-Schwarz

where the two last inequalities are valid for λ\lambda small enough. Hence t↦f(t)t\mapsto f(t) is decreasing and f′(t)≤−a2+f(0)2f^{\prime}(t)\leq-a^{2}+f(0)^{2}, that is if we call b:=a2−f(0)2>0b:=a^{2}-f(0)^{2}>0, we have f(t)≤−bt+f(0)f(t)\leq-bt+f(0), and for τ=f(0)/b\tau=f(0)/b, we have f(τ)=0f(\tau)=0. And as τ<t+1\tau<t_{+1}, we have f(t+1)=0f(t_{+1})=0. This concludes the proof of condition (iii) and of the lemma. ∎

Now, we show that the exact same conclusion is also valid for the neurons in S−,1S_{-,1} i.e., after a sub-logarithmic time, they eventually align with −D−-D_{-} and deactivates with respect to positive outputs. We define here a ending time similar to (12):

We prove the following lemma analogous to Lemma 7.

If λ≤λ∗\lambda\leq\lambda_{*}, then we have the following inequalities,

∀j∈⟦m⟧,∥wjt−1∥≤2cλ1−rε\forall j\in\llbracket m\rrbracket,\|w_{j}^{t_{-1}}\|\leq 2c\lambda^{1-r\varepsilon}

∀j∈S−,1,⟨wjt−1,−D−⟩≥(1−2λε)∥D−∥\forall j\in S_{-,1},\langle\mathsf{w}_{j}^{t_{-1}},-D_{-}\rangle\geq(1-2\lambda^{\varepsilon})\|D_{-}\|.

For all j∈S−,1j\in S_{-,1}, let kk such that yk>0y_{k}>0, then ⟨wjt−1,xk⟩=0\langle w_{j}^{t_{-1}},x_{k}\rangle=0.

The proof is essentially the same than the proof of Lemma 7. We will be short and underline solely the main differences.

First step: we show condition (i), i.e. that for t≤t−1t\leq t_{-1}, we have ∥wjt∥≤2cλ1−rε\|w_{j}^{t}\|\leq 2c\lambda^{1-r\varepsilon}. Indeed, let τλr:=inf⁡{t≥0∣∃j∈⟦m⟧,∥wjt∥>2cλ1−rε}\tau^{r}_{\lambda}:=\inf\{t\geq 0\mid\exists j\in\llbracket m\rrbracket,\|w_{j}^{t}\|>2c\lambda^{1-r\varepsilon}\}. For all t≤τλrt\leq\tau^{r}_{\lambda}, we have ∣hθt(xk)∣≤4mc2λ2(1−rε)\left|h_{\theta^{t}}(x_{k})\right|\leq 4mc^{2}\lambda^{2(1-r\varepsilon)} and similarly to the proof of Lemma 7, we have that for t≤τλrt\leq\tau^{r}_{\lambda}, ∥wjt∥≤∥wj0∥e(∥D+∥+4mc2λ2(1−rε))t\|w_{j}^{t}\|\leq\|w_{j}^{0}\|e^{\left(\|D_{+}\|+4mc^{2}\lambda^{2(1-r\varepsilon)}\right)t}. This shows that for t≤min⁡(τλr, t−1)t\leq\min(\tau^{r}_{\lambda},\,t_{-1}),

where λ∗\lambda^{*} has been taken small enough. Hence τλr≥t−1\tau^{r}_{\lambda}\geq t_{-1}.

Second step: we show (ii), i.e. that the neurons almost align after time t−1t_{-1}. Indeed, for j∈S−,1j\in S_{-,1}, similarly to Lemma 7, we have that for all t≤t−1t\leq t_{-1}:

Denoting ar2=∥D−∥2−4mc2∥D−∥λ2(1−rε)a_{r}^{2}=\|D_{-}\|^{2}-4mc^{2}\|D_{-}\|\lambda^{2(1-r\varepsilon)}, we have by Grönwall comparison

Now, this gives ⟨wjt−1,−D−⟩>ar(1−2e2arrε∥D+∥ln⁡(λ))\langle\mathsf{w}_{j}^{t_{-1}},-D_{-}\rangle>a_{r}(1-2e^{2a_{r}\frac{r\varepsilon}{\|D_{+}\|}\ln(\lambda)}) and lower bounding ara_{r} as before, we have

Third step: we show (iii). Consider j∈S−,1j\in S_{-,1} and kk such that yk>0y_{k}>0. If ⟨wj0,xk⟩<0\langle w_{j}^{0},x_{k}\rangle<0, then by Lemma 6, there is nothing to prove. Otherwise, as for t≤t−1t\leq t_{-1} we have ∥wjt∥≤2cλ1−rε\|w_{j}^{t}\|\leq 2c\lambda^{1-r\varepsilon}, it yields (as long as ⟨wjt,xk⟩>0\langle\mathsf{w}_{j}^{t},x_{k}\rangle>0)

Hence, if for some time τ≤t−1\tau\leq t_{-1}, we have ⟨wjτ,xk⟩=0\langle\mathsf{w}_{j}^{\tau},x_{k}\rangle=0, then for all times t∈[τ,t−1]t\in[\tau,t_{-1}], this quantity will remain . We now we continue to upperbound the derivative of ⟨wjτ,xk⟩\langle\mathsf{w}_{j}^{\tau},x_{k}\rangle:

From this, we sum all the yk>0y_{k}>0, and noting f(t):=1n ⁣ ⁣∑k∣⟨wjt,xk⟩>0 ⁣ ⁣ ⁣ ⁣yk⟨wjt,xk⟩\mathds1yk>0\displaystyle f(t):=\frac{1}{n}\!\!\sum_{k\mid\langle\mathsf{w}_{j}^{t},x_{k}\rangle>0}\!\!\!\!y_{k}\langle\mathsf{w}_{j}^{t},x_{k}\rangle\mathds{1}_{y_{k}>0}, we have

If we let a2:=1n2∑k∣yk>0yk2−5mc2λ2(1−rε)1n∑k∣yk>0yk>0a^{2}:=\frac{1}{n^{2}}\sum_{k\mid y_{k}>0}y^{2}_{k}-5mc^{2}\lambda^{2(1-r\varepsilon)}\frac{1}{n}\sum_{k\mid y_{k}>0}y_{k}>0, we have f′(t)≤−a2+f(t)2f^{\prime}(t)\leq-a^{2}+f(t)^{2}, and as by Cauchy-Schwarz

where the two last inequalities are valid for λ\lambda small enough. Hence t↦f(t)t\mapsto f(t) is decreasing and f′(t)≤−a2+f(0)2f^{\prime}(t)\leq-a^{2}+f(0)^{2}, that is if we call b:=a2−f(0)2>0b:=a^{2}-f(0)^{2}>0, we have f(t)≤−bt+f(0)f(t)\leq-bt+f(0), and for τ=f(0)/b\tau=f(0)/b, we have f(τ)=0f(\tau)=0. And as τ<t−1\tau<t_{-1}, we have f(t−1)=0f(t_{-1})=0. This concludes the proof of condition (iii) and of the Lemma. ∎

We end this subsection by a lemma that shows that until t±1t_{\pm 1}, the neurons wjtw^{t}_{j} with j∈S±1j\in S_{\pm 1}, could not have collapsed to .

Define c‾=min⁡∥wj0∥/λ\underline{c}=\min\|w_{j}^{0}\|/\lambda, then, if λ≤λ∗\lambda\leq\lambda^{*},

for all t≤t+1t\leq t_{+1}, ∀j∈S+,1\forall j\in S_{+,1}, ∥wjt∥>c‾λ1+ε/2\displaystyle\|w_{j}^{t}\|>\underline{c}\lambda^{1+\varepsilon}/2,

for all t≤t−1t\leq t_{-1}, ∀j∈S−,1\forall j\in S_{-,1}, ∥wjt∥>c‾λ1+rε/2\displaystyle\|w_{j}^{t}\|>\underline{c}\lambda^{1+r\varepsilon}/2.

Let us begin with the first point. As stated in the proof of Lemma 7, yk−hθt(xk)y_{k}-h_{\theta^{t}}(x_{k}) and yky_{k} have the same sign for any t≤t+1t\leq t_{+1}. As a consequence, for j∈S+,1j\in S_{+,1}, it yields

Then, by Grönwall’s lemma, this gives for any t≤t+1t\leq t_{+1}, ∥wjt∥≥∥wj0∥e−(∥D+∥+4mc2λ2(1−ε))t\|w_{j}^{t}\|\geq\|w_{j}^{0}\|e^{-\left(\|D_{+}\|+4mc^{2}\lambda^{2(1-\varepsilon)}\right)t}. This shows that, for t≤t+1t\leq t_{+1},

This concludes the first point of the lemma. The second point is very similar to the first one. Indeed, in this case also, for t≤t−1t\leq t_{-1}, yky_{k} rules the sign of the residual so that for j∈S−,1j\in S_{-,1},

Then, by Grönwall’s lemma, this gives for any t≤t−1t\leq t_{-1}, ∥wjt∥≥∥wj0∥e−(∥D+∥+4mc2λ2(1−rε))t\|w_{j}^{t}\|\geq\|w_{j}^{0}\|e^{-\left(\|D_{+}\|+4mc^{2}\lambda^{2(1-r\varepsilon)}\right)t}. This shows that, as t≤t−1t\leq t_{-1},

B.4 Phase 2: proof of Lemma 3

During the second phase, the norm of the neurons in S+,1S_{+,1} (which are aligned with D+D_{+}) grows until perfectly fitting all the positive labels of the training points. Meanwhile, all the other neurons do not move significantly. We define the ending time of the second phase as

The following lemma corresponds to Lemma 3 of the main text.

If λ≤λ∗\lambda\leq\lambda_{*}, then we have the following inequalities

t2≤−1+3ε∥D+∥ln⁡(λ),t_{2}\leq-\frac{1+3\varepsilon}{\|D_{+}\|}\ln(\lambda),

∀j∈⟦m⟧∖S+,1,∥wjt2∥<2cλε,\forall j\in\llbracket m\rrbracket\setminus S_{+,1},\|w_{j}^{t_{2}}\|<2c\lambda^{\varepsilon},

∀j∈S+,1,⟨wjt2,D+⟩>∥D+∥−λε2.\forall j\in S_{+,1},\langle\mathsf{w}_{j}^{t_{2}},D_{+}\rangle>\|D_{+}\|-\lambda^{\frac{\varepsilon}{2}}.

The first point of the lemma states that the second phase lasts a time of order −ln⁡(λ)-\ln(\lambda). The other two points state that the first and third conditions in the definition of t2t_{2} do not hold at the end of the phase. Thus, the second condition in Equation 17 holds at t2t_{2}, meaning that the norm of the neurons in S+,1S_{+,1} have grown enough to fit the positive labels:

A direct consequence of thisThis comes from the decomposition given by Equation 19 in the proof. is that at the end of the second phase, the training loss on the positive labels is of order λ2ε5\lambda^{\frac{2\varepsilon}{5}} (at least).

Preliminaries: define for this proof h(t)=∥D+∥−min⁡j∈S+,1⟨wjt,D+⟩h(t)=\|D_{+}\|-\min_{j\in S_{+,1}}\langle\mathsf{w}_{j}^{t},D_{+}\rangle. By definition of the second phase, we have h(t)≤λε2h(t)\leq\lambda^{\frac{\varepsilon}{2}} for any t∈[t+1,t2]t\in[t_{+1},t_{2}].

We first show that during the whole second phase, D+tD_{+}^{t} is almost colinear with D+D_{+}. For any j∈S+,1j\in S_{+,1} and t∈[t+1,t2]t\in[t_{+1},t_{2}], define djt=wjt−D+∥D+∥d_{j}^{t}=\mathsf{w}_{j}^{t}-\frac{D_{+}}{\|D_{+}\|}. Since wjt\mathsf{w}_{j}^{t} and D+∥D+∥\frac{D_{+}}{\|D_{+}\|} are both of norm 11, we have ∥djt∥2=−2⟨djt,D+∥D+∥⟩\|d_{j}^{t}\|^{2}=-2\langle d_{j}^{t},\frac{D_{+}}{\|D_{+}\|}\rangle. By definition of h(t)h(t),

So finally, we have for any j∈S+,1j\in S_{+,1} and t∈[t+1,t2]t\in[t_{+1},t_{2}] the decomposition

Now let any kk such that yk>0y_{k}>0. Using Equation 18 and the fact that ∥wjt∥≤2cλε\|w_{j}^{t}\|\leq 2c\lambda^{\varepsilon} for any j∉S+,1j\not\in S_{+,1}, we have for any t∈[t+1,t2]t\in[t_{+1},t_{2}]:

where ∣hk(t)∣≤∑j∈S+,1∥wjt∥2∣⟨djt,xk⟩∣+4mc2λ2ε\displaystyle|h_{k}(t)|\leq\sum_{j\in S_{+,1}}\|w_{j}^{t}\|^{2}|\langle d_{j}^{t},x_{k}\rangle|+4mc^{2}\lambda^{2\varepsilon}. It follows that

Roughly, this means that as long as 1−∑j∈S+,1∥wjt∥2n∥D+∥1-\frac{\sum_{j\in S_{+,1}}\|w_{j}^{t}\|^{2}}{n\|D_{+}\|} is large enough, D+tD_{+}^{t} is almost colinear with D+D_{+}. Precisely, we have for any t∈[t+1,t2]t\in[t_{+1},t_{2}] the following decomposition

First point: denote u+(t)=∑j∈S+,1∥wjt∥2u_{+}(t)=\sum_{j\in S_{+,1}}\|w_{j}^{t}\|^{2}. We then have by balancedness and Equation 3

So we have the following growth comparison

By definition of the second phase, u+(t)≤n∥D+∥u_{+}(t)\leq n\|D_{+}\| for any t∈[t+1,t2]t\in[t_{+1},t_{2}]. We chose λ\lambda small enough, so that 2∥D+∥λε4≥4mc2λ2ε\sqrt{2\|D_{+}\|}\lambda^{\frac{\varepsilon}{4}}\geq 4mc^{2}\lambda^{2\varepsilon}. This implies that ∥h⊥(t)∥∨∣h+(t)∣∥D+∥≤22∥D+∥λε4\|h_{\bot}(t)\|\lor|h_{+}(t)|\|D_{+}\|\leq 2\sqrt{2\|D_{+}\|}\lambda^{\frac{\varepsilon}{4}}. So we finally have the following growth comparison during the second phase

Solution of the ODE f′(t)=af(t)−bf(t)2f^{\prime}(t)=af(t)-bf(t)^{2} with f(0)∈(0,ab)f(0)\in(0,\frac{a}{b}) are of the form f(t)=abea(t−τ)1+ea(t−τ)f(t)=\frac{a}{b}\frac{e^{a(t-\tau)}}{1+e^{a(t-\tau)}}. Note in the following

By Grönwall comparison, for any t∈[t+1,t2]t\in[t_{+1},t_{2}],

Thanks to Lemma 9, u+(t+1)≥c‾24λ2(1+ε)u_{+}(t_{+1})\geq\frac{\underline{c}^{2}}{4}\lambda^{2(1+\varepsilon)}. This implies that

In particular, if t2−t+1>−1+2ε∥D+∥ln⁡(λ)t_{2}-t_{+1}>-\frac{1+2\varepsilon}{\|D_{+}\|}\ln(\lambda), then

Note that a(λ)b(λ)=n∥D+∥−O(λε4)\frac{a(\lambda)}{b(\lambda)}=n\|D_{+}\|-\mathcal{O}\left(\lambda^{\frac{\varepsilon}{4}}\right) and lim⁡λ→0a(λ)∥D+∥(1+2ε)−2(1+ε)=2ε\lim_{\lambda\to 0}\frac{a(\lambda)}{\|D_{+}\|}(1+2\varepsilon)-2(1+\varepsilon)=2\varepsilon. As a consequence, for λ\lambda small enough, we have u+(t2)>n∥D+∥−λε5u_{+}(t_{2})>n\|D_{+}\|-\lambda^{\frac{\varepsilon}{5}} if t2−t+1>−1+2ε∥D+∥ln⁡(λ)t_{2}-t_{+1}>-\frac{1+2\varepsilon}{\|D_{+}\|}\ln(\lambda). This would break the second condition in the definition of the second phase Equation 17, so that t2−t+1≤−1+2ε∥D+∥ln⁡(λ)t_{2}-t_{+1}\leq-\frac{1+2\varepsilon}{\|D_{+}\|}\ln(\lambda) and thanks to Lemma 7, t2≤−1+3ε∥D+∥ln⁡(λ)t_{2}\leq-\frac{1+3\varepsilon}{\|D_{+}\|}\ln(\lambda).

Second point: let j∈⟦m⟧∖S+,1j\in\llbracket m\rrbracket\setminus S_{+,1}. Similarly to the proof of Lemma 8, we can show if sj=−1\mathsf{s}_{j}=-1 during the second phase that

If sj=1\mathsf{s}_{j}=1 instead, there is some kjk_{j} such that ykj>0y_{k_{j}}>0 and ⟨wjt,xkj⟩<0\langle w_{j}^{t},x_{k_{j}}\rangle<0 thanks to Lemma 6 and the continuous initialisation. In that case we have the following inequalities

where we recall α=min⁡k∣yk>0yk22∥D+∥2>0\alpha=\frac{\min_{k\mid y_{k}>0}y_{k}^{2}}{2\|D_{+}\|^{2}}>0. The previous inequalities are also valid during the first phase. In any case, for any j∉S+,1j\not\in S_{+,1} and t≤t2t\leq t_{2}:

Note that we chose ε\varepsilon small enough in Section B.1, so that (1+3ε)max⁡((1−α)∥D+∥,∥D−∥)∥D+∥≤1−ε(1+3\varepsilon)\frac{\max\left((1-\alpha)\|D_{+}\|,\|D_{-}\|\right)}{\|D_{+}\|}\leq 1-\varepsilon. Since t2≤−1+3ε∥D+∥ln⁡(λ)t_{2}\leq-\frac{1+3\varepsilon}{\|D_{+}\|}\ln(\lambda), Grönwall inequality yields for any j∉S+,1j\not\in S_{+,1} and λ\lambda small enough

Third point: let j∈S+,1j\in S_{+,1}. Recall that we have for any t∈[t+1,t2]t\in[t_{+1},t_{2}]

For any t∈[t+1,tu]t\in[t_{+1},t_{u}], we have for λ\lambda small enough

The positive solution of the ODE f′(t)=a2−b2f(t)2f^{\prime}(t)=a^{2}-b^{2}f(t)^{2} is either increasing if f(0)≤abf(0)\leq\frac{a}{b} or remains larger than ab\frac{a}{b}. The following inequality thus implies by Grönwall comparison for any t∈[t+1,tu]t\in[t_{+1},t_{u}] and λ\lambda small enough

We thus have h(tu)<5∥D+∥λεh(t_{u})<5\|D_{+}\|\lambda^{\varepsilon}. So u+(tu)=n∥D+∥7u_{+}(t_{u})=\frac{n\|D_{+}\|}{7}. Let us now bound t2−tut_{2}-t_{u}. Similarly to Equation 21, we actually have for any t∈[tu,t2]t\in[t_{u},t_{2}]

We showed u+(tu)=n∥D+∥7u_{+}(t_{u})=\frac{n\|D_{+}\|}{7} and so τu≤tu−1a(λ)ln⁡(b(λ)a(λ)n∥D+∥7)\tau_{u}\leq t_{u}-\frac{1}{a(\lambda)}\ln(\frac{b(\lambda)}{a(\lambda)}\frac{n\|D_{+}\|}{7}), i.e. for any t∈[tu,t2]t\in[t_{u},t_{2}]:

Similarly to the proof of the first point, we can then show that for λ\lambda small enough, t2−tu≤−ε5∥D+∥ln⁡(λ)t_{2}-t_{u}\leq-\frac{\varepsilon}{5\|D_{+}\|}\ln(\lambda). Now note that for any t∈[tu,t2]t\in[t_{u},t_{2}], Equation 22 yields:

By Grönwall’s comparison, we thus have for λ\lambda small enough

B.5 Phase 3: proof of Lemma 4

Similarly to the second phase with the positive labels, the third phase aims at fitting the negative labels. During this third phase, the norm of the neurons in S−,1S_{-,1} (which are aligned with −D−-D_{-}) grows until perfectly fitting all the negative labels of the training points; while all the other neurons do not change significantly. We define the ending time of the second phase as

The following lemma corresponds to Lemma 4 of the main text.

If λ≤λ∗\lambda\leq\lambda_{*},then the following inequalities hold

t3≤−1+3rε∥D−∥ln⁡(λ),t_{3}\leq-\frac{1+3r\varepsilon}{\|D_{-}\|}\ln(\lambda),

∀j∉S+,1∪S−,1,∥wjt3∥<3cλε,\forall j\not\in S_{+,1}\cup S_{-,1},\|w_{j}^{t_{3}}\|<3c\lambda^{\varepsilon},

∀j∈S−,1,⟨wjt3,−D−⟩>∥D−∥−λε14,\forall j\in S_{-,1},\langle\mathsf{w}_{j}^{t_{3}},-D_{-}\rangle>\|D_{-}\|-\lambda^{\frac{\varepsilon}{14}},

∥D+t3∥<λε14\|D_{+}^{t_{3}}\|<\lambda^{\frac{\varepsilon}{14}}.

The first point of the lemma states that the third phase also lasts a time of order −ln⁡(λ)-\ln(\lambda). It actually ends after the second one (t3>t2t_{3}>t_{2}), since the neurons in S−,1S_{-,1} do not grow in norm during the second one. The last point states that the positive labels remain fitted during the third phase. As a consequence, this also means that the neurons of S+,1S_{+,1} do not change a lot after t2t_{2}. The other points imply that the second condition in Equation 23 holds at t3t_{3}, meaning that the norm of the neurons in S−,1S_{-,1} have grown enough to fit the negative labels.

Preliminaries: the three first points of this proof share similarities with the proof of Lemma 10. For conciseness and clarity, parts of the proof are shortened, as they follow the same lines as the proof of Lemma 10. Similarly to the proof of the second phase, we define h(t)=∥D−∥−min⁡j∈S−,1⟨wjt,−D−⟩h(t)=\|D_{-}\|-\min_{j\in S_{-,1}}\langle w_{j}^{t},-D_{-}\rangle and we have for any j∈S−,1j\in S_{-,1} the decomposition

Since we chose λ\lambda small enough, this implies that hθt(xk)≥ykh_{\theta^{t}}(x_{k})\geq y_{k} for kk such that yk<0y_{k}<0 and t∈[t−1,t3]t\in[t_{-1},t_{3}]. Moreover, thanks to Lemma 7, ⟨wjt+1,xk⟩=0\langle w_{j}^{t_{+1}},x_{k}\rangle=0 for any such kk and j∈S+,1j\in S_{+,1}. Because of this, we have for any [t−1,t3][t_{-1},t_{3}]

Equation 25 is thus an equality on [t−1,t3][t_{-1},t_{3}]. Similarly to the second phase for D+tD_{+}^{t}, we can now decompose D−tD_{-}^{t} as

where ⟨h⊥(t),D−⟩=0\langle h_{\bot}(t),D_{-}\rangle=0 and ∣h−(t)∣∥D−∥∨∥h⊥(t)∥≤1n∑j∈S−,1∥wjt∥22h(t)∥D−∥+9mc2λ2ε|h_{-}(t)|\|D_{-}\|\lor\|h_{\bot}(t)\|\leq\frac{1}{n}\sum_{j\in S_{-,1}}\|w_{j}^{t}\|^{2}\sqrt{\frac{2h(t)}{\|D_{-}\|}}+9mc^{2}\lambda^{2\varepsilon}.

Using the previous decompositions, we also have the following inequalities for any j∈S−,1j\in S_{-,1}

where ∣gj(t)∣≤∥djt∥∥D+t∥|g_{j}(t)|\leq\|d_{j}^{t}\|\|D_{+}^{t}\|.

First point: with the previous decompositions, we can now prove the first point of Lemma 11. Define u−(t)=∑j∈S−,1∥wjt∥2u_{-}(t)=\sum_{j\in S_{-,1}}\|w_{j}^{t}\|^{2}. Based on the previous inequalities, we have for any t∈[t−1,t3]t\in[t_{-1},t_{3}]

From there, we can show similarly to the proof of the second phase that t3≤−1+3rε∥D−∥ln⁡(λ)t_{3}\leq-\frac{1+3r\varepsilon}{\|D_{-}\|}\ln(\lambda).

Second point: first consider j∉S+,1∪S−,1j\not\in S_{+,1}\cup S_{-,1} such that sj=1\mathsf{s}_{j}=1. Thanks to Lemma 10, we already have ∥wjt∥<2cλε\|w_{j}^{t}\|<2c\lambda^{\varepsilon} for any t≤t2t\leq t_{2}. For t∈[t2,t3]t\in[t_{2},t_{3}], we then have

Grönwall inequality then gives for λ\lambda small enough: ∥wjt3∥≤2cλεe−λε141+3rε∥D−∥ln⁡(λ)<3cλε\|w_{j}^{t_{3}}\|\leq 2c\lambda^{\varepsilon}e^{-\lambda^{\frac{\varepsilon}{14}}\frac{1+3r\varepsilon}{\|D_{-}\|}\ln(\lambda)}<3c\lambda^{\varepsilon}.

Let now j∉S+,1∪S−,1j\not\in S_{+,1}\cup S_{-,1} such that sj=−1\mathsf{s}_{j}=-1. By definition, there is some kjk_{j} such that ykj<0y_{k_{j}}<0 and ⟨wjt,xkj⟩<0\langle w_{j}^{t},x_{k_{j}}\rangle<0. Similarly to the proof of Lemma 10, we then have for any t∈[0,t3]t\in[0,t_{3}]:

where we recall β=min⁡k,yk<0yk22∥D−∥2>0\beta=\frac{\min_{k,y_{k}<0}y_{k}^{2}}{2\|D_{-}\|^{2}}>0. Note that we chose ε\varepsilon small enough, so that (1+3rε)(1−β)≤1−ε(1+3r\varepsilon)(1-\beta)\leq 1-\varepsilon. Grönwall inequality then yields for λ\lambda small enough ∥wjt3∥<3cλε\|w_{j}^{t_{3}}\|<3c\lambda^{\varepsilon}.

Third point: let j∈S−,1j\in S_{-,1}. Recall that for any t∈[t−1,t3]t\in[t_{-1},t_{3}], ⟨wjt,Djθt⟩=⟨wjt,D−t⟩+gj(t)\langle\mathsf{w}_{j}^{t},D_{j}^{\theta^{t}}\rangle=\langle\mathsf{w}_{j}^{t},D_{-}^{t}\rangle+g_{j}(t). This yields for any t∈[t−1,t3]t\in[t_{-1},t_{3}]

Similarly to the third phase, we can show for λ\lambda small enough the following sequence of properties

h(tu)<5∥D−∥λε8h(t_{u})<5\|D_{-}\|\lambda^{\frac{\varepsilon}{8}}

t3−tu≤−ε57∥D−∥ln⁡(λ)t_{3}-t_{u}\leq-\frac{\varepsilon}{57\|D_{-}\|}\ln(\lambda)

⟨wjt3,−D−⟩>∥D−∥−λε14\langle\mathsf{w}_{j}^{t_{3}},-D_{-}\rangle>\|D_{-}\|-\lambda^{\frac{\varepsilon}{14}}.

Thanks to Lemma 10, τ≥t2\tau\geq t_{2}. For any kk such that yk>0y_{k}>0, we have

Recall the for any j∈S+,1j\in S_{+,1}, ⟨wjt,xk⟩=0\langle w_{j}^{t},x_{k}\rangle=0 for any kk such that yk<0y_{k}<0. From there, we have ⟨wjt,Djθt⟩=⟨wjt,D+t⟩\langle w_{j}^{t},D_{j}^{\theta^{t}}\rangle=\langle w_{j}^{t},D_{+}^{t}\rangle. This leads to the following inequalities

Given the bounds on ∥D+t∥\|D_{+}^{t}\| and t3t_{3}, a simple Grönwall argument implies that ∑j∈S+,1∥wjt∥2\sum_{j\in S_{+,1}}\|w_{j}^{t}\|^{2} did not change significantly between t2t_{2} and t3t_{3}. In particular for t∈[t2,τ]t\in[t_{2},\tau] and λ\lambda small enough:

Since ∥D+t2∥≤2∥D+∥λε5\|D_{+}^{t_{2}}\|\leq 2\|D_{+}\|\lambda^{\frac{\varepsilon}{5}} thanks to Lemma 10, the above inequality implies by Grönwall comparison that ∥D+t∥≤(1+2∥D+∥)λε5\|D_{+}^{t}\|\leq\left(1+2\|D_{+}\|\right)\lambda^{\frac{\varepsilon}{5}} for any t∈[t2,τ]t\in[t_{2},\tau] and λ\lambda small enough.

Note that for λ\lambda small enough ⟨Djθt,wjt⟩≤−∥D−∥2\langle D_{j}^{\theta^{t}},\mathsf{w}_{j}^{t}\rangle\leq-\frac{\|D_{-}\|}{2} for any j∈S−,1j\in S_{-,1} and t∈[t−1,τ]t\in[t_{-1},\tau]. We thus have for any t∈[t2,τ]t\in[t_{2},\tau]

Moreover, recall that ⟨D+t,xk⟩>0\langle D_{+}^{t},x_{k}\rangle>0 before t2t_{2}. Thanks to Lemma 8, this actually implies that αjt2=0\alpha_{j}^{t_{2}}=0 and by Grönwall comparison: αjτ≤(4r+2∥D−∥)λε5\alpha_{j}^{\tau}\leq\left(4r+\frac{2}{\|D_{-}\|}\right)\lambda^{\frac{\varepsilon}{5}}.

Now note that ⟨Djθt,wjt⟩≤0\langle D_{j}^{\theta^{t}},\mathsf{w}_{j}^{t}\rangle\leq 0 for any t∈[τ,t3]t\in[\tau,t_{3}], which leads for any t∈[τ,t3]t\in[\tau,t_{3}] to

Equation 27 then becomes for any t∈[τ,t3]t\in[\tau,t_{3}] and λ\lambda small enough

Thanks to Lemma 16, this implies for any t∈[τ,t3]t\in[\tau,t_{3}] and λ\lambda small enough

Similarly to the first point, we can also show that t3−τ≤−ε8∥D−∥ln⁡(λ)t_{3}-\tau\leq-\frac{\varepsilon}{8\|D_{-}\|}\ln(\lambda), which finally yields that ∥D+t∥≤λε14\|D_{+}^{t}\|\leq\lambda^{\frac{\varepsilon}{14}} on [t2,t3][t_{2},t_{3}] for λ\lambda small enough. ∎

B.6 Final phase: proof of Theorem 1

To prove 1, we need to first prove some auxiliary lemmas. Lemma 12 shows that at time t3t_{3} the neural network is in the vicinity of some identifiable (i.e. independent of λ\lambda) interpolator. Lemmas 13, 14 and 15 allow to apply Chatterjee convergence result when a local PL is satisfied. Finally, we restate the main theorem in 2 for the sake of clearness and prove it.

Lemmas 10 and 11 directly imply Lemma 12 for some θλ∗\theta^{*}_{\lambda} that depends on λ\lambda. Extra work is required to prove that this minimal norm interpolator does not actually depend on λ\lambda.

Lemmas 10 and 11 imply the following properties

for any j∈S+,1j\in S_{+,1}, ∥wjt3−D+∥D+∥∥≤λε15\|\mathsf{w}_{j}^{t_{3}}-\frac{D_{+}}{\|D_{+}\|}\|\leq\lambda^{\frac{\varepsilon}{15}},

∣∑j∈S+,1∥wjt∥2−n∥D+∥∣≤λε15\left|\sum_{j\in S_{+,1}}\|w_{j}^{t}\|^{2}-n\|D_{+}\|\right|\leq\lambda^{\frac{\varepsilon}{15}},

for any j∈S−,1j\in S_{-,1}, ∥wjt3+D−∥D−∥∥≤2∥D−∥λε28\|\mathsf{w}_{j}^{t_{3}}+\frac{D_{-}}{\|D_{-}\|}\|\leq\sqrt{\frac{2}{\|D_{-}\|}}\lambda^{\frac{\varepsilon}{28}},

∣∑j∈S−,1∥wjt∥2−n∥D−∥∣≤λε29\left|\sum_{j\in S_{-,1}}\|w_{j}^{t}\|^{2}-n\|D_{-}\|\right|\leq\lambda^{\frac{\varepsilon}{29}},

for any j∉S+,1∪S−,1j\not\in S_{+,1}\cup S_{-,1}, ∥wjt∥≤3cλε\|w_{j}^{t}\|\leq 3c\lambda^{\varepsilon}.

Consider any pair j,j′∈S+,1j,j^{\prime}\in S_{+,1}. We recall Equation 9:

We note in the following ρ~jt,w~jt\widetilde{\rho}_{j}^{t},\widetilde{\mathsf{w}}_{j}^{t} any solutions of the ODEs

where D~jt=1n∑kykxk\mathds1⟨w~jt,xk⟩>0\widetilde{D}_{j}^{t}=\frac{1}{n}\sum_{k}y_{k}x_{k}\mathds{1}_{\langle\widetilde{\mathsf{w}}_{j}^{t},x_{k}\rangle>0}. Remark that the process (w~,ρ~)(\widetilde{\mathsf{w}},\widetilde{\rho}) does not depend on λ\lambda as ∥wj0∥=λgj\|w_{j}^{0}\|=\lambda g_{j} with the gjg_{j}’s being standard Gaussian vectors. We first have

For yk>0y_{k}>0, both \mathds1⟨w~jt,xk⟩>0\mathds{1}_{\langle\widetilde{\mathsf{w}}_{j}^{t},x_{k}\rangle>0} and \mathds1⟨wjt,xk⟩>0\mathds{1}_{\langle\mathsf{w}_{j}^{t},x_{k}\rangle>0} remain positive until t+1t_{+1}. As a consequence, the first sum is non-positive, which leads to

Since ∣hθt(xk)∣≤4mc2λ2(1−ε)|h_{\theta^{t}}(x_{k})|\leq 4mc^{2}\lambda^{2(1-\varepsilon)} during the first phase, Grönwall lemma implies that

We thus have ∥wjt+1−w~jt+1∥≤4mc2n(∥D+2+D−∥2)λ2−2(1+2)ε\|\mathsf{w}_{j}^{t_{+1}}-\widetilde{\mathsf{w}}_{j}^{t_{+1}}\|\leq\frac{4mc^{2}}{\sqrt{n\left(\|D_{+}^{2}+D_{-}\|^{2}\right)}}\lambda^{2-2(1+\sqrt{2})\varepsilon}. From there, note that we also have

The quantity ∣ρ~jt+1−ρ~j′t+1∣|\widetilde{\rho}_{j}^{t_{+1}}-\widetilde{\rho}_{j^{\prime}}^{t_{+1}}| only depends on λ\lambda because of the t+1t_{+1} term. The ρ~\widetilde{\rho} are indeed independent of λ\lambda.

Similarly to the proof of Lemma 7, we can show that after some time tjt_{j}, D~jt=D+\widetilde{D}_{j}^{t}=D_{+} for all t≥tjt\geq t_{j}. From there, Equation 28 imply for any t≥tjt\geq t_{j}

Solving these ODEs, we thus have for some constants τj,cj\tau_{j},c_{j} and any t≥tjt\geq t_{j}:

For λ\lambda small enough, t+1≥tj∨tj′t_{+1}\geq t_{j}\vee t_{j^{\prime}} and then:

where dj,j′=cj−cj′+∥D+∥(τj′−τj)d_{j,j^{\prime}}=c_{j}-c_{j^{\prime}}+\|D_{+}\|\left(\tau_{j^{\prime}}-\tau_{j}\right) and ∣hj,j′(λ)∣≤e2(τj∨τj′)λ2ε|h_{j,j^{\prime}}(\lambda)|\leq e^{2(\tau_{j}\vee\tau_{j^{\prime}})}\lambda^{2\varepsilon}. Using Equation 29, this leads for λ\lambda small enough to

where ∣gj,j′(λ)∣≤4λε15|g_{j,j^{\prime}}(\lambda)|\leq 4\lambda^{\frac{\varepsilon}{15}}. As the sum of the norms of the wjw_{j} is known, this actually fixes the norm of each individual neuron. Precisely we have for ∣f(λ)∣≤λε15|f(\lambda)|\leq\lambda^{\frac{\varepsilon}{15}}

And so we finally have for any j∈S+,1j\in S_{+,1}, ∥wjt3∥=n∥D+∥∑i∈S+,1e2di,j+O(λε15)\|w_{j}^{t_{3}}\|=\sqrt{\frac{n\|D_{+}\|}{\sum_{i\in S_{+,1}}e^{2d_{i,j}}}}+\mathcal{O}\left(\lambda^{\frac{\varepsilon}{15}}\right). Similarly, we can show that for any j∈S−,1j\in S_{-,1}, ∥wjt3∥=n∥D−∥∑i∈S−,1e2di,j+O(λε30)\|w_{j}^{t_{3}}\|=\sqrt{\frac{n\|D_{-}\|}{\sum_{i\in S_{-,1}}e^{2d_{i,j}}}}+\mathcal{O}\left(\lambda^{\frac{\varepsilon}{30}}\right). We can now define θ∗\theta^{*} as follows:

wj∗=n∥D+∥∑i∈S+,1e2di,jD+∥D+∥w_{j}^{*}=\sqrt{\frac{n\|D_{+}\|}{\sum_{i\in S_{+,1}}e^{2d_{i,j}}}}\frac{D_{+}}{\|D_{+}\|} and aj∗=∥wj∗∥a_{j}^{*}=\|w_{j}^{*}\| if j∈S+,1j\in S_{+,1}

wj∗=−n∥D−∥∑i∈S−,1e2di,jD−∥D−∥w_{j}^{*}=-\sqrt{\frac{n\|D_{-}\|}{\sum_{i\in S_{-,1}}e^{2d_{i,j}}}}\frac{D_{-}}{\|D_{-}\|} and aj∗=−∥wj∗∥a_{j}^{*}=-\|w_{j}^{*}\| if j∈S−,1j\in S_{-,1}

wj∗=0w_{j}^{*}=0 and aj∗=0a_{j}^{*}=0 if j∉S+,1∪S−,1j\not\in S_{+,1}\cup S_{-,1}.

For this final phase, we know that at time t3t_{3}, given Lemma 11, the neural net is arrived at a point θt3\theta^{t_{3}} satisfying the following: for ε′=ε30\varepsilon^{\prime}=\frac{\varepsilon}{30},

For all t>0t>0, θt∈Θ:={θ=(aj,wj)j≤m, such that ∀j∈⟦m⟧, ∣aj∣2=∥wj∥2}\theta^{t}\in\Theta:=\{\theta=(a_{j},w_{j})_{j\leq m},\textrm{ such that }\forall j\in\llbracket m\rrbracket,\ |a_{j}|^{2}=\|w_{j}\|^{2}\}.

For all j∈⟦m⟧∖(S+,1∪S−,1),∥wjt3∥≤λε′\displaystyle j\in\llbracket m\rrbracket\setminus\left(S_{+,1}\cup S_{-,1}\right),\|w_{j}^{t_{3}}\|\leq\lambda^{\varepsilon^{\prime}}.

We have ∣1n∑j∈S+,1∥wjt3∥2−∥D+∥∣≤λϵ′\left|\frac{1}{n}\sum_{j\in S_{+,1}}\|w_{j}^{t_{3}}\|^{2}-\|D_{+}\|\right|\leq\lambda^{\epsilon^{\prime}} and ∣1n∑j∈S−,1∥wjt3∥2−∥D−∥∣≤λϵ′\left|\frac{1}{n}\sum_{j\in S_{-,1}}\|w_{j}^{t_{3}}\|^{2}-\|D_{-}\|\right|\leq\lambda^{\epsilon^{\prime}}.

For all j∈S+,1j\in S_{+,1}, ∣⟨wjt3,D+⟩−∥D+∥∣≤λϵ′\left|\langle\mathsf{w}_{j}^{t_{3}},D_{+}\rangle-\|D_{+}\|\right|\leq\lambda^{\epsilon^{\prime}} and j∈S−,1j\in S_{-,1}, ∣⟨wjt3,D−⟩−∥D−∥∣≤λϵ′\left|\langle\mathsf{w}_{j}^{t_{3}},D_{-}\rangle-\|D_{-}\|\right|\leq\lambda^{\epsilon^{\prime}}.

Let us fist show an auxiliary Lemma that states that when these four conditions are satisfied, then the loss is almost zero.

For θ\theta such that conditions (i), (ii), (iii) and (iv) are satisfied then, for λ\lambda small enough,

Assume that θ\theta satisfies conditions (i), (ii), (iii) and (iv). Then, for kk such that yk>0y_{k}>0,

Using the (ii) property, on the one hand, the second term is upper bounded by mλ2ε′m\lambda^{2\varepsilon^{\prime}}, and, on the other hand, as ∥wj−D+∥D+∥∥2≤2λε′∥D+∥\|\mathsf{w}_{j}-\frac{D_{+}}{\|D_{+}\|}\|^{2}\leq 2\frac{\lambda^{\varepsilon^{\prime}}}{\|D_{+}\|}, we have by adding and subtracting D+∥D+∥\frac{D_{+}}{\|D_{+}\|} in the inner product of the first term

for λ\lambda small enough. And the same goes similarly for yk<0y_{k}<0. Hence,

We show a local estimate of the PL inequality when balancedness, i.e. (i), is assumed.

Indeed, thanks to the balancedness property we have the following calculation

where the last inequality is implied by the definition of the set S+,1S_{+,1}. As the same goes for S−,1S_{-,1} replacing the sum over the positive (yk)(y_{k})’s by the negative ones, we have that

This concludes the proof of the claimed inequality. ∎

Second we show that on a neighbourhood of θt3\theta^{t_{3}} intersected with Θ\Theta, the local PL constant can be lower bounded, where we recall Θ\Theta is the set of balanced parameters.

For λ\lambda small enough, we have the following lower bound on the PL constant

Indeed, fix any r>0r>0 and take θ∈B(θt3,r)∩Θ\theta\in B(\theta^{t_{3}},r)\cap\Theta. Let us denote by (wjt3)j(w_{j}^{t_{3}})_{j} the components of the hidden layer of θt3\theta^{t_{3}}. We have

and for r=λε′/8r=\lambda^{\varepsilon^{\prime}/8} and λ\lambda small enough such that λε′+2λε′/8mn(∥D+∥+λε′/2)≤∥D+∥/2\lambda^{\varepsilon^{\prime}}+2\lambda^{\varepsilon^{\prime}/8}\sqrt{\frac{m}{n}}\left(\sqrt{\|D_{+}\|}+\lambda^{\varepsilon^{\prime}/2}\right)\leq\|D_{+}\|/2, we have

and the exact same inequality stands for the sum over j∈S−,1j\in S_{-,1} with alignment vector D−D_{-}. ∎

Thanks to the lower bound given by Lemma 15, we can now conclude that the gradient flow will not go out B(θt3, λε′8)B(\theta^{t_{3}},\ \lambda^{\frac{\varepsilon^{\prime}}{8}}) and will converge exponentially fast to some θλ∞\theta^{\infty}_{\lambda}. Then we can take the limit λ→0\lambda\to 0 to characterise in this low initialisation regime the limit of the gradient flow.

The gradient flow (θt)t>0(\theta^{t})_{t>0} converges to some θλ∞\theta^{\infty}_{\lambda} of zero training loss , i.e L(θλ∞)=0L(\theta^{\infty}_{\lambda})=0.

There exists θ∗\displaystyle\theta^{*} such that we have the following limit:

and in its last phase, the gradient flow (θt)t≥t3(\theta^{t})_{t\geq t_{3}} stays in B(θ∗, λε240) ∩ ΘB(\theta^{*},\,\lambda^{\frac{\varepsilon}{240}})\,\cap\,\Theta for which the convergence is exponential.

From Lemma 15, Equation 32, as ε′=ε/30\varepsilon^{\prime}=\varepsilon/30 we have

and for λ\lambda small enough, L(θt3)/λε/120≤λε/60/λε/120=λε/120≤min⁡{∥D+∥,∥D−∥}L(\theta^{t_{3}})/\lambda^{\varepsilon/120}\leq\lambda^{\varepsilon/60}/\lambda^{\varepsilon/120}=\lambda^{\varepsilon/120}\leq\min\left\{\|D_{+}\|,\|D_{-}\|\right\}. Hence for r=λε/240r=\lambda^{\varepsilon/240},

Then, Theorem 2.1 of Chatterjee applies (at least a benign modification of it restricting the flow to Θ\Theta) and this shows that the gradient flow (θt)t≥t3(\theta^{t})_{t\geq{t_{3}}} stays in B(θt3, λε240)∩ΘB(\theta^{t_{3}},\ \lambda^{\frac{\varepsilon}{240}})\cap\Theta, and converges towards some θλ∞\theta^{\infty}_{\lambda} of zero loss at exponential speed. Furthermore, from Lemma 12, there exists θ∗\theta^{*}, independent of λ\lambda, and belonging to argminL(θ)=0∥θ∥2\textrm{argmin}_{L(\theta)=0}\|\theta\|^{2}, such that θt3∈B(θ∗, λε31)∩Θ\theta^{t_{3}}\in B(\theta^{*},\ \lambda^{\frac{\varepsilon}{31}})\cap\Theta. Hence, (θt)t≥t3∈B(θ∗, 2λε240)∩Θ(\theta^{t})_{t\geq{t_{3}}}\in B(\theta^{*},\ 2\lambda^{\frac{\varepsilon}{240}})\cap\Theta and finally: lim⁡λ→0lim⁡t→∞θt=lim⁡λ→0θλ∞=θ∗\lim_{\lambda\to 0}\lim_{t\to\infty}\theta^{t}=\lim_{\lambda\to 0}\theta^{\infty}_{\lambda}=\theta^{*}. ∎

B.7 Auxiliary Lemmas

This section stats the auxiliary Lemma 16 used in the proof of Lemma 11.

Suppose a non-negative function ff verifies for non-negative constants a,ba,b and cc

The previous inequality f(t)≤f(0)+ca+baF(t)f(t)\leq f(0)+\frac{c}{a}+\frac{b}{a}F(t) yields the lemma. ∎

Appendix C On global solutions of the minimum norm problem

In this section, we study the following optimisation problem:

An important remark is that the optimisation problem in Equation 34 is not convex because of the constraint set. Hence, the KKT conditions stated below are only necessary conditions: there exist real numbers (λk)k∈⟦n⟧(\lambda_{k})_{k\in\llbracket n\rrbracket} such that

In the orthonormal case, it is possible to solve explicitly the non-convex optimisation problem defined in (34). Recall the definition of balanced networks: Θ={(a,W) such that ∀j∈⟦m⟧,∣aj∣=∥wj∥}\Theta=\{(a,W)\text{ such that }\forall j\in\llbracket m\rrbracket,|a_{j}|=\|w_{j}\|\}. For θ∈Θ\theta\in\Theta, this means that there exists (sj)j∈⟦m⟧∈{−1,1}m(\mathsf{s}_{j})_{j\in\llbracket m\rrbracket}\in\{-1,1\}^{m} such that aj=sj∥wj∥a_{j}=\mathsf{s}_{j}\|w_{j}\|. Let us call S0θ={j∈⟦m⟧ ∣ wj=0}\mathsf{S}^{\theta}_{0}=\{j\in\llbracket m\rrbracket\,|\,w_{j}=0\}, S+θ=(S0θ)c∩{j∈⟦m⟧ ∣ sj=+1}\mathsf{S}^{\theta}_{+}=(\mathsf{S}^{\theta}_{0})^{c}\cap\{j\in\llbracket m\rrbracket\,|\,\mathsf{s}_{j}=+1\} and S−θ=(S0θ)c∩{j∈⟦m⟧ ∣ sj=−1}\mathsf{S}^{\theta}_{-}=(\mathsf{S}^{\theta}_{0})^{c}\cap\{j\in\llbracket m\rrbracket\,|\,\mathsf{s}_{j}=-1\}, where (S0θ)c:=⟦m⟧∖S0θ(\mathsf{S}^{\theta}_{0})^{c}:=\llbracket m\rrbracket\setminus\mathsf{S}^{\theta}_{0}. Note that the family S0θ\mathsf{S}^{\theta}_{0}, S+θ\mathsf{S}^{\theta}_{+} and S−θ\mathsf{S}^{\theta}_{-} form a partition of ⟦m⟧\llbracket m\rrbracket.

All KKT points of the problem (34) that are balanced, i.e. θ∈Θ\theta\in\Theta, are in fact global minimisers with objective value 2n(∥D+∥+∥D−∥)\displaystyle 2n(\|D_{+}\|+\|D_{-}\|). More precisely, we have the following description of the global minimisers of Equation 34:

First the constraint set is non-empty if m≥2m\geq 2: one can consider the hidden weights defined as W=(n∥D+∥D+,−n∥D−∥D−,0,…,0)W=(\sqrt{\frac{n}{\|D_{+}\|}}D_{+},-\sqrt{\frac{n}{\|D_{-}\|}}D_{-},0,\ldots,0) and the outputs a=(n∥D+∥,−n∥D−∥,0,…,0)a=(\sqrt{n\|D_{+}\|},-\sqrt{n\|D_{-}\|},0,\ldots,0) so that for all k∈⟦m⟧k\in\llbracket m\rrbracket, h(a,W)(xk)=⟨a,σ(Wxk)⟩=ykh_{(a,W)}(x_{k})=\langle a,\sigma(Wx_{k})\rangle=y_{k}. Hence the minimum is attained in the closed ball centred in the origin and of radius ∥(a,W)∥\|(a,W)\|. Moreover, the set C:={θ∣∀k≤n, hθ(xk)=yk}=∩k≤nCkC:=\{\theta|\forall k\leq n,\ h_{\theta}(x_{k})=y_{k}\}=\cap_{k\leq n}C_{k}, where each CkC_{k} is a closed subset as it is the pre-image of yky_{k} by the continuous function: θ↦hθ(xk)\theta\mapsto h_{\theta}(x_{k}). Hence the minimum is to be found in the intersection of a compact and a close subset, that is, a compact set overall. By continuity of the norm to be minimised over this compact set, there exists a global minimum.

Moreover, let (a∗,W∗)(a^{*},W^{*}) be a global minimiser of the optimisation problem. Note that if there exists j∈⟦m⟧j\in\llbracket m\rrbracket such that ∣aj∗∣2≠∥wj∗∥2|a^{*}_{j}|^{2}\neq\|w^{*}_{j}\|^{2}, then, as for c=∣aj∗∣/∥wj∗∥c=|a^{*}_{j}|/\|w^{*}_{j}\|, ∣aj∗∣2/c+c∥wj∗∥2<∣aj∗∣2+∥wj∗∥2|a^{*}_{j}|^{2}/c+c\|w^{*}_{j}\|^{2}<|a^{*}_{j}|^{2}+\|w^{*}_{j}\|^{2}, without changing the constraint set, we have found a strictly better minimum. Hence, for all j∈⟦m⟧j\in\llbracket m\rrbracket we have: ∣aj∗∣2=∥wj∗∥2|a^{*}_{j}|^{2}=\|w^{*}_{j}\|^{2}.

On the other side, for all k≤nk\leq n, hθ∗(xk)=ykh_{\theta^{*}}(x_{k})=y_{k}, and thus

From there, we deduce that λk\lambda_{k} and yky_{k} have the same sign and

Finally define W+∗=∑j∈S+θ∥wj∗∥2\displaystyle\mathsf{W}^{*}_{+}=\sum_{j\in\mathsf{S}^{\theta}_{+}}\|w_{j}^{*}\|^{2} and W−∗=∑j∈S−θ∥wj∗∥2\displaystyle\mathsf{W}^{*}_{-}=\sum_{j\in\mathsf{S}^{\theta}_{-}}\|w_{j}^{*}\|^{2}, we have

And the KKT condition (36) reads, e.g. for j∈S+θj\in\mathsf{S}^{\theta}_{+},

Hence, ∑k∣yk>0λk2=1\sum_{k|y_{k}>0}\lambda_{k}^{2}=1 and a similar reasoning on S−θ\mathsf{S}^{\theta}_{-} gives that ∑k∣yk<0λk2=1\sum_{k|y_{k}<0}\lambda_{k}^{2}=1. Overall this gives that ∑k∣yk>0yk2=(W+∗)2\sum_{k|y_{k}>0}y_{k}^{2}=(\mathsf{W}^{*}_{+})^{2} and ∑k∣yk<0yk2=(W−∗)2\sum_{k|y_{k}<0}y_{k}^{2}=(\mathsf{W}^{*}_{-})^{2}. And as S+θ,S−θ\mathsf{S}^{\theta}_{+},\mathsf{S}^{\theta}_{-} are a partition of ⟦m⟧\llbracket m\rrbracket,

Hence all KKT points that are balanced have the same objective value and hence are global minimisers of the objective. The description of the set of minimisers directly follows from the above proof. ∎

Appendix D On the existence of gradient flows

This section shows that a global solution of the ODE followed by the gradient flow only exists for the choice of subdifferential σ′(0)=0\sigma^{\prime}(0)=0. Precisely, the gradient flow follows the following ODE

where ∂f\partial f is the Clarke subdifferential of ff. The subdifferential of the loss is uniquely defined up to the choice of the subdifferential of the activation function at . If we consistently choose a fixed value σ0∈\sigma_{0}\in for the latter, we then have

Note that an empirical study of the influence of σ0\sigma_{0} on the dynamics has been conducted by Bertoin et al. . We have the following proposition, justifying the choice σ′(0)=0\sigma^{\prime}(0)=0.

The ODE (39) admits (at least) one solution on [0,∞)[0,\infty) if and only if σ0=0\sigma_{0}=0.

First of all note that ⟨wj0,xk⟩≠0\langle w_{j}^{0},x_{k}\rangle\neq 0 for all jj and kk. The Peano theorem then implies there exists a local solution of this ODE, i.e., there exists a time t0t_{0} such that Equation 39 admits a continuous solution θt\theta^{t} on [0,t0)[0,t_{0}). Assume now that t0=∞t_{0}=\infty. As long as ⟨wj0,xk⟩≠0\langle w_{j}^{0},x_{k}\rangle\neq 0 for all jj and kk, the analysis does not depend on the choice of σ0\sigma_{0}. We can thus show similarly to the proof of the first phase that for some time τ\tau and some j,kj,k: ⟨wjτ,xk⟩=0\langle w_{j}^{\tau},x_{k}\rangle=0. Moreover, still following the lines of the first phase, we have some δ>0\delta>0 such that ∣hθt(xk)∣<∣yk∣2|h_{\theta^{t}}(x_{k})|<\frac{|y_{k}|}{2} on [0,τ+δ][0,\tau+\delta]. Recall that

As a consequence, ⟨wjt,xk⟩\langle w_{j}^{t},x_{k}\rangle is monotone on [0,τ+δ][0,\tau+\delta], and thus decreasing. In particular, ⟨wjt,xk⟩≤0\langle w_{j}^{t},x_{k}\rangle\leq 0 on [τ,τ+δ][\tau,\tau+\delta]. Moreover, note that ⟨wjt,xk⟩\langle w_{j}^{t},x_{k}\rangle can not become (strictly) negative, as its derivative is as soon as it becomes negative. Indeed, if t−≔inf⁡{t∣⟨wjt,xk⟩<0}<τ+δt_{-}\coloneqq\inf\{t\mid\langle w_{j}^{t},x_{k}\rangle<0\}<\tau+\delta, then the scalar product has a zero derivative on (t−,τ+δ)(t_{-},\tau+\delta) and is thus constant, equal to by continuity, on this interval. So we finally have ⟨wjt,xk⟩=0\langle w_{j}^{t},x_{k}\rangle=0 on [τ,τ+δ][\tau,\tau+\delta] and θt\theta^{t} is then a solution of Equation 39 (almost everywhere) only if σ0=0\sigma_{0}=0.