Revisiting Small Batch Training for Deep Neural Networks

Dominic Masters, Carlo Luschi

Introduction

The use of deep neural networks has recently enabled significant advances in a number of applications, including computer vision, speech recognition and natural language processing, and in reinforcement learning for robotic control and game playing (LeCun et al., 2015; Schmidhuber, 2015; Goodfellow et al., 2016; Arulkumaran et al., 2017).

Deep learning optimization is typically based on Stochastic Gradient Descent (SGD) or one of its variants (Bottou et al., 2016; Goodfellow et al., 2016). The SGD update rule relies on a stochastic approximation of the expected value of the gradient of the loss function over the training set, based on a small subset of training examples, or mini-batch.

The recent drive to employ progressively larger batch sizes is motivated by the desire to improve the parallelism of SGD, both to increase the efficiency of current processors and to allow distributed processing on a larger number of nodes (Dean et al., 2012; Das et al., 2016). In contrast, the use of small batch sizes has been shown to improve generalization performance and optimization convergence (Wilson & Martinez, 2003; LeCun et al., 2012; Keskar et al., 2016). The use of small batch sizes also has the advantage of requiring a significantly smaller memory footprint, which affords an opportunity to design processors which gain efficiency by exploiting memory locality. This motivates a fundamental trade-off in the choice of batch size.

Hoffer et al. (2017) have shown empirically that it is possible to maintain generalization performance with large batch training by performing the same number of SGD updates. However, this implies a computational overhead proportional to the mini-batch size, which negates the effect of improved hardware efficiency due to increased parallelism.

In order to improve large batch training performance for a given training computation cost, it has been proposed to scale the learning rate linearly with the batch size (Krizhevsky, 2014; Chen et al., 2016; Bottou et al., 2016; Goyal et al., 2017). In this work, we suggest that the discussion about the scaling of the learning rate with the batch size depends on how the problem is formalized, and revert to the view of Wilson & Martinez (2003), that considers the SGD weight update formulation based on the sum, instead of the average, of the gradients over the mini-batch. From this perspective, using a fixed learning rate keeps the expectation of the weight update per training example constant for any choice of batch size. At the same time, as will be discussed in Section 2, holding the expected value of the weight update per gradient calculation constant while increasing the batch size implies a linear increase of the variance of the weight update with the batch size.

To investigate these issues, we have performed a comprehensive set of experiments for a range of network architectures. The results provide evidence that increasing the batch size results in both a degradation of the test performance and a progressively smaller range of learning rates that allows stable training, where the notion of stability here refers to the robust convergence of the training algorithm.

The paper is organized as follows. Section 2 briefly reviews the main work on batch training and normalization. Section 3 presents a range of experimental results on training and generalization performance for the CIFAR-10, CIFAR-100 and ImageNet datasets. Previous work has compared training performance using batch sizes of the order of 128–256 with that of very large batch sizes of up to 4096 (Hoffer et al., 2017) or 8192 (Goyal et al., 2017), while we consider the possibility of using a wider range of values, down to batch sizes as small as 2 or 4. Section 4 provides a discussion of the main results presented in the paper. Finally, conclusions are drawn in Section 5.

Background: Batch Training and Batch Normalization

We assume a deep network model with parameters θ\boldsymbol{\theta}, and consider the non-convex optimization problem corresponding to the minimization of the loss function L(θ)L(\boldsymbol{\theta}), with respect to θ\boldsymbol{\theta}. L(θ)L(\boldsymbol{\theta}) is defined as the sample average of the loss per training example Li(θ)L_{i}(\boldsymbol{\theta}) over the training set,

where MM denotes the size of the training set. The above empirical loss is used as a proxy for the expected value of the loss with respect to the true data generating distribution.

Batch gradient optimization originally referred to the case where the gradient computations were accumulated over one presentation of the entire training set before being applied to the parameters, and stochastic gradient methods were typically online methods with parameter update based on a single training example. Current deep learning stochastic gradient algorithms instead use parameter updates based on gradient averages over small subsets of the full training set, or mini-batches, and the term batch size is commonly used to refer to the size of a mini-batch (Goodfellow et al., 2016).

Optimization based on the SGD algorithm uses a stochastic approximation of the gradient of the loss L(θ)L(\boldsymbol{\theta}) obtained from a mini-batch B\mathcal{B} of mm training examples, resulting in the weight update rule

This implies that, when increasing the batch size, a linear increase of the learning rate η\eta with the batch size mm is required to keep the mean SGD weight update per training example constant.

This linear scaling rule has been widely adopted, e.g., in Krizhevsky (2014), Chen et al. (2016), Bottou et al. (2016), Smith et al. (2017) and Jastrzebski et al. (2017).

This implies that, adopting the linear scaling rule, an increase in the batch size would also result in a linear increase in the covariance matrix of the weight update η Δθ\eta\,\Delta\boldsymbol{\theta}. Conversely, to keep the scaling of the covariance of the weight update vector η Δθ\eta\,\Delta\boldsymbol{\theta} constant would require scaling η\eta with the square root of the batch size mm (Krizhevsky, 2014; Hoffer et al., 2017).

2 A Different Perspective on Learning Rate Scaling

We suggest that the discussion about the linear or sub-linear increase of the learning rate with the batch size is the purely formal result of assuming the use of the average of the local gradients over a mini-batch in the SGD weight update.

As discussed in Wilson & Martinez (2003), current batch training implementations typically use a weight correction based on the average of the local gradients according to equation (3), while earlier work on batch optimization often assumed a weight correction based on the sum of the local gradients (Wilson & Martinez, 2003). Using the sum of the gradients at the point θk\boldsymbol{\theta}_{k}, the SGD parameter update rule can be rewritten as

If we now consider a sequence of updates from the point θk\boldsymbol{\theta}_{k} with a batch size mm, from (5) the value of the weights at step k+nk+n is expressed as

while increasing the batch size by a factor nn implies that the weights at step k+1 k+1\, (corresponding to the same number of gradient calculations) are instead given by

From (7), (8), the end point values of the weights θk+n\boldsymbol{\theta}_{k+n} and θk+1\boldsymbol{\theta}_{k+1} after nmnm gradient calculations will generally be different, except in the case where one could assume

as also pointed out in Wilson & Martinez (2003) and Goyal et al. (2017). Therefore, for larger batch sizes, the update rule will be governed by progressively different dynamics for larger values of nn, especially during the initial part of the training, during which the network parameters are typically changing rapidly over the non-convex loss surface (Goyal et al., 2017). It is important to note that (9) is a better approximation if the batch size and/or the base learning rate are small.

The purely formal difference of using the average of the local gradients instead of the sum has favoured the conclusion that using a larger batch size could provide more ‘accurate’ gradient estimates and allow the use of larger learning rates. However, the above analysis shows that, from the perspective of maintaining the expected value of the weight update per unit cost of computation, this may not be true. In fact, using smaller batch sizes allows gradients based on more up-to-date weights to be calculated, which in turn allows the use of higher base learning rates, as each SGD update has lower variance. Both of these factors potentially allow for faster and more robust convergence.

3 Effect of Batch Normalization

The training of modern deep networks commonly employs Batch Normalization (Ioffe & Szegedy, 2015). This technique has been shown to significantly improve training performance and has now become a standard component of many state-of-the-art networks.

Batch Normalization (BN) addresses the problem of internal covariate shift by reducing the dependency of the distribution of the input activations of each layer on all the preceding layers. This is achieved by normalizing activation xix_{i} for each feature,

where μ^B\hat{\mu}_{\mathcal{B}} and σ^B2\hat{\sigma}_{\mathcal{B}}^{2} denote respectively the sample mean and sample variance calculated over the batch B\mathcal{B} for one feature. The normalized values x^i\hat{x}_{i} are then further scaled and shifted by the learned parameters γ\gamma and β\beta (Ioffe & Szegedy, 2015)

For the case of a convolutional layer, with a feature map of size p×qp\times q and batch size mm, the sample size for the estimate of μ^B\hat{\mu}_{\mathcal{B}} and σ^B2\hat{\sigma}_{\mathcal{B}}^{2} is given by m⋅p⋅qm\cdot p\cdot q, while for a fully-connected layer the sample size is simply equal to mm. For very small batches, the estimation of the batch mean and variance can be very noisy, which may limit the effectiveness of BN in reducing the covariate shift. Moreover, as pointed out in Ioffe (2017), with very small batch sizes the estimates of the batch mean and variance used during training become a less accurate approximation of the mean and variance used for testing. The influence of BN on the performance for different batch sizes is investigated in Section 3.

We observe that the calculation of the mean and variance across the batch makes the loss calculated for a particular example dependent on other examples of the same batch. This intrinsically ties the empirical loss (1) that is minimized to the choice of batch size. In this situation, the analysis of (7), (8) in Section 2.2 is only strictly applicable for a fixed value of the batch size used for BN. In some cases the overall SGD batch is split into smaller sub-batches used for BN. A common example is in the case of distributed processing, where BN is often implemented independently on each separate worker to reduce communication costs (e.g. Goyal et al. (2017)). For this scenario, the discussion relating to (7), (8) in Section 2.2 is directly applicable assuming a fixed batch size for BN. The effect of using different batch sizes for BN and for the weight update of the optimization algorithm will be explored in Section 3.6.

Different types of normalization have also been proposed (Ba et al., 2016; Salimans & Kingma, 2016; Arpit et al., 2016; Ioffe, 2017; Wu & He, 2018). In particular, Batch Renormalization (Ioffe, 2017) and Group Normalization (Wu & He, 2018) have reported improved performance for small batch sizes.

4 Other Related Work

Hoffer et al. (2017) have shown empirically that the reduced generalization performance of large batch training is connected to the reduced number of parameter updates over the same number of epochs (which corresponds to the same computation cost in number of gradient calculations). Hoffer et al. (2017) present evidence that it is possible to achieve the same generalization performance with large batch size, by increasing the training duration to perform the same number of SGD updates. Since from (3) or (5) the number of gradient calculations per parameter update is proportional to the batch size mm, this implies an increase in computation proportional to mm.

Jastrzebski et al. (2017) claim that both the SGD generalization performance and training dynamics are controlled by a noise factor given by the ratio between the learning rate and the batch size, which corresponds to linearly scaling the learning rate. The paper also suggests that the invariance to the simultaneous rescaling of both the learning rate and the batch size breaks if the learning rate becomes too large or the batch size becomes too small. However, the evidence presented in this paper only shows that the linear scaling rule breaks when it is applied to progressively larger batch sizes.

Smith et al. (2017) have recently suggested using the linear scaling rule to increase the batch size instead of decreasing the learning rate during training. While this strategy simply results from a direct application of (4) and (6), and guarantees a constant mean value of the weight update per gradient calculation, as discussed in Section 2.2 decreasing the learning rate and increasing the batch size are only equivalent if one can assume (9), which does not generally hold for large batch sizes. It may however be approximately applicable in the last part of the convergence trajectory, after having reached the region corresponding to the final minimum. Finally, it is worth noting that in practice increasing the batch size during training is more difficult than decreasing the learning rate, since it may require, for example, modifying the computation graph.

Batch Training Performance

The experiments have been performed on three different datasets: CIFAR-10, CIFAR-100 (Krizhevsky, 2009) and ImageNet (Krizhevsky et al., 2012). The CIFAR-10 and CIFAR-100 experiments have been run for different AlexNet and ResNet models (Krizhevsky et al., 2012; He et al., 2015a) exploring the main factors that affect generalization performance, including BN, data augmentation and network depth. Further experiments have investigated the training performance of the ResNet-50 model (He et al., 2015a) on the ImageNet dataset, to confirm significant conclusions on a more complex task.

For all the experiments, training has been based on the standard SGD optimization. While SGD with momentum (Sutskever et al., 2013) has often been used to train the network models considered in this work (Krizhevsky et al., 2012; He et al., 2015a; Goyal et al., 2017), it has been shown that the optimum momentum coefficient is dependent on the batch size (Smith & Le, 2017; Goyal et al., 2017). For this reason momentum was not used, with the purpose of isolating the interaction between batch size and learning rate.

For the experiments using the CIFAR-10 and CIFAR-100 datasets, we have investigated a reduced version of the AlexNet model of Krizhevsky et al. (2012), and the ResNet-8, ResNet-20 and ResNet-32 models as described in He et al. (2015a). The reduced AlexNet implementation uses convolutional layers with stride equal to 1, kernel sizes equal to ,numberofchannelsperlayerequalto, number of channels per layer equal to, max pool layers with 3×33\times 3 kernels and stride 2, and 256 hidden nodes per fully-connected layer. Unless stated otherwise, all the experiments have been run for a total of 82 epochs for CIFAR-10 and 164 epochs for CIFAR-100, with a learning rate schedule based on a reduction by a factor of 10 at 50% and at 75% of the total number of iterations. Weight decay has been also applied, with λ=5⋅10−4\lambda=5\cdot 10^{-4} for the AlexNet model (Krizhevsky et al., 2012) and λ=10−4\lambda=10^{-4} for the ResNet models (He et al., 2015a). In all cases, the weights have been initialized using the method described in He et al. (2015b).

The CIFAR datasets are split into a 50,000 example training set and a 10,000 example test set (Krizhevsky, 2009), from which we have obtained the reported results. The experiments have been run with and without augmentation of the training data. This has been performed following the procedure of Lin et al. (2013) and He et al. (2015a), which pads the images by 4 pixels on each side, takes a random 32×\times32 crop, and then randomly applies a horizontal flip.

The ImageNet experiments were based on the ResNet-50 model of He et al. (2015a), with preprocessing that follows the method of Simonyan & Zisserman (2014) and He et al. (2015a). As with the CIFAR experiments, the weights have been initialized based on the method of He et al. (2015b), with a weight decay parameter λ=10−4\lambda=10^{-4}. Training has been run for a total of 90 epochs, with a learning rate reduction by a factor of 10 at 30, 60 and 80 epochs, following the procedure of Goyal et al. (2017).

For all the results, the reported test or validation accuracy is the median of the final 5 epochs of training (Goyal et al., 2017).

2 Performance Without Batch Normalization

3 Performance With Batch Normalization

We now present experimental results with BN. As discussed in Section 2.3, BN has been shown to significantly improve the training convergence and performance of deep networks.

The results reported in Figures 9–9 indicate in each of the different cases a clear optimum value of base learning rate, which is only achievable for batch sizes m=16m=16 or smaller for CIFAR-10, and m=8m=8 or smaller for CIFAR-100.

4 Performance With Gradual Warm-up

As discussed in Section 2.4, warm-up strategies have previously been employed to improve the performance of large batch training (Goyal et al., 2017). The results reported in Goyal et al. (2017) have shown that the ImageNet validation accuracy for ResNet-50 for batch size m=256m=256 could be matched over the same number of epochs with batch sizes of up to 81928192, by applying a suitable gradual warm-up procedure. In our experiments, the same warm-up strategy has been used to investigate a possible similar improvement for the case of small batch sizes.

Following the procedure of Goyal et al. (2017), the implementation of the gradual warm-up corresponds to the use of an initial learning rate of η/32\eta/32, with a linear increase from η/32\eta/32 to η\eta over the first 5% of training. The standard learning rate schedule is then used for the rest of training. Here, the above gradual warm-up has been applied to the training of the ResNet-32 model, for both CIFAR-10 and CIFAR-100 datasets with BN and data augmentation. The corresponding performance results are reported in Figures 12 and 12.

As discussed in Section 2.4, the gradual warm-up helps to maintain consistent training dynamics in the early stages of optimization, where the approximation (9) is expected to be particularly weak. From Figures 12 and 12, the use of the gradual warm-up strategy proposed in Goyal et al. (2017) generally improves the training performance. However, the results also show that the best performance is still obtained for smaller batch sizes. For the CIFAR-10 results (Figure 12), the best test accuracy corresponds to batch sizes m=4m=4 and m=8m=8, although quite good results are maintained out to m=128m=128, while for CIFAR-100 (Figure 12) the best test performance is obtained for batch sizes between m=4m=4 and m=16m=16, with a steeper degradation for larger batch sizes.

5 ImageNet Performance

The ResNet-50 performance for the ImageNet dataset is reported in Figure 12, confirming again that the best performance is achieved with smaller batches. The best validation accuracy is obtained with batch sizes between m=16m=16 and m=64m=64.

Overall, the ImageNet results confirm that the use of small batch sizes provides performance advantages and reliable training convergence for an increased range of learning rates.

6 Different Batch Size for Weight Update and Batch Normalization

In the experimental results presented in Sections 3.2 to 3.5 (Figures 3–12), the same batch size has been used both for the SGD weight updates and for BN. We now consider the effect of using small sub-batches for BN, and larger batches for SGD. This is common practice for the case of data parallel distributed processing, where BN is often implemented independently on each individual worker, while the overall batch size for the SGD weight updates is the aggregate of the BN sub-batches across all workers.

The results reported in Figure 14 show a general performance improvement by reducing the overall batch size for the SGD weight updates, in line with the analysis of Section 2. We also note that, for the CIFAR-100 results, the best test accuracy for a given overall batch size is consistently obtained when even smaller batches are used for BN. This is against the intuition that BN should benefit from estimating the normalization statistics over larger batches. This evidence suggests that, for a given overall batch size and fixed base learning rate, the best values of the batch size for BN are typically smaller that the batch size used for SGD. Moreover, the results indicate that the best values of the batch size for BN are only weakly related to the SGD batch size, possibly even independent of it. For example, a BN batch size m=4m=4 or m=8m=8 appears to give best results for all SGD batch sizes tested.

The results presented in Figure 14 show that the performance with BN generally improves by reducing the overall batch size for the SGD weight updates. The collected data supports the conclusion that the possibility of achieving the best test performance depends on the use of a small batch size for the overall SGD optimization.

The plots also indicate that the best performance is obtained when smaller batches are used for BN in the range between 1616 and 6464, for all SGD batch sizes tested. From these results, the advantages of using a small batch size for both SGD weight updates and BN suggests that the best solution for a distributed implementation would be to use a small batch size per worker, and distribute both BN and SGD over multiple workers.

Discussion

The reported experiments have explored the training dynamics and generalization performance of small batch training for different datasets and neural networks.

Figure 15 summarizes the best CIFAR-10 results achieved for a selection of the network architectures discussed in Sections 3. The curves show that the best test performance is always achieved with a batch size below m=32m=32. In the cases without BN, the performance is seen to continue to improve down to batch size m=2m=2. This is in line with the analysis of Section 2.2, that suggests that, with smaller batch sizes, the possibility of using more up-to-date gradient information always has a positive impact on training. The operation with BN is affected by the issues highlighted in Section 2.3. However, the performance with BN appears to be maintained for batch sizes even as small as m=4m=4 and m=8m=8.

For the more difficult ImageNet training, Figure 12 shows that a similar trend can be identified. The best performance is achieved with batch sizes between m=16m=16 and m=64m=64, although for m=64m=64 SGD becomes acutely sensitive to the choice of learning rate. The performance for batch sizes m≤8m\leq 8 depends on the discussed effect of BN. As mentioned in Section 2.3, the issue of the BN performance for very small batch sizes can be addressed by adopting alternative normalization techniques like Group Normalization (Wu & He, 2018).

Goyal et al. (2017) have shown that the accuracy of ResNet-50 training on ImageNet for batch size 256256 could be maintained with batch sizes of up to 81928192 by using a gradual warm-up scheme, and with a BN batch size of 3232 whatever the SGD batch size. This warm-up strategy has been also applied to the CIFAR-10 and CIFAR-100 cases investigated here. In our experiments, the use of a gradual warm-up did improve the performance of the large batch sizes, but did not fully recover the performance achieved with the smallest batch sizes.

While the presented results advocate the use of small batch sizes to improve the convergence and accuracy, the use of small batch sizes also has the effect of reducing the computational parallelism available. This motivates the need to consider the trade-off between hardware efficiency and test performance.

We should also observe that the experiments performed in this paper refer to some of the most successful deep network architectures that have been recently proposed. These models have typically been designed under the presumption that large batches are necessary for good machine performance. The realisation that small batch training is preferable for both generalization performance and training stability should motivate re-examination of machine architectural choices which perform well with small batches; there exists at least the opportunity to exploit the small memory footprint of small batches. Moreover, it can be expected that small batch training may provide further benefits for model architectures that have inherent stability issues.

We have also investigated the performance with different batch sizes for SGD and BN, which has often been used to reduce communication in a distributed processing setting. The evidence presented in this work shows that in these circumstances the largest contributor to the performance was the value of the overall batch size for SGD. For the overall SGD batch sizes providing the best test accuracy, the best performance was achieved using values of batch size for BN between m=4m=4 and m=8m=8 for the CIFAR-10 and CIFAR-100 datasets, or between m=16m=16 and m=64m=64 for the ImageNet dataset. These results further confirm the need to consider the trade-off between accuracy and efficiency, based on the fact that to achieve the best test performance it is important to use a small overall batch size for SGD. In order to widely distribute the work, this implies that the best solution would be to distribute the implementation of both BN and SGD optimization over multiple workers.

Conclusions

We have presented an empirical study of the performance of mini-batch stochastic gradient descent, and reviewed the underlying theoretical assumptions relating training duration and learning rate scaling to mini-batch size.

The presented results confirm that using small batch sizes achieves the best training stability and generalization performance, for a given computational cost, across a wide range of experiments. In all cases the best results have been obtained with batch sizes m=32m=32 or smaller, often as small as m=2m=2 or m=4m=4.

With BN and larger datasets, larger batch sizes can be useful, up to batch size m=32m=32 or m=64m=64. However, these datasets would typically require a distributed implementation to avoid excessively long training. In these cases, the best solution would be to implement both BN and stochastic gradient optimization over multiple processors, which would imply the use of a small batch size per worker. We have also observed that the best values of the batch size for BN are often smaller than the overall SGD batch size.

The results also highlight the optimization difficulties associated with large batch sizes. The range of usable base learning rates significantly decreases for larger batch sizes, often to the extent that the optimal learning rate could not be used. We suggest that this can be attributed to a linear increase in the variance of the weight updates with the batch size.

Overall, the experimental results support the broad conclusion that using small batch sizes for training provides benefits both in terms of range of learning rates that provide stable convergence and achieved test performance for a given number of epochs.

References

Appendix A Additional Results with Batch Normalization

Figures 18 and 18 show the effect of the training length in number of epochs on the results of Section 3.3. The CIFAR-10 ResNet-32 experiments with data augmentation have been repeated for half the number of epochs (41 epochs) and for twice the number of epochs (164 epochs). In both cases, the learning rate reductions have been implemented at 50% and at 75% of the total number of epochs.

The collected results show a consistent performance improvement for increasing training length. For lower base learning rates, the improvements are more pronounced, both for very small batch sizes and for large batch sizes. However, for the base learning rates corresponding to the best test accuracy, the benefits of longer training are important mainly for large batch sizes, in line with Hoffer et al. (2017). Even increasing the training length, and therefore the total computational cost, the performance of large batch training remains inferior compared to the small batch performance.