Stop Wasting My Time! Saving Days of ImageNet and BERT Training with Latest Weight Averaging
Jean Kaddour
Introduction
The arsenal of deep learning methods (e.g., architectures, regularizers, pre-trainers, etc.) has been growing rapidly; for the last decade, thousands of them have been proposed yearly. Arguably, many are brittle and not as universally effective as initially claimed . One way to filter \saywhat really works is by testing methods on large datasets. For example, for vision tasks, methods that demonstrated success on ImageNet have often proven to be successful in other tasks .
Large datasets, however, require access to expensive multi-GPU machines to enable data parallelism and reasonable training durations. Less well-funded researchers do not have access to supercomputers, and lengthy training runs make quick, iterative experimentation of research ideas difficult. Simple task-, model-, and optimizer-agnostic methods that can be easily added to existing training pipelines and speed up training time have the potential to make deep learning research more accessible.
In the 90s, Polyak & Juditsky studied how to accelerate the convergence speed of stochastic gradient descent in the convex loss function regime. They proved that the running average of the model weights iterates converges to the loss minimizer asymptotically with the highest possible rate. When visualizing a convex loss function, the geometric intuition is simple: whenever the optimizer oscillates around a minimum, the average of the iterates will be closer to it.
However, in deep learning, loss functions are highly non-convex . Weight averaging has been mainly used to improve the model’s generalization performance at the end of or after training .
We revisit weight averaging applied to neural networks from a convergence speed perspective. Inspired by Li et al. , we focus on the middle stage of training: after the dramatic changes of the local loss landscape during the very first training steps but before the optimizer converges. Because the weights still undergo substantial change in that middle phase, averaging all collected models, e.g., by maintaining a moving average , can be sub-optimal. Therefore, we propose to average the latest checkpoints (each collected at the end of an epoch) throughout training, which we refer to as LAtest Weight Averaging (LAWA).
LAtest Weight Averaging (LAWA)
The key idea is to collect model checkpoints once at the end of each epoch in a queue. LAWA ’s solution at the end of epoch is .
Requirements include few training loop modifications, as shown in Algorithm 1, and additional memory. In practice, we store the checkpoints in RAM or on disk and only transfer them to the GPU once we want to evaluate . To improve the time complexity of the averaging operation, one can use a circular queue Coding interview preparers might remember this Leetcode problem..
The number of latest weights is a hyper-parameter, and we achieve good results across both experiments with default value , as shown in Figure 2. However, we observe that averaging too many checkpoints () results in worse performance.
The averaging coefficients can also follow a different pattern, e.g., in Appendix C, we experiment with an exponential moving average (assigning higher weights to the more recent checkpoints) and observe that this works worse than uniform coefficients. One can also learn the averaging coefficients , but for simplicity and to avoid additional computational costs, we do not do so.
The checkpoint saving frequency might be thought of as a hyper-parameter; however, in this work, we always set it to one epoch. When there is so much data that we are in a sub-one-epoch training regime , we may collect checkpoints every steps and need to tune . Another heuristic might be to collect a checkpoint whenever the validation loss has not improved .
If the network includes batch norm layers, then their statistics for are unknown. Prior work has suggested computing them by an inference pass through the training dataset. We do not observe a large effect of doing so compared to simply copying ’s statistics, possibly because we only average the latest weights instead of keeping one running average over many epochs .
Results
We run all experiments on a machine with 4x NVIDIA 3090s and report its wall-clock time.
We consider the ImageNet -classes classification task , which includes M training images and k validation images. To train a ResNet50 , we use the official PyTorch implementation and train for epochs using SGD with a momentum value of and a cosine learning rate schedule. Our 4-GPU machine takes ~min for one epoch. For , we re-compute the batch norm layer statistics with a full inference pass through the training dataset before evaluating .
In Figure 1(a), we observe that LAWA reaches a high accuracy dramatically faster, e.g., validation accuracy of around ~% (the final accuracy is ~%) is reached ~ epochs (~ hours) earlier than the baseline optimizer (SGD). However, we also note that its head start decreases towards the end of the training, and the highest reached accuracy is not reached much earlier. This observation raises the question of whether we can use LAWA to “jump forward” and continue the training from to reach the optimal accuracy faster, which we further discuss in Section 4.
2 Masked Language Modeling: RoBERTa-Base on WikiText-103
Next, we pre-train a (Ro)BERT(a)-Base model with masked language modeling (MLM) objective on the WikiText-103 dataset with M and k tokens for training and validation set, respectively. We follow the training recipe provided by fairseq : We train with Adam for epochs, using a batch size of , a polynomial learning rate decay with k warmup steps and a peak learning rate of . Our 4-GPU machine needs ~min for one training epoch.
In Figure 1(b), we report the training and validation MLM cross-entropy (CE) losses as a function of the number of training steps (as typically done in NLP). We observe that LAWA consistently improves the losses, and it reaches Adam’s final best validation loss ~ epochs ahead, saving ~ GPU hours. Interestingly, ’s final validation performance is noticeably better than ’s, confirming previous results on improved generalization obtained with weight averaging .
Future Work
It is tempting to think that we may “jump forward” training by applying the LAWA procedure and then continue training from there if some target accuracy has not been reached yet. One issue is that we would need to adjust the learning rate each time we “jump”. In practice, we may not know by how much (if at all) we accelerated the training progression. Hence, it remains unclear how to adjust a learning rate scheduler or the state variables of an adaptive optimizer.
k𝑘k scheduler.
In Figure 2, we observe that at different times, different values perform better; e.g., during the end of the training, higher performs better; motivating a scheduler for .
Accelerating training from the very beginning.
We focus on speeding up training during the middle stage of training: after the first training steps but long before the optimizer converges. The reason for that is that in the very early training phase, the gradient typically moves with large magnitudes until it converges to a smaller subspace of the loss function’s Hessian, in which it then remains over long periods of training (middle stage). We empirically confirm that averaging during the early phase worsens the baseline’s performance, as can be seen in Figure 4.
Combining LAWA with other acceleration techniques.
As we will discuss in the next section, there are several other techniques available to accelerate neural network training. For example, the SAM optimizer can accelerate training too , and Kaddour et al. show that SAM combined with weight averaging can further boost the final test performance.
Relationships between LAWA and optimization hyper-parameters.
For example, SGD becomes unstable for certain learning rates ; can we similarly characterize when LAWA is effective?
Applying other operations to a set of checkpoints.
For example, by learning a hyper-network that takes in one or more checkpoints and predicts the model parameters at later training stages.
When does it not work?
LAWA may not always cause speed-ups because Kaddour et al. reported some negative results on using weight averaging to improve the model’s final performance.
Related work
The idea of weight averaging is not novel; it has been studied widely in linear settings .
Szegedy et al. used weight averaging to create the GoogLeNet model, which, at that time, set a new state of the art in the ImageNet 2014 challenge . Izmailov et al. introduce Stochastic Weight Averaging (SWA), a weight averaging strategy starting from pre-trained models to move them to better-generalizing regions in the same loss basin. Kaddour et al. extensively study SWA’s effectiveness, including non-typical domains like graph-structured data, and suggest combining it with SAM to boost its final performance further. Wortsman et al. propose to average weights of multiple models with different hyper-parameter configurations. All three works (i) average weights toward the end or even after convergence, (ii) focus on the models’ final test performances, and (iii) incorporate one moving-averaged model, while we show in Figure 2 that too large can result in suboptimal results, especially at earlier training times.
This work is heavily inspired by Li et al. ’s Trainable Weight Averaging (TWA), who propose to learn averaging coefficients for training speed-ups. Concurrently, Guo et al. observe that running the SWA procedure multiple times accelerates convergence. In some sense, LAWA generalizes their procedure by keeping an average of the latest checkpoints instead of running SWA sequentially. Another related optimizer utilizing an auxiliary set of “fast weights” before updating the weights of interest is the Lookahead (LA) optimizer . We compare LAWA and LA in Appendix B.
Another line of work has shown that training data re-weighting can speed up training. For example, some re-weighting methods focus on proxy models , importance sampling or removing spurious correlations .
Acknowledgements
I thank Matt J. Kusner and Mingtian Zhang for feedback and fruitful discussions. I acknowledge support from the Engineering and Physical Sciences Research Council with grant number EP/S021566/1.
References
Appendix A Losses and Top-5 Accuracies
For completeness, we also plot the training and validation losses for both experiments and the top-5 accuracies for the ImageNet experiment.
Figure 3 shows similar speed up trends of LAWA over SGD as discussed in the main body (Figure 1).
Figure 4 shows the training and validation losses for RoBERTa-Base trained on WikiText103. Here, we also include losses during earlier stages of training and point out that during these more fluctuant phases, LAWA performs worse. We expect this behavior because previous works pointed out that the network undergoes dramatic changes in early phases .
Appendix B LAWA vs. Lookahead
We compare LAWA () against Lookahead on the moderately-sized CIFAR-100 dataset (50k training images) and train a ResNet34 . We follow commonly used hyper-parameters (see e.g., ), and train with SGD for epochs, using a batch size of , a momentum value of and a cosine learning rate scheduler with initial learning rate . For LA, we use and , as suggested by the authors for this particular CIFAR100 dataset.
Figure 5 shows the training/test accuracy/loss as a function of the number of epochs. LAWA reaches high test accuracy around epochs earlier than SGD/LA.
Initially, we started experimenting with this learning task before scaling up to larger datasets. Since we only observed slight but not dramatic improvements in LA over the baseline, we did not evaluate LA in the larger-scale ImageNet and BERT experiments. However, note that we apply LAWA to the SGD checkpoints; an interesting future direction can be to combine LAWA with LA, i.e., to average over checkpoints obtained with LA.
Appendix C Uniform vs. Exponentially Decayed Averaging Coefficients
We compare uniform (UNI, corresponding to by default) and exponentially-decaying (EXP) weight coefficients. We follow the same ResNet34 / CIFAR100 setup as in the previous section.
We set for both strategies. Figure 6 shows that UNI slightly outperforms EXP; however, the difference is not large.