Predicting Training Time Without Training

Luca Zancato, Alessandro Achille, Avinash Ravichandran, Rahul Bhotika, Stefano Soatto

Introduction

Say you are a researcher with many more ideas than available time and compute resources to test them. You are pondering to launch thousands of experiments but, as the deadline approaches, you wonder whether they will finish in time, and before your computational budget is exhausted. Could you predict the time it takes for a network to converge, before even starting to train it?

We look to efficiently estimate the number of training steps a Deep Neural Network (DNN) needs to converge to a given value of the loss function, without actually having to train the network. This problem has received little attention thus far, possibly due to the fact that the initial training dynamics of a randomly initialized DNN are highly non-trivial to characterize and analyze. However, in most practical applications, it is common to not start from scratch, but from a pre-trained model. This may simplify the analysis, since the final solution obtained by fine-tuning is typically not too far from the initial solution obtained after pre-training. In fact, it is known that the dynamics of overparametrized DNNs during fine-tuning tends to be more predictable and close to convex .

We therefore characterize the training dynamics of a pre-trained network and provide a computationally efficient procedure to estimate the expected profile of the loss curve over time. In particular, we provide qualitative interpretation and quantitative prediction of the convergence speed of a DNN as a function of the network pre-training, the target task, and the optimization hyper-parameters.

We use a linearized version of the DNN model around pre-trained weights to study its actual dynamics. In a similar technique is used to describe the learning trajectories of randomly initialized wide neural networks. Such an approach is inspired by the Neural Tangent Kernel (NTK) for infinitely wide networks . While we note that NTK theory may not correctly predict the dynamics of real (finite size) randomly initialized networks , we show that our linearized approach can be extended to fine-tuning of real networks in a similar vein to . In order to predict fine-tuning Training Time (TT) without training we introduce a Stochastic Differential Equation (SDE) (similar to ) to approximate the behavior of SGD: we do so for a linearized DNN and in function space rather than in weight space. That is, rather than trying to predict the evolution of the weights of the network (a DD-dimensional vector), we aim to predict the evolution of the outputs of the network on the training set (a N×CN\times C-dimensional vector, where NN is the size of the dataset and CC the number of network’s outputs). This drastically reduces the dimensionality of the problem for over-parametrized networks (that is, when NC≪DNC\ll D).

A possible limiting factor of our approach is that the memory requirement to predict the dynamics scales as O(DC2N2)O(DC^{2}N^{2}). This would rapidly become infeasible for datasets of moderate size and for real architectures (DD is in the order of millions). To mitigate this, we show that we can use random projections to restrict to a much smaller D0D_{0}-dimensional subspace with only minimal loss in prediction accuracy. We also show how to estimate Training Time using a small subset of N0N_{0} samples, which reduces the total complexity to O(D0 C2N02)O(D_{0}\,C^{2}N_{0}^{2}). We do this by exploiting the spectral properties of the Gram matrix of the gradients. Under mild assumptions the same tools can be used to estimate Training Time on a larger dataset without actually seeing the data.

To summarize, our main contributions are:

We present both a qualitative and quantitative analysis of the fine-tuning Training Time as a function of the Gram-Matrix Θ\Theta of the gradients at initialization (empirical NTK matrix).

We show how to reduce the cost of estimating the matrix Θ\Theta using random projections of the gradients, which makes the method efficient for common architectures and large datasets.

We introduce a method to estimate how much longer a network will need to train if we increase the size of the dataset without actually having to see the data (under the hypothesis that new data is sampled from the same distribution).

We test the accuracy of our predictions on off-the-shelf state-of-the-art models trained on real datasets. We are able to predict the correct training time within a 20% error with 95% confidence over several different datasets and hyperparameters at only a small fraction of the time it would require to actually run the training (30-45x faster in our experiments).

Related Work

Predicting the training time of a state-of-the-art architecture on large scale datasets is a relatively understudied topic. In this direction, Justus et al. try to estimate the wall-clock time required for a forward and backward pass on given hardware. We focus instead on a complementary aspect: estimating the number of fine-tuning steps necessary for the loss to converge below a given threshold. Once this has been estimated we can combine it with the average time for the forward and backward pass to get a final estimate of the wall clock time to fine-tune a DNN model without training it.

Hence, we are interested in predicting the learning dynamics of a pre-trained DNN trained with either Gradient Descent (GD) or Stochastic Gradient Descent (SGD). While different results are known to describe training dynamics under a variety of assumptions (e.g. ), in the following we are mainly interested on recent developments which describe the optimization dynamics of a DNN using a linearization approach. Several works suggest that in the over-parametrized regime wide DNNs behave similar to linear models, and in particular they are fully characterized by the Gram-Matrix of the gradients, also known as empirical Neural Tangent Kernel (NTK).

Under these assumptions, derive a simple connection between training time and spectral decomposition of the NTK matrix. However, their results are limited to Gradient Descend dynamics and to simple architectures which are not directly applicable to real scenarios. In particular, their arguments hinge on the assumption of using a randomly initialized very wide two-layer or infinitely wide neural network . We take this direction a step further, providing a unified framework which allows us to describe training time for both SGD and GD on common architectures.

Again, we rely on a linear approximation of the model, but while the practical validity of such linear approximation for randomly initialized state-of-the-art architectures (such as ResNets) is still discussed , we follow Mu et al. and argue that the fine-tuning dynamics of over-parametrized DNNs can be closely described by a linearization. We expect such an approximation to hold true since the network does not move much in parameters space during fine-tuning and over-parametrization leads to smooth and regular loss function around the pre-trained weights . Under this premise, we tackle both GD and SGD in an unified framework and build on to model training of a linear model using a Stochastic Differential Equation in function space. We show that, as also hypothesized by , linearization can provide an accurate approximation of fine-tuning dynamics and therefore can be used for training time prediction.

Predicting training time

In this section we look at how to efficiently approximate the training time of a DNN without actual training. By Training Time (TT) we mean the number of optimization steps – of either Gradient Descent (GD) or Stochastic Gradient Descent (SGD) – needed to bring the loss on the training set below a certain threshold.

To address both problems, building on top of in the Supplementary we prove the following result.

In the limit of small learning rate η\eta, the output on the training set of a linearized network ftlinf_{t}^{lin} trained with SGD evolves according to the following Stochastic Differential Equation (SDE):

where X\mathcal{X} is the set of training images, ∣B∣|B| the batch-size and dndn is a DD-dimensional Brownian motion. We have defined the Gram gradients matrix Θ\Theta (i.e., the empirical Neural Tangent Kernel matrix) and the covariance matrix Σ\Sigma of the gradients as follows:

where gi≡∇wf0(xi)g_{i}\equiv\nabla_{w}f_{0}(x_{i}). Note both Θ\Theta and Σ\Sigma only require gradients w.r.t. parameters computed at initialization.

The first term of eq. 1 is an ordinary differential equation (ODE) describing the deterministic part of the optimization, while the second stochastic term accounts for the noise. In Figure 2 (left) we show the qualitative different behaviour of the solution to the deterministic part of eq. 1 and the complete SDE eq. 1. While several related results are known in the literature for the dynamics of the network in weight space , note that eq. 1 completely characterizes the training dynamics of the linearized model by looking at the evolution of the output ftlin(X)f_{t}^{lin}(\mathcal{X}) of the model on the training samples – a N×CN\times C-dimensional vector – rather than looking at the evolution of the weights wtw_{t} – a DD-dimensional vector. When the number of data points is much smaller than the number of weights (which are in the order of millions for ResNets), this can result in a drastic dimensionality reduction, which allows easy estimation of the solution to eq. 1. Solving eq. 1 still comes with some challenges, particularly in computing Θ\Theta efficiently on large datasets and architectures. We tackle these in Section 4. Before that, we take a look at how different hyper-parameters and different pre-trainings affect the training time of a DNN on a given task.

Effective learning rate. From 1 we can gauge how hyper-parameters will affect the optimization process of the linearized model and, by proxy, of the original model it approximates. One thing that should be noted is that 1 assumes the network is trained with momentum m=0m=0. Using a non-zero momentum leads to a second order differential equation in weight space, that is not captured by 1. We can however, introduce heuristics to handle the effect of momentum: Smith et al. note that the momentum acts on the stochastic part shrinking it by a factor 1/(1−m)\sqrt{1/(1-m)}. Meanwhile, under the assumptions we used in 1 (small learning rate), we can show (see Supplementary Material) the main effect of momentum on the deterministic part is to re-scale the learning rates by a factor 1/(1−m)1/(1-m). Given these results, we define the effective learning rate (ELR) η^=η/(1−m)\hat{\eta}=\eta/(1-m) and claim that, in first approximation, we can simulate the effect of momentum by using η^\hat{\eta} instead of η\eta in eq. 1. In particular, models with different learning rates and momentum coefficients will have similar (up to noise) dynamics (and hence training time) as long as the effective learning rate η^\hat{\eta} remains the same. In Figure 2 we show empirically that indeed same effective learning rate implies similar loss curve. That similar effective learning rate gives similar test performance has also been observed in .

Batch size. The batch size appears only in the stochastic part of the equation, its main effect is to decrease the scale of the SDE noise term. In particular, when the batch size goes to infinity ∣B∣→∞|B|\to\infty we recover the deterministic gradient flow also studied by . Note that we need the batch size ∣B∣|B| to go to infinity, rather than being as large as the dataset since we assumed random batch sampling with replacement. If we assume extraction without replacement the stochasticity is annihilated as soon as ∣B∣=N|B|=N (see for a more in depth discussion).

2 Effect of pre-training on training time

We now use the SDE in eq. 1 to analyze how the combination of different pre-trainings of the model – that is, different w0w_{0}’s – and different tasks affect the training time. In particular, we show that a necessary condition for fast convergence is that the gradients after pre-training cluster well with respect to the labels. We conduct this analysis for a binary classification task with yi=±1y_{i}=\pm 1, but the extension is straightforward for multi-class classification, under the simplifying assumptions that we are operating in the limit of large batch size (GD) so that only the deterministic part of eq. 1 remains. Under these assumptions, eq. 1 can be solved analitically and the loss of the linearized model at time tt can be written in closed form as (see Supplementary Material):

The following characterization can easily be obtained using an eigen-decomposition of the matrix Θ\Theta.

Let S=∇wfw(X)T∇wfw(X)S=\nabla_{w}f_{w}(\mathcal{X})^{T}\nabla_{w}f_{w}(\mathcal{X}) be the second moment matrix of the gradients and let S=UΣUTS=U\Sigma U^{T} be the uncentered PCA of the gradients, where Σ=diag⁡(λ1,…,λn,0,…,0)\Sigma=\operatorname{diag}({\lambda_{1},\ldots,\lambda_{n},0,\ldots,0}) is a D×DD\times D diagonal matrix, n≤min⁡(N,D)n\leq\min(N,D) is the rank of SS and λi\lambda_{i} are the eigenvalues sorted in descending order. Then we have

where λkvk=(gi⋅uk)i=1N\lambda_{k}\mathbf{v}_{k}=(g_{i}\cdot\mathbf{u}_{k})_{i=1}^{N} is the NN-dimensional vector containing the value of the kk-th principal component of gradients gig_{i} and δy:=Y−f0(X)\delta\mathbf{y}:=\mathcal{Y}-f_{0}(\mathcal{X}).

Training speed and gradient clustering. We can give the following intuitive interpretation: consider the gradient vector gig_{i} as a representation of the sample xix_{i}. If the first principal components of gig_{i} are sufficient to separate the classes (i.e., cluster them), then convergence is faster (see Figure 3). Conversely, if we need to use the higher components (associated to small λk\lambda_{k}) to separate the data, then convergence will be exponentially slower. Arora et al. also use the eigen-decomposition of Θ\Theta to explain the slower convergence observed for a randomly initialized two-layer network trained with random labels. This is straightforward since the projection of a random vector will be uniform on all eigenvectors, rather than concentrated on the first few, leading to slower convergence. However, we note that the exponential dynamics predicted by do not hold for more general networks trained from scratch (see Section 6). In particular, eq. 5 mandates that the loss curve is always convex (it is sum of convex functions), which may not be the case for deep networks trained from scratch.

Efficient numerical estimation of training time

In 2 we have shown a closed form solution to the SDE in eq. 1 in the limit of large batch size, and for the MSE loss. Unfortunately, in general eq. 1 does not have a closed form expression when using the cross-entropy loss . A numerical solution is however possible, enabled by the fact that we describe the network training in function space, which is much smaller than weight space for over-parametrized models. The main computational cost is to create the matrix Θ\Theta in eq. 1 – which has cost O(DC2N2)O(DC^{2}N^{2}) – and to compute the noise in the stochastic term. Here we show how to reduce the cost of Θ\Theta to O(D0C2N2)O(D_{0}C^{2}N^{2}) for D0≪DD_{0}\ll D using a random projection approximation. Then, we propose a fast approximation for the stochastic part. Finally, we describe how to reduce the cost in NN by using only a subset N′<NN^{\prime}<N of samples to predict training time.

Random projection. To keep the notation uncluttered, here we assume w.l.o.g. C=1C=1. In this case the matrix Θ\Theta contains N2N^{2} pairwise dot-products of the gradients (a DD-dimensional vector) for each of the NN training samples (see eq. 2). Since DD can be very large (in the order of millions) storing and multiplying all gradients can be expensive as NN grows. Hence, we look at a dimensionality reduction technique. The optimal dimensionality reduction that preserves the dot-product is obtained by projecting on the first principal components of SVD, which however are themselves expensive to obtain. A simpler technique is to project the gradients on a set of D′D^{\prime} standard Gaussian random vectors: it is known that such random projections preserve (in expectation) pairwise product between vectors, and hence allow us to reconstruct the Gram matrix while storing only D′D^{\prime}-dimensional vector, with D′≪DD^{\prime}\ll D. We further increase computational efficiency using multinomial random vectors {-1,0,+1} as proposed in which further reduce the computational cost by avoiding floating point multiplications. In Figure 4 we show that the entries of Θ\Theta and its spectrum are well approximated using this method, while the computational time becomes much smaller.

Computing the noise. The noise covariance matrix Σ\Sigma is a D×DD\times D-matrix that changes over time. Both computing it at each step and storing it is prohibitive. Estimating Σ\Sigma correctly is important to describe the dynamics of SGD , however we claim that a simple approximation may suffice to describe the simpler dynamic in function space. We approximate ∇wf0lin(X)Σ1/2\nabla_{w}f^{\text{lin}}_{0}(\mathcal{X})\Sigma^{1/2} approximating Σ\Sigma with its diagonal (so that the we only need to store a DD-dimensional vector). Rather than computing the whole Σ\Sigma at each step, we estimate the value of the diagonal at the beginning of the training. Then, by exploiting eq. 3, we see that the only change to Σ\Sigma is due to ∇ftlinL\nabla_{f^{\text{lin}}_{t}}\mathcal{L}, whose norm decreases over time. Therefore we use the easy-to-compute ∇ftlinL\nabla_{f^{\text{lin}}_{t}}\mathcal{L} to re-scale our initial estimate of Σ\Sigma.

Larger datasets. In the MSE case from eq. 4, knowing the eigenvalues λk\lambda_{k} and the corresponding residual projections pk=(δy⋅vk)2p_{k}=(\delta\mathbf{y}\cdot\mathbf{v}_{k})^{2} we can predict in closed form the whole training curve. Is it possible to predict λk\lambda_{k} and pkp_{k} using only a subset of the dataset? It is known that the eigenvalues of the Gram matrix of Gaussian data follow a power-law distribution of the form λk=ck−s\lambda_{k}=ck^{-s}. Moreover, by standard concentration argument, one can prove that the eigenvalues should converge to a given limit as the number of datapoints increases. We verify that a similar power-law and convergence result also holds for real data (see Figure 4). Exploiting this result, we can estimate cc and ss from the spectrum computed on a subset of the data, and then predict the remaining eigenvalues. A similar argument holds for the projections pkp_{k}, which also follow a power-law (albeit with slower convergence). We describe the complete estimation in the Supplementary Material.

Results

We now empirically validate the accuracy of 1 in approximating the loss curve of an actual deep neural network fine-tuned on a large scale dataset. We also validate the goodness of the numerical approximations described in Section 4. Due to the lack of a standard and well established benchmark to test Training Time estimation algorithms we developed one with the main goal to closely resemble fine-tuning common practice for a wide spectrum of different tasks.

Experimental setup. We define training time as the first time the (smoothed) loss is below a given threshold. However, since different datasets converge at different speeds, the same threshold can be too high (it is hit immediately) for some datasets, and too low for others (it may take hundreds of epochs to be reached). To solve this, and have cleaner readings, we define a ‘normalized’ threshold as follows: we fix the total number of fine-tuning steps TT, and measure instead the first time the loss is within ϵ\epsilon from the final value at time TT. This measure takes into account the ‘asymptotic’ loss reached by the DNN within the computational budget (which may not be close to zero if the budget is low), and naturally adapts the threshold to the difficulty of the dataset. We compute both the real loss curve and the predicted training curve using 1 and compare the ϵ\epsilon-training-time measured on both. We report the absolute prediction error, that is ∣tpredicted−treal∣|t_{\text{predicted}}-t_{\text{real}}|. For all the experiments we extract 5 random classes from each dataset (Table 1) and sample 150 images (or the maximum available for the specific dataset). Then we fine-tuned ResNet18/34 using either GD or SGD.

Accuracy of the prediction. In Figure 1 we show TT estimates errors (for different ϵ∈{1,...,40}\epsilon\in\{1,...,40\}) under a plethora of different conditions ranging from different learning rates, batch sizes, datasets and optimization methods. For all the experiments we choose a multi-class classification problem with Cross Entropy (CE) Loss unless specified otherwise, and fixed computational budget of T=150T=150 steps both for GD and SGD. We note that our estimates are consistently within respectively a 13% and 20% relative error around the actual training time 95% of the times.

In Table 1 we describe the sensitivity of our estimates to different thresholds ϵ\epsilon both when our assumptions do and do not hold (high and low learning rates regimes). Note that a larger threshold ϵ\epsilon is hit during the initial convergence phase of the network, when a small number of iterations corresponds a large change in the loss. Correspondingly, the hitting time can be measured more accurately and our errors are lower. A smaller ϵ\epsilon depends more on correct prediction of the slower asymptotic phase, for which exact hitting time is more difficult to estimate.

Wall-clock run-time. In Figure 5 we show the wall-clock runtime of our training time prediction method compared to the time to actually train the network for TT steps. Our method is 30-40 times faster. Moreover, we note that it can be run completely on CPU without a drastic drop in performance. This allows to cheaply estimate TT and allocate/manage resources even without access to a GPU.

Effect of dataset distance. We note that the average error for Surfaces (Figure 6) is uniformily higher than the other datasets. This may be due to the texture classification task being quite different from ImageNet, on which the network is pretrained. In this case we can expect that the linearization assumption is partially violated since the features must adjust more during fine-tuning.

Effect of hyper-parameters on prediction accuracy. We derived 1 under several assumptions, importantly: small learning rate and wtw_{t} close to w0w_{0}. In Figure 6 (left) we show that indeed increasing the learning rate decreases the accuracy of our prediction, albeit the accuracy remains good even at larger learning rates. Fine-tuning on larger dataset makes the weights move farther away from the initialization w0w_{0}. In Figure 6 (right) we show that this slightly increases the prediction error. Finally, we observe in Figure 6 (center) that using a smaller batch size, which makes the stochastic part of 1 larger also slightly increases the error. This can be ascribed to the approximation of the noise term (Figure 4). On the other hand, in Figure 2 (right) we see that the effect of momentum on a fine-tuned network is very well captured by the effective learning rate (Section 3.1), as long as the learning rate is reasonably small, which is the case for fine-tuning. Hence the SDE approximation is robust to different values of the momentum. In general, we note that even when our assumptions are not fully met training time can still be approximated with only a slightly higher error. This suggest that point-wise proximity of the training trajectory of linear and real models is not necessary as long as their behavior (decay-rate) is similar (see also Supplementary Material).

Discussion and conclusions

We have shown that we can predict with a 13-20% accuracy the time that it will take for a pre-trained network to reach a given loss, in only a small fraction of the time that it would require to actually train the model. We do this by studying the training dynamics of a linearized version of the model – using the SDE in eq. 1 – which, being in the smaller function space compared to parameters space, can be solved numerically. We have also studied the dependency of training time from pre-training and hyper-parameters (Section 3.1), and how to make the computation feasible for larger datasets and architectures (Section 4).

While we do not necessarily expect a linear approximation around a random initialization to hold during training of a real (non wide) network, we exploit the fact that when using a pre-trained network the weights are more likely to remain close to initialization , improving the quality of the approximation. However, in the Supplementary Material we show that even when using a pre-trained network, the trajectories of the weights of linearized model and of the real model can differ substantially. On the other hand, we also show that the linearized model correctly predicts the outputs (not the weights) of the real model throughout the training, which is enough to compute the loss. We hypothesise that this is the reason why eq. 1 can accurately predict the training time using a linear approximation.

The procedure described so far can be considered as an open loop procedure meaning that, since we are estimating training time before any fine-tuning step is performed, we are not gaining any feedback from the actual training. How to perform training time prediction during the actual training, and use training feedback (e.g., gradients updates) to improve the prediction in real time, is an interesting future direction of research.

References

Predicting Training Time Without Training: Supplementary Material

In the Supplementary Material we give the pseudo-code for the training time prediction algorithm (Appendix A) together with implementation details, show additional results including prediction of training time using only a subset of samples, and comparison of real and predicted loss curves in a variety of conditions (Appendix C). Finally, we give proofs of all statements.

Appendix A Algorithm

We can compute the estimate on training time based also on the accuracy of the model: we straightforwardly modify the above algorithm and use the predictions ftlin(X)f_{t}^{\text{lin}}(\mathcal{X}) to compute the error instead of the loss (e.g. fig. 10).

We now briefly describe some implementations details regarding the numerical solution of ODE and SDE. Both of them can be solved by means of standard algorithms: in the ODE case we used LSODA (which is the default integrator in scipy.integrate.odeint), in the SDE case we used Euler-Maruyama algorithm for Ito equations.

We observe removing batch normalization (preventing the statistics to be updated) and removing data augmentation improve linearization approximation both in the case of GD and SGD. Interestingly data augmentation only marginally alters the spectrum of the Gram matrix Θ\Theta and has little impact on the linearization approximation w.r.t. batch normalization. observed similar effects but, differently from us, their analysis has been carried out using randomly initialized ResNets.

Appendix B Target datasets

Appendix C Additional Experiments

Prediction of training time using a subset of samples. In Section 4 we suggest that in the case of MSE loss, it is possible to predict the training time on a large dataset using a smaller subset of samples (we discuss the details in Appendix D). In Figure 7 we show the result of predicting the loss curve on a dataset of N=4000N=4000 samples using a subset of N=1000N=1000 samples. Similarly, in Figure 11 (top row) we show the more difficult example of predicting the loss curve on N=1000N=1000 samples using a very small subset of N0=100N_{0}=100 samples. In both cases we correctly predict that training on a larger dataset is slower, in particular we correctly predict the asymptotic convergence phase. Note in the case N0=100N_{0}=100 the prediction is less accurate, this is in part due to the eigenspectrum of Θ\Theta being still far from its limiting behaviour achieved for large number of data (see Appendix D).

Comparison of predicted and real error curve. In Figure 8 we compare the error curve predicted by our method and the actual train error of the model as a function of the number of optimization steps. The model is trained on a subset of 2 classes of CIFAR-10 with 150 samples. We run the comparison for both gradient descent (left) and SGD (right), using learning rate η=0.001\eta=0.001, momentum m=0m=0 and (in the case of SGD) batch size 100. In both cases we observe that the predicted curve is reasonably close to the actual curve, more so at the beginning of the training (which is expected, since the linear approximation is more likely to hold). We also perform an ablation study to see the effect of different approximation of SGD noise in the SDE in eq. 1. In Figure 8 (center) we estimate the variance of the noise of SGD at the beginning of the training, and then assume it is constant to solve the SDE. Notice that this predicts the wrong asymptotic behavior, in particular the predicted error does not converge to zero as SGD does. In Figure 8 (right) we rescale the noise as we suggest in Section 4: once the noise is rescaled the SDE is able to predict the right asymptotic behavior of SGD.

Prediction accuracy in weight space and function space. In Section 3 and Section 6 we argue that using a differential equation to predict the dynamics in function space rather than weight space is not only faster (in the over-parametrized case), but also more accurate. In Figure 9 we show empirically that solving the corresponding ODE in weight space leads to a substantially larger prediction error.

Point-wise similarity of predicted and observed loss curve. In some cases, we observe that the predicted and observed loss curves can differ. This is especially the case when using cross-entropy loss (Figure 10). We hypothesize that this may be due to improper prediction of the dynamics when the softmax output saturates, as the dynamic becomes less linear . However, the train error curve (which only depends on the relative order of the outputs) remains relatively correct. We should also notice that prediction of the ϵ\epsilon-training-time T^ϵ\hat{T}_{\epsilon} can be accurate even if the curves are not point-wise close. The ϵ\epsilon-training-time seeks to find the first time after which the loss or the error is within an ϵ\epsilon threshold. Hence, as long as the real and predicted loss curves have a similar asymptotic slope the prediction will be correct, as we indeed verify in Figure 10 (bottom).

Appendix D Prediction of training time on larger datasets

In Section 4 we suggest that, in the case of MSE loss, it is possible to predict the training time on a large dataset using a subset of the samples. To do so we leverage the fact that the eigenvalues of Θ\Theta follows a power-law which is independent on the size of the dataset for large enough sizes (see Figure 4, right). More precisely, from 2, we know that given the eigenvalues λk\lambda_{k} of Θ\Theta and the projections pk=δy⋅vkp_{k}=\delta\mathbf{y}\cdot\mathbf{v}_{k} it is possible to predict the loss curve using

Let Θ0\Theta_{0} be the Gram-matrix of the gradients computed on the small subset of N0N_{0} samples, and let Θ\Theta be the Gram-matrix of the whole dataset of size NN. Using the fact that, as we increase the number of samples, the eigenvalues (once normalized by the dataset size) converge to a fixed limit (Figure 4, right), we estimate the eigenvalues λk\lambda_{k} of Θ\Theta as follow: we fit the coefficients ss and cc of a power law λk=ck−s\lambda_{k}=ck^{-s} to the eigenvalues of Θ0\Theta_{0}, and use the same coefficients to predict the eigenvalues of Θ\Theta. However, we notice that the coefficient ss (slope of the power law) estimated using a small subset of the data is often smaller than the slope observed on larger datase (note in Figure 4 (right) that the curves for smaller datasets are more flat). We found that using the following corrected power law increases the precision of the prediction:

Empirically, we determined α∈[0.1,0.2]\alpha\in[0.1,0.2] to give a good fit over different combinations of NN and N0N_{0}. In Figure 11 (center) we compare the predicted eigenspectrum of Θ\Theta with the actual eigenspectrum of Θ\Theta .

The projections pkp_{k} follow a similar power-law – albeit more noisy (see Figure 11, right) – so directly fitting the data may give an incorrect result. However, notice that in this case we can exploit an additional constraint, namely that ∑kpk=∥δy∥2\sum_{k}p_{k}=\|\delta\mathbf{y}\|^{2} (∥δy∥2\|\delta\mathbf{y}\|^{2} is a known quantity: labels and initial model predictions on the large dataset). Let pk=δy⋅vkp_{k}=\delta\mathbf{y}\cdot\mathbf{v}_{k} and let pk′=δy⋅vk′p^{\prime}_{k}=\delta\mathbf{y}\cdot\mathbf{v}_{k}^{\prime} where vk\mathbf{v}_{k} and vk′\mathbf{v}_{k}^{\prime} are the eigenvectors of Θ\Theta and Θ0\Theta_{0} respectively. Fix a small k0k_{0} (in our experiments, k0=100k_{0}=100). By convergence laws , we have that pk′≃pkp^{\prime}_{k}\simeq p_{k} when k<k0k<k_{0}. The remaining tail of pkp_{k} for k>k0k>k_{0} must now follow a power-law and also be such that ∑kpk=∥δy∥2\sum_{k}p_{k}=\|\delta\mathbf{y}\|^{2}. This uniquely identify the coefficients of a power law. Hence, we use the following prediction rule for pkp_{k}:

where aa and bb are such that p^k0=pk0′\hat{p}_{k_{0}}=p^{\prime}_{k_{0}} and ∑kp^k=∥δy∥2\sum_{k}\hat{p}_{k}=\|\delta\mathbf{y}\|^{2}.

In Figure 11 (left), we use the approximated λ^k\hat{\lambda}_{k} and p^k\hat{p}_{k} to predict the loss curve on a dataset of N=1000N=1000 samples using a smaller subset of N0=100N_{0}=100 samples. Notice that we correctly predict that the convergence is slower on the larger dataset. Moreover, while training on the smaller dataset quickly reaches zero, we correctly estimate the much slower asymptotic phase on the larger dataset. Increasing both NN and N0N_{0} increases the accuracy of the estimate, since the eigenspectrum of Θ\Theta is closer to convergence: In Figure 7 we show the same experiment as Figure 11 with N0=1000N_{0}=1000 and N=4000N=4000. Note the increase in accuracy on the predicted curve.

Appendix E Effective learning rate

We now show that having a momentum term has the effect of increasing the effective learning rate in the deterministic part of eq. 1. A similar treatment of the momentum term is also in [28, Appendix D]. Consider the update rule of SGD with momentum:

If η\eta is small, the weights wtw_{t} will change slowly and we can consider gtg_{t} to be approximately constant on short time periods, that is gt+1=gg_{t+1}=g. Under these assumptions, the gradient accumulator ata_{t} satisfies the following recursive equation:

which is solved by (assuming a0=0a_{0}=0 as common in most implementations):

In particular, ata_{t} converges exponentially fast to the asymptotic value a∗=g/(1−m)a^{*}=g/(1-m). Replacing this asymptotic value in the weight update equation above gives:

Appendix F Proof of theorems

To describe SGD dynamics in function space we start from deriving the SDE in parameter space. In order to derive the SDE required to model SGD we will start describing the discrete update of SGD as done in .

where LB(θt)=L(fθt(XB),YB)\mathcal{L}^{B}(\theta_{t})=\mathcal{L}(f_{\theta_{t}}(\mathcal{X}^{B}),\mathcal{Y}^{B}) is the average loss on a mini-batch BB (for simplicity, we assume that BB is a set of indexes sampled with replacement).

The mini-batch gradient ∇θLB(θt)\nabla_{\theta}\mathcal{L}^{B}(\theta_{t}) is an unbiased estimator of the full gradient, in particular the following holds:

Where we defined the covariance of the gradients as:

and gi:=∇wf0(xi)g_{i}:=\nabla_{w}f_{0}(x_{i}). The first term in the covariance is the second order moment matrix while the second term is the outer product of the average gradient.

Following standard approximation arguments (see and references there in) in the limit of small learning rate η\eta we can approximate the discrete stochastic equation eq. 6 with the SDE:

Given this result, we are going now to describe how to derive the SDE for the output ft(X)f_{t}(\mathcal{X}) of the network on the train set X\mathcal{X}. Using Ito’s lemma (see and references there in), given a random variable θ\theta that evolves according to an SDE, we can obtain a corresponding SDE that describes the evolution of a function of θ\theta. Applying the lemma to fθ(X)f_{\theta}(\mathcal{X}) we obtain:

Using the fact that in our case the model is linearized, so fθ(x)f_{\theta}(x) is a linear function of θ\theta, we have that ∇θ2f(j)(x)=0\nabla^{2}_{\theta}f^{(j)}(x)=0 and hence A=0A=0. This leaves us with the SDE:

F.2 Proposition 2: Loss decomposition

Let ∇wfw(X)=VΛU\nabla_{w}f_{w}(\mathcal{X})=V\Lambda U be the singular value decomposition of ∇wfw(X)\nabla_{w}f_{w}(\mathcal{X}) where Λ\Lambda is a rectangular matrix (of the same size of ∇wfw(X)\nabla_{w}f_{w}(\mathcal{X})) containing the singular values {σ1,…,σN}\{\sigma_{1},\ldots,\sigma_{N}\} on the diagonal. Both UU and VV are orthogonal matrices. Note that we have

We now use the singular value decomposition to derive an expression for Lt\mathcal{L}_{t} in case of gradient descent and MSE loss (which we call LtL_{t}). In this case, the differential equation eq. 1 reduces to:

which is a linear ordinary differential equation that can be solved in closed form. In particular, we have:

Replacing this in the expression for the MSE loss at time tt we have:

Now recall that, by the properties of the matrix exponential, we have:

where e−2ΛΛTt=diag⁡(e−2ηλ1t,e−2ηλ2t,…)e^{-2\Lambda\Lambda^{T}t}=\operatorname{diag}(e^{-2\eta\lambda_{1}t},e^{-2\eta\lambda_{2}t},\ldots) with λk:=σk2\lambda_{k}:=\sigma_{k}^{2}. Then, defining δy=Y−f0(X)\delta\mathbf{y}=\mathcal{Y}-f_{0}(\mathcal{X}) and denoting with vk\mathbf{v}_{k} the kk-th column of VV we have:

Now let uk\mathbf{u}_{k} denote the kk-th column of UTU^{T} and gig_{i} the ii-th column of ∇wfw(X)T\nabla_{w}f_{w}(\mathcal{X})^{T} (that is, the gradient of the ii-th sample). To conclude the proof we only need to show that λkvk=(gi⋅uk)i=1N\lambda_{k}\mathbf{v}_{k}=(g_{i}\cdot\mathbf{u}_{k})_{i=1}^{N}. But this follows directly from the SVD decompostion ∇wfw(X)=VΛU\nabla_{w}f_{w}(\mathcal{X})=V\Lambda U, since then VΛ=∇wfw(X)UTV\Lambda=\nabla_{w}f_{w}(\mathcal{X})U^{T}.