Toward Communication Efficient Adaptive Gradient Method
Xiangyi Chen, Xiaoyun Li, Ping Li
Introduction
Distributed training has been proven to be a successful way of accelerating training large scale machine learning models, for example, training ultra-large-scale CTR (click-through rate) models in commercial search engines Fan et al. (2019); Zhao et al. (2020); Xu et al. (2021). With the advances of computing power and algorithmic design, one can now train models that need to be trained for days even weeks in the past within just a few minutes (You et al., 2020). When the computing power is high compared with the network bandwidth connecting different machines in distributed training, the training speed can be bottlenecked by the transmission of gradients and parameters. Such situation occurs increasingly more often in recent years due to the rapid growth in power of GPUs. Therefore, reducing communication overhead is gradually becoming an important research direction in distributed training (Alistarh et al., 2017; Wangni et al., 2018; Lin et al., 2018). In addition, a new training paradigm called Federated Learning (Konečnỳ et al., 2016; McMahan et al., 2017) was proposed recently, where models are trained distributively with mobile devices as workers and data holders. Consider the case where the data is stored and the model is trained on each user’s mobile device (e.g., Wang et al. (2020)). The existence of limited bandwidth necessitates the development of communication efficient training algorithms. Moreover, it is unpractical for every user to keep communicating with the central server, due to the power condition or wireless connection of the device. To cope with communication issues, an SGD-based algorithm with periodic model averaging called Federated Averaging is proposed in McMahan et al. (2017).
Federated learning extends the traditional parameter server setting, where the data are located on different workers and information is aggregated at a central parameter server to coordinate the training. In the parameter server setting, many effective communication reduction techniques were proposed such as gradient compression (Lin et al., 2018; Bernstein et al., 2018) and quantization (Alistarh et al., 2017; Wangni et al., 2018; Wen et al., 2017) for distributed SGD. In federated learning, one can substantially reduce communication cost by avoiding frequent transmission between local workers and central server. The workers train and maintain their own models locally, and the central server aggregates and averages the model parameters of all workers periodically. After averaging, new model parameter is fed back to each local worker, which starts another round of “local training + global averaging”. Some of the aforementioned communication reduction techniques can also be incorporated into Federated Averaging to further reduce the communication cost (Reisizadeh et al., 2020).
Despite the active efforts on improving algorithms based on periodic model averaging (Haddadpour et al., 2019; Reisizadeh et al., 2020; Haddadpour et al., 2020), the prototype algorithm is still SGD. On the other hand, we know that adaptive gradient methods such as AdaGrad (Duchi et al., 2011), Adam (Kingma and Ba, 2015) and AMSGrad (Reddi et al., 2019) often perform better than SGD when training neural nets, in terms of difficulty of parameter tuning and convergence speed at early stages. This motivates us to study adaptive gradient methods in federated learning.
Our contributions. We study how to incorporate adaptive gradient method into federated learning. Specifically, we show that unlike SGD, a naive combination of adaptive gradient methods and periodic model averaging results in algorithms that may fail to converge. Based on federated averaging and decentralized training, we propose an adaptive gradient method with communication cost sublinear in that is guaranteed to converge to stationary points with a rate of , where is the dimension of the problem, is the number of iterations, and is the number of workers (nodes). Our proposed method enjoys the benefit from both worlds: the fast convergence performance of adaptive gradient methods and low communication cost of federated learning.
Related Work
Federated learning. A classical framework for distributed training is the parameter server framework. In such a setting, a parameter server is used to coordinate training and the main computation (e.g., gradient computation) is offloaded to workers in parallel. For SGD under this framework (Recht et al., 2011; Li et al., 2014; Zinkevich et al., 2010; Zhao et al., 2020), the gradients can be computed by workers on their local data and sent to the parameter server which will aggregate the gradients and update the model parameters. Recently, a variant of the parameter server setting called federated learning (McMahan et al., 2017; Konečnỳ et al., 2016) draws great attention. One of the key features of federated learning is that workers are likely to be mobile devices which share a low bandwidth with the parameter server. Thus, communication cost plays a more important role in federated learning compared with the traditional parameter server setting. To reduce the communication cost, McMahan et al. (2017) proposed an algorithm called Federated Averaging which is a version of parallel SGD with local updates. In Federated Averaging, each worker updates their own model parameters locally using SGD, and the local models are synchronized by periodic averaging through the parameter server. The algorithm is also called local SGD or K-step SGD in some other papers (Yu et al., 2019; Stich, 2019; Zhou and Cong, 2018). Theoretically, it is proven in Yu et al. (2019) that local SGD can save a communication factor of while achieving the same convergence rate as vanilla SGD.
Adaptive gradient methods. Adaptive gradient methods usually refer to the class of gradient based optimization algorithms that adaptively update their learning rate (for each parameter coordinate) using historical gradients. Adaptive gradient methods such as Adam Kingma and Ba (2015), AdaGrad (Duchi et al., 2011), AdaDelta (Zeiler, 2012) are commonly used for training deep neural networks. It has been observed empirically that in many cases adaptive methods can outperform SGD or other methods in terms of convergence speed. There are also many variants trying to improve different aspects of these algorithms, e.g., Reddi et al. (2019); Keskar and Socher (2017); Luo et al. (2019); Chen et al. (2019b); Agarwal et al. (2019). The work Reddi et al. (2019) pointed out the divergence issue of Adam, and proposed AMSGrad algorithm for a fix. Moreover, increasing efforts are investigated into theoretical analysis of these algorithms Chen et al. (2019a); Ward et al. (2019); Li and Orabona (2019); Staib et al. (2019); Zhou et al. (2018); Zou et al. (2019); Zou and Shen (2018). In the federated learning setting, Xie et al. (2019) proposed a variant of AdaGrad, and Reddi et al. (2021) proposed a framework of adaptive gradient methods that includes variants of AdaGrad, Adam, and Yogi Zaheer et al. (2018), along with a convergence for the framework. There are also a few recent works trying to apply adaptive gradient methods in distributed optimization Xu et al. (2020); Chen et al. (2021). In this paper, we embed new adaptive gradient methods into federated learning and provide rigorous convergence analysis. The proposed algorithm can achieve the same convergence rate as its vanilla version, while enjoying the communication reduction brought by periodic model averaging.
Distributed Training with Periodic Model Averaging
In this section, we introduce our problem setting and the periodic model averaging framework for federated learning.
Notation. Throughout the paper, denotes model parameter at node and iteration . is element-wise division when and are vectors of the same dimension, and and denote element-wise multiplication and power, respectively. denotes the -th coordinate of vector .
In this paper, we consider the following formulation for distributed training, with works (nodes):
where can be considered as the averaged loss over data at worker and the function can only be accessed by the worker itself. For instance, for training neural nets, can be viewed as the average loss of data located at the -th node.
We consider the case where ’s might be nonconvex (e.g., deep nets). Our convergence analysis needs the following assumptions.
A1: Lipschitz property, is differentiable and -smooth, i.e., .
A4: Bounded gradient estimator, .
The assumptions A1, A2, A3 are standard in stochastic optimization. A4 is a little stronger than bounded variance assumption (A2), and is commonly used in analysis for adaptive gradient methods (Chen et al., 2019a; Ward et al., 2019) to simplify the convergence analysis by bounding possible adaptive learning rates.
2 Periodic Model Averaging
Recently, a new trend in algorithm design for distributed training is using periodic model averaging to reduce communication cost. This is motivated by the fact that in some circumstances where distributed optimization is used, the computation time is dominated by communication time. This phenomenon is significantly exacerbated in federated learning, where the bandwidth is relatively small (e.g. wireless networks on mobile devices).
Local SGD (i.e., Federated Averaging (Konečnỳ et al., 2016; Zhou and Cong, 2018)) is featured by the use of periodic model averaging with SGD (see Algorithm 1). Periodic model averaging can reduce the number of communication rounds and it is shown in Yu et al. (2019); Stich (2019) that by using periodic model averaging, one can achieve the same convergence rate as distributed SGD with a communication cost sublinear in . Note that, in practical scenarios, samples distribution for each node may not be i.i.d.. In the example of training on mobile devices, if the data on each device is collected from each single user, samples on a node will no longer be randomly drawn from the population.
Since local SGD is heavily used for training neural nets in federated learning, it is natural to consider using adaptive gradient methods in such setting to integrate advantageous aspects of adaptive gradient methods. In the remaining sections of this paper, we will study how to use periodic model averaging with adaptive gradient methods.
Adaptive gradient methods with periodic model averaging
In this section, we explore the possibilities of combining adaptive gradient method with periodic model averaging. We use AMSGrad (Reddi et al., 2019) as our prototype algorithm due to its nice convergence guarantee and superior empirical performance. The proposed scheme will be called local AMSGrad.
Similar to local SGD, the most straightforward way to combine AMSGrad with periodic model averaging works as follows:
2. The variables are averaged every iterations.
The algorithm’s pseudo code is shown in Algorithm 2.
Since Algorithm 2 is similar to Algorithm 1 except for the use of adaptive learning rate, given that Algorithm 1 is guaranteed to converge to stationary points, one may expect that Algorithm 2 is also guaranteed to converge. However, this is not necessarily the case. Algorithm 2 can fail to converge to stationary points, due to the possibility that the adaptive learning rates on different nodes are different. We show this possibility in Theorem 4.1.
There exists a problem where Algorithm 2 converges to non-stationary points no matter how small the stepsize is.
Proof: We prove by providing a counter example. Consider a simple 1-dimensional case where with
It is clear that has a unique stationary point at such that . Suppose , , and the initial point is for . Also suppose that , i.e., we average local parameters after every iteration. At , for the first node associated with , we have , and . For , we have and . Since in the naive method every node keeps its own learning rate, after the first update we have , . By Algorithm 2, we have , which heads towards the opposite direction of the true stationary point. We can then continue to show that for , , and , for . Therefore, we always update by , while updating and by . As a result, the averaged model parameter will head towards , instead of 0. The above argument can be trivially extended to arbitrary stepsize since the gradients do not change in the linear region of the function.
Given the example of divergence shown in the proof of Theorem 4.1, we know that a naive combination of periodic model averaging and adaptive gradient methods may not be valid even in a very simple case. By diving into the example where the algorithm fails, one can notice that the divergence is caused by the non-consensus of adaptive learning rates on different nodes. This suggests that we should keep the adaptive learning rate the same at different nodes. Next, we incorporate this idea into algorithm design to use shared adaptive learning rate on different nodes.
2 Local AMSGrad with shared adaptive learning rates
In the last section, we have showed an example where a naive combination of AMSGrad and periodic model averaging may diverge. The key divergence mechanism is due to the use of different adaptive learning rates on different nodes. A natural way to improve it is to force different nodes to have the same adaptive learning rate and we instantiate this idea in Figure 1 and Algorithm 3.
Compared with Algorithm 2, Algorithm 3 introduces a periodic averaging step for and updates at the server side, the same is used for local updates of different nodes. The intuition behind the design is that, since can be viewed as second moment estimation of the gradients and the average of is also an estimation of second moment, we expect the performance of the proposed method to be close to original AMSGrad. Note that this is not the only way to synchronize adaptive learning rate at different nodes, e.g., one can keep locally and use the average of obtained at the previous averaging step as the adaptive learning rate during local updates. With the synchronization of adaptive learning rates, the divergence example in Theorem 4.1 is no longer valid. Nevertheless, the convergence guarantee of the proposed algorithm is still not clear since it uses periodic averaging with adaptive learning rates and momentum. In the next section, we will establish the convergence guarantee of the proposed algorithm.
Convergence of Local AMSGrad
In this section, we analyze the convergence behavior of Algorithm 3. The main result is summarized in Theorem 5.1.
For Algorithm 3, if A1 - A4 are satisfied, define , set , we have for any ,
For Algorithm 3, if A1 - A4 are satisfied, set , for when and for when , we have
Again, the factor is due to the bounded coordinate-wise variance assumption A3, one can easily remove the dependency by assuming bounded total variance. The proof of Theorem 5.1 can be found in Appendix A.
Experiments
We compare the performance of local SGD, local AMSGrad, and naive local AMSGrad on a synthetic Gaussian mixture dataset (Sagun et al., 2017), the letter recogintion dataset (Frey and Slate, 1991) and the standard MNIST dataset. Experiments were conducted using the PaddlePaddle deep learning platform.
In the first experiment, we use the synthetic dataset (Gaussian cluster data) from Sagun et al. (2017). We use 10 isotropic 100 dimensional Gaussian distributions with different mean and same standard deviation to generate the data. The standard deviation of each dimension is 1. The mean of each cluster is generated from an isotropic Gaussian distribution with marginal for each dimension. The labels are the corresponding indices of the cluster from which the data are drawn. The model is a neural network with 2 hidden layers with 50 nodes per layer, and the activation function is ReLU for both layers. The batch size in training is 256. We use workers with each worker containing data from two classes. This assignment of data corresponds to a non-i.i.d. distribution of data on different nodes. The average local period is set to 10. We perform the learning rate search on a log scale, and increase the learning rate starting from 1e-6 until the algorithm diverges or the performance deteriorates significantly. Specifically, the maximum learning rate is 1 for local SGD and 1e-2 for both local AMSGrad and naive local AMSGrad.
We compare the performance of the algorithms with their best learning rate in Figure 2. It can be seen that local SGD and local AMSGrad perform very similarly. Naive local AMSGrad is worse than the other two algorithms by a small margin.
The performance of different algorithms with different learning rate is shown in Figure 3. It can be seen that all algorithms perform well with suitable learning rate. From these results, it seems that local AMSGrad do not have clear advantages over other algorithms. In particular, naive local AMSGrad performs not bad albeit it lacks convergence guarantee. We conjecture that this is due to that such a simple dataset cannot make the adaptive learning rate on different nodes differ significantly. In the next sets of experiments, we use more complicated real-world datasets to test the performance of different algorithms.
2 MNIST dataset (non-i.i.d. case)
In this section, we compare the algorithms on training a convolutional neural network (CNN) on MNIST. Similar to last set of experiments, we perform learning rate search starting from 1e-6 until the algorithm diverges or deteriorates significantly. The average period is set to 10. We set to be 1e-4 for both Adam and AMSGrad. The data is distributed on 5 nodes and each node contains data from two classes, and there is no overlap on labels between different nodes. We expect such an allocation of data can creates a highly non-i.i.d. data distribution, leading to significantly different adaptive learning rates on different nodes. The neural network in the experiments consists of 3 convolution+pooling layers with ReLu activation, followed by a 10 nodes fully connected layer with softmax activation. The first convolution+pooling layer has 20 5x5 filters followed by 2x2 max pooling with stride 2. The second and third convolution+polling layer has 50 filters with other parameters being the same as the first layer.
Figure 4 shows the training and testing performance of different algorithms. Specifically, 1e-3 is chosen for local AMSGrad and naive local AMSGrad while 1e-4 is chosen for local SGD. It can be seen that local AMSGrad outperforms local SGD by a large margin.
The performance of algorithms with different learning rate is shown in Figure 5. We observe that naive local AMSGrad has very poor performance with all learning rate choices, and local AMSGrad tend to perform better than local SGD on average. The slow convergence of local SGD is also observed in McMahan et al. (2017) when the data distribution is non-i.i.d. While the sampling on nodes in McMahan et al. (2017) may somehow reduce the influence of non-i.i.d. data, the convergence speed of local SGD is significantly impacted by the non-i.i.d. distribution in our experiment since we always use all nodes for parameter update.
3 Letter recognition dataset (i.i.d. case)
In this section, we use the letter recognition dataset Frey and Slate (1991) to test the performance of different algorithms. We train a fully connected neural network with two hidden layers on the dataset. The first hidden layer has 300 nodes and the second layer has 200 nodes, both of which use ReLU as activation. The learning rate search strategy and other parameter settings are the same as the MNIST experiments. Different from the previous two sets of experiments, the data on the 5 workers are randomly assigned. This corresponds to an i.i.d. data distribution. Thus, all algorithms are expected to work well in this set of experiments. The performance comparison of algorithms with their best learning rate is provided in Figure 6. We can see that all algorithms achieve over 90% accuracy and local AMSGrad again performs the best, with 2% higher test accuracy than local SGD. In this case, naive local AMSGrad is also better than local SGD. Thus, when data distribution is i.i.d., it might be okay to use the naive version of local AMSGrad in practice, considering that it involves even less communication. The performance of algorithms with different learning rate is shown in Figure 7, which shows that all algorithms perform quite stable, while the two adaptive gradient methods still outperform local SGD in general.
Conclusion
In this paper, we study how to design adaptive gradient methods for federated learning by utilizing periodic model averaging. We first construct counter examples to illustrate how a naive combination of adaptive gradient methods and periodic model averaging can fail to converge. Then, by utilizing the insights from the study of non-convergence, we propose an adaptive gradient method local AMSGrad in the setting of federated learning with proved convergence guarantee. Local AMSGrad enjoys the sublinear communication cost of periodic model averaging as well as the superb empirical performance of adaptive gradient methods. Experiments show that local AMSGrad can often significantly outperform both the local SGD and a naive design of local adaptive gradient methods, especially when the dada distribution on different nodes is non-i.i.d.
Appendix A Proof of Theorem 5.1
To prove the convergence of local AMSGrad, we first define an auxiliary sequence of iterates
where and we define .
We have the following property for the new sequence .
where and .
By the updating rule of Algorithm 3, we have . Thus,
In what follows, we use the auxiliary sequence in Lemma A.1 to prove convergence of the algorithm.
where the expectation is taken over all the randomness of stochastic gradients until iteration .
It remains to upper bound the second-order term on RHS of (4) and characterize the effective descent in the first order term (LHS of (4)). We first characterize the effective descent.
By (3), we can write the first-order term as
where .
Using the fact , we have
where the first two quantities on RHS of (6) will contribute to the descent of objective in a single optimization step, while the last term is the possible ascent introduced by the bias on the stochastic gradients. The bound of the last term is given by Lemma A.2.
where the last inequality is due to Cauchy-Schwartz.
Using Lipschitz property (Assumption A1) of , we can further bound the differences of gradients on RHS of (8) by
It remains to bound and using the update rule of and . For the difference between and , we have
where . For the second term containing the consensus error , we have
For iterates produced by Algorithm 3, we have
Proof: Let be the largest multiple of that is less than . By the updating rule of Algorithm 3, we have
because , and , and .
Combining (8), (9), (10), (A) and (12), we obtain
which is the desired bound. This completes the proof.
Substituting (13) into (A) and (4) yields
Summing over from 1 to and divide both sides by , we get
What remains is to bounded and .
In addition, conditioned on all randomness in (gradients from iteration until ), we have
where the last inequality is due to non-decreasing property of .
where the second inequality is because is non-decreasing.
Substituting the bounds on and into (A), we obtain
Further, by choosing , we have
and when , we have
At this point, we have obtained the convergence rate (when is sufficiently large), which matches the convergence rate of SGD. One remaining item is to convert the convergence measure to the norm of gradients of . We do this by the following Lemma.
where the first inequality is due to Cauchy-Schwartz, the second inequality is due to Jensen’s inequality, and the last inequality is due to Lemma A.3 and L-smoothness of (A1).
by Lemma A.5. Multiplying both sides of the above inequality by 8 and using the fact that complete the overall proof.