Momentum Benefits Non-IID Federated Learning Simply and Provably
Ziheng Cheng, Xinmeng Huang, Pengfei Wu, Kun Yuan
Introduction
Federated learning (FL) is a powerful paradigm for large-scale machine learning (Konečnỳ et al., 2016; McMahan et al., 2017a). In situations where data and computational resources are dispersed among a diverse range of clients, including phones, tablets, sensors, hospitals, and other devices and agents, federated learning facilitates local data processing and collaboration among these clients (Kairouz et al., 2021). Consequently, a centralized model can be trained without transmitting decentralized data from clients directly to servers, thereby ensuring a fundamental level of privacy.
Federated learning encounters several significant challenges in algorithmic development. Firstly, the reliability and relatively slow nature of network connections between the server and clients pose obstacles to efficient communication during the training process. Secondly, the dynamic availability of only a small subset of clients for training at any given time demands strategies that can adapt to this variable environment. Lastly, the presence of substantial heterogeneity of non-iid data across different clients further complicates the training process.
FedAvg (Konečnỳ et al., 2016; McMahan et al., 2017a; Stich, 2019; Yu et al., 2019a; Lin et al., 2020; Wang & Joshi, 2021) has emerged as a prevalent algorithm for federated learning, leveraging multiple stochastic gradient descent (SGD) steps within each client before communicating with a central server. While FedAvg is readily implementable and has demonstrated success in certain applications, its performance is notably hindered by the presence of data heterogeneity, i.e., non-iid clients, even when all clients participate in the training process (Li et al., 2019; Yang et al., 2021). To mitigate the influence of data heterogeneity, SCAFFOLD (Karimireddy et al., 2020b) maintains a control variable on each client to compensate for “client drift” in its local SGD updates, making convergence more robust to data heterogeneity and client sampling. Due to their practicality and effectiveness, FedAvg and SCAFFOLD have become foundational algorithms in federated learning, leading to the development of numerous variants that cater to decentralized (Koloskova et al., 2020; Rizk et al., 2022; Nguyen et al., 2022; Alghunaim, 2023), compressed (Haddadpour et al., 2021; Reisizadeh et al., 2020; Mitra et al., 2021), asynchronous (Chen et al., 2020a, b; Xu et al., 2021a), and personalized (Fallah et al., 2020; Pillutla et al., 2022; Tan et al., 2022; T Dinh et al., 2020) federated learning scenarios.
Various methods have been proposed to enhance the convergence of FedAvg, SCAFFOLD, and their variance-reducedThroughout the paper, variance reduction refers to techniques aiming to mitigate the influence of within-client gradient stochasticity, as opposed to the inter-client data heterogeneity. extensions. While exhibiting superior convergence rates, these approaches typically make impractical adjustments to algorithmic structures. For instance, STEM (Khanduri et al., 2021) requires increasing either the batch size or the number of local steps with algorithmic iterations. Similarly, CE-LSGD (Patel et al., 2022) and MIME (Karimireddy et al., 2020a) mandate computing a large-batch or even full-batch local gradient per round for each client. Additionally, FedProx (Li et al., 2020), FedPD (Zhang et al., 2021), and FedDyn (Durmus et al., 2021) rely on solving “local problems” to an extremely high precision. These adjustments may not align with the practical constraints in federated learning setups.
Furthermore, many of these algorithms, including FedAvg, STEM, FedProx, MIME, and CE-SGD, still rely on the assumption of bounded data heterogeneity. When this assumption is violated, the theoretical analyses of these algorithms become invalid. While some algorithms, such as LED (Alghunaim, 2023) and VRL-SGD (Liang et al., 2019), can handle unbounded data heterogeneity, their convergence rates are not state-of-the-art, as demonstrated in Table 1. These limitations motivate us to develop novel strategies that are easy to implement, robust to data heterogeneity, and exhibit superior theoretical convergence rates.
This paper examines the utilization of momentum to enhance the performance of FedAvg and SCAFFOLD.
In order to ensure simplicity and practicality in implementations, we only introduce momentum to the local SGD steps, avoiding any inclusion of impractical elements, such as gradient computation of large batchsizes or solving local problems to high precision. Remarkably, this straightforward approach effectively alleviates the necessity for stringent assumptions of bounded data heterogeneity, leading to noteworthy improvements in convergence rates. The main findings and contributions of this paper are summarized below.
First, when all clients participate in the training process:
We demonstrate that incorporating momentum allows FedAvg and its variance-reduced extension to converge without relying on the assumption of bounded data heterogeneity, even using constant local learning rates. This is rather surprising as, to our knowledge, all existing analyses for FedAvg, e.g., (Karimireddy et al., 2020b; Yang et al., 2021; Wang et al., 2020b), require bounded data heterogeneity even with diminishing local learning rates.
We further establish that, by effectively removing the influence of data heterogeneity on convergence, momentum empowers FedAvg and its variance-reduced extension with state-of-the-art convergence rates in the context of full client participation.
Second, when partial clients participate in the training process per iteration:
The proposed SCAFFOLD-M that incorporates momentum into SCAFFOLD achieves a provably faster convergence rate. To our knowledge, this is the first result that improves upon SCAFFOLD without imposing any additional assumptions beyond those used in (Karimireddy et al., 2020b).
We further introduce momentum to SCAFFOLD with variance reduction, achieving the first variance-reduced federated learning algorithm that does not rely on the assumption of bounded data heterogeneity. This algorithm attains a state-of-the-art convergence rate in the context of partial client participation and unbounded data heterogeneity.
Tables 1 and 2 present a comprehensive comparison of the convergence rates and associated assumptions of existing algorithms, as well as our newly proposed approaches.
It is observed that by simply adding momentum to local steps, FedAvg, SCAFFOLD, and their variance-reduced extensions all attain state-of-the-art convergence rates without resorting to further assumptions such as bounded data heterogeneity. We support our theoretical findings with extensive numerical experiments.
2 Related work
FedAvg is a well-known algorithm introduced by (McMahan et al., 2017b) as a heuristic to enhance communication efficiency and data privacy in federated learning. Numerous subsequent studies have focused on analyzing its convergence under the assumption of homogeneous datasets, where clients are independent and identically distributed (iid) and all clients participate fully (Stich, 2019; Yu et al., 2019b; Wang & Joshi, 2021; Lin et al., 2020; Zhou & Cong, 2017). However, when dealing with heterogeneous clients and partial client participation, FedAvg is found to be vulnerable to data heterogeneity because of the ”client drift” effect (Karimireddy et al., 2020b; Yang et al., 2021; Wang et al., 2020b; Li et al., 2019).
Considerable research efforts have been devoted to mitigating the impact of data heterogeneity in federated learning. For example, Li et al. (2020) propose FedProx, which introduces a proximal term to the objective function. Yang et al. (2021) utilize a two-sided learning rate approach, while Wang et al. (2020a) propose FedNova, a normalized averaging method. Additionally, Zhang et al. (2021) present FedPD, which addresses data heterogeneity from a primal-dual optimization perspective. Notably, Karimireddy et al. (2020b) introduces SCAFFOLD, an effective algorithm that employs control variables to mitigate the influence of data heterogeneity and partial client participation. FedGate (Haddadpour et al., 2021) and LED (Alghunaim, 2023) are two recent effective algorithms that have alleviated the impact of data heterogeneity, utilizing gradient tracking (Xu et al., 2015; Di Lorenzo & Scutari, 2016; Pu & Nedić, 2020; Xin et al., 2020; Alghunaim & Yuan, 2021) and exact-diffusion (Yuan et al., 2019, 2020, 2021a) techniques, respectively.
The momentum mechanism dates back to Nesterov’s acceleration (Yurri, 2004) and Polyak’s heavy-ball method (Polyak, 1964) in deterministic optimization, which later flourishes in the stochastic scenario (Yan et al., 2018; Yu et al., 2019a; Liu et al., 2020) and other communication efficient algorithms (Yuan et al., 2021b; He et al., 2023b, a). Extensive research has explored incorporating momentum into federated learning (Reddi et al., 2021; Wang et al., 2020b; Karimireddy et al., 2020a; Khanduri et al., 2021; Patel et al., 2022; Das et al., 2022; Yu et al., 2019a), and numerous empirical studies have demonstrated its substantial impact on enhancing the performance of federated learning algorithms (Wang et al., 2020b; Xu et al., 2021b; Reddi et al., 2021; Jin et al., 2022; Kim et al., 2022). However, whether momentum can offer theoretical benefits to federated learning remains underexplored. This work demonstrates that momentum can improve non-iid federated learning simply and provably. It is worth noting that the theoretical utility of momentum has been demonstrated in various scenarios beyond federated learning. For instance, Guo et al. (2021) proved that momentum can correct the bias experienced by the Adam method, while a very recent work (Fatkhullin et al., 2023) demonstrated that momentum can improve the error feedback technique in communication compression. The analysis presented in this work is different from (Guo et al., 2021) and (Fatkhullin et al., 2023) due to the unique challenges encountered in federated learning including multiple local updates, data heterogeneity, and partial client participation.
Problem setup
This section formulates the problem of non-iid federated learning. Formally, we consider minimizing the following objective with the fewest number of client-server communication rounds:
Here, the random variable represents a local datapoint available at client , while the function denotes the non-convex local loss function associated with client . This function takes expectation with respect to the local data distribution . In practice, the local data distributions among different clients typically differ from each other, resulting in the inequality for any pair of nodes and . This phenomenon is commonly referred to as data heterogeneity. If all local clients were homogeneous, meaning that all local data samples follow the same distribution , we would have for any and . In addition, throughout the paper we assume that the function is bounded from below and possesses a global minimum . To facilitate convergence analysis, we also introduce the following standard assumptions.
It is worth noting that Assumption 2 implies Assumption 1, which is typically used in variance-reduced algorithms, e.g., (Karimireddy et al., 2020a; Khanduri et al., 2021; Fang et al., 2018; Cutkosky & Orabona, 2019). We will utilize either Assumption 1 or 2 in different algorithms.
Accelerating FedAvg with momentum
This section focuses on the scenario in which all clients participate in the training process of federated learning. We will introduce momentum to both FedAvg and its variance-reduced extension. Additionally, we will demonstrate that the incorporation of momentum effectively eliminates the impact of data heterogeneity, leading to improved convergence rates.
where is the momentum coefficient, and represents an global gradient estimate updated in the outer loop . It is important to note that FedAvg-M will reduce to the vanilla FedAvg when . Furthermore, FedAvg-M is easy to implement, as it maintains the same algorithmic structure and incurs no additional uplink communication overhead compared to FedAvg.
The inclusion of momentum in FedAvg yields notable theoretical improvements. Firstly, it eliminates the need for the data heterogeneity assumption, also known as the gradient similarity assumption. The assumption can be expressed as
where measures the magnitude of data heterogeneity. By incorporating momentum, the above assumption is no longer required for the convergence analysis of FedAvg. Secondly, momentum enables FedAvg to converge at a state-of-the-art rate. These improvements are justified as follows:
Under Assumption 1 and 3, if we set , ,
where and , then FedAvg-M satisfies
where notation denotes inequalities that hold up to a numeric number.
Table 1 compares FedAvg-M with existing algorithms when all clients participate in the training process. The results demonstrate that FedAvg-M attains the most favorable convergence rate without relying on any assumption of data heterogeneity. Moreover, this rate matches the lower bound provided by (Arjevani et al., 2019).
Based on Theorem 3.1, it can be inferred that when , FedAvg-M allows the utilization of constant local learning rate which does not decay as the number of communication rounds increases. This characteristic eases the tuning of the local learning rate and improves empirical performance. In contrast, many existing convergence results of FedAvg necessitate the adoption of local learning rates that diminish as increases, as exemplified by e.g., (Yang et al., 2021; Li et al., 2019; Karimireddy et al., 2020b; Koloskova et al., 2020).
The momentum mechanism relies on an accumulated gradient estimate , which, although biased, exhibits reduced variance due to its accumulation nature compared to the stochastic gradient computed with a single data batch. Importantly, by utilizing directions for local updates, an “anchoring” effect is achieved, effectively mitigating the “client-drift” phenomenon. In the extreme case where , all clients remain synchronized in their local updates, eliminating any drift. By appropriately tuning the coefficient , FedAvg-M maintains the same convergence rate as (Yang et al., 2021) while removing the requirement of data heterogeneity assumption utilized in their analysis.
2 Variance-reduced FedAvg with momentum
When each local loss function is further assumed to be sample-wise smooth (i.e., Assumption 2), we can replace the local descent direction in Algorithm 1 with a variance-reduced momentum direction
to further enhance convergence, leading to variance-reduced FedAvg with momentum, or FedAvg-M-VR for short, see the detailed algorithm in Appendix B.2. The variable is the last-iterate global model maintained in the server. The construction of the variance-reduced direction equation 3.1 effectively mitigates the influence of within-client gradient noise and can be traced back to SARAH (Nguyen et al., 2017) and STORM (Cutkosky & Orabona, 2019) in stochastic optimization; more discussion can be found in (Tan et al., 2022). Same as FedAvg-M, turning off the variance-reduced momentum of FedAvg-M-VR, i.e., setting , recovers FedAvg. FedAvg-M-VR shares the same algorithmic structure and uplink communication workload as FedAvg.
Under Assumption 2 and 3, if we take with and set , , , and
FedAvg-M-VR surpasses all existing variance-reduced federated learning methods in terms of convergence rate, as demonstrated by the results presented in Table 1. Additionally, when compared to BVR-L-SGD (Murata & Suzuki, 2021) and CE-LSGD (Patel et al., 2022), FedAvg-M-VR computes each local stochastic gradient using a batchsize of , contrasting with the batchsize employed by BVR-L-SGD and CE-LSGD. Furthermore, in comparison to STEM (Khanduri et al., 2021), FedAvg-M-VR does not rely on the assumption of bounded data heterogeneity.
Based on discussions in Sections 3.1 and 3.2, we demonstrate that FedAvg-M and FedAvg-M-VR, in the context of full client participation, can achieve the state-of-the-art convergence rate without resorting to any stronger assumption, e.g., bounded data heterogeneity or impractical algorithmic structures such as large batchsizes.
Accelerating SCAFFOLD with momentum
This section addresses the scenario where a random subset of clients participates in the training process per iteration. To tackle the challenges arising from partial participation, SCAFFOLD employs a control variable in each client to counteract the “client drift” effect during local updates. To further enhance the convergence performance, we will introduce momentum to both SCAFFOLD and its variance-reduced extension. Through our analysis, we will demonstrate that the incorporation of momentum results in new state-of-the-art convergence rates for these algorithms.
We introduce momentum to enhance the estimation of the stochastic gradient, resulting in the newly proposed algorithm SCAFFOLD-M, outlined in Algorithm 2. In SCAFFOLD-M, clients are randomly selected from a pool of clients for each iteration of trainining. The control variables and are maintained by the client and server, respectively. In SCAFFOLD, the local descent direction is given by . In contrast, SCAFFOLD-M incorporates momentum directions for local updates:
where represents the global stochastic gradient vector maintained by the server. It is worth noting that SCAFFOLD-M can reduce to SCAFFOLD by setting .
The inclusion of momentum in SCAFFOLD yields notable theoretical improvements, as justified by the following theorem.
Under Assumption 1 and 3, if we take , with , and set , , , , then SCAFFOLD-M converges as
Compared to SCAFFOLD, SCAFFOLD-M exhibits provably faster convergence under partial participation, as demonstrated in the comparison presented in Table 2. Specifically, when the gradients are noiseless (i.e., ), achieving the same level of precision requires a ratio, between SCAFFOLD-M and SCAFFOLD, of communication rounds given by
where notation indicates that the equality holds up to a numeric number. Consequently, if , SCAFFOLD-M achieves up to times improvement in comparison to the vanilla SCAFFOLD, when aiming for the same precision. This improvement is particularly significant as , the number of clients, is typically very large. It is also worth noting that prior to the introduction of SCAFFOLD-M, SCAFFOLD was the only known non-iid federated learning method, to the best of our knowledge, that is robust to both unbounded data heterogeneity and partial client sampling, and capable of attaining linear speedup without relying on impractical algorithmic structures. Consequently, the development of SCAFFOLD-M provides an alternative and superior choice.
2 Variance-reduced SCAFFOLD with momentum
Similar to FedAvg-M-VR, when the loss functions further enjoy the sample-wise smoothness property, we can obtain SCAFFOLD-M-VR by replacing momentum directions in Algorithm 2 with variance-reduced momentum directions
The detailed algorithm is in Appendix C.2, and the convergence is shown below.
Under Assumption 2 and 3, if we take with , and set , , , , then SCAFFOLD-M-VR converges as
SCAFFOLD-M-VR outperforms all existing variance-reduced federated learning methods under partial participation in terms of convergence rate when data heterogeneity is severe (i.e., is large), see results listed in Table 2. Moreover, SCAFFOLD-M-VR has the following additional advantages. Compared to MimeLiteMVR (Karimireddy et al., 2020a), SCAFFOLD-M-VR does not need access to noiseless (full-batch) local gradients per iteration. Compared to MB-STORM (Patel et al., 2022) and CE-LSGD (Patel et al., 2022), SCAFFOLD-M-VR does not require bounded data heterogeneity and computes each gradient efficiently with batchsize , as opposed to batchsize .
Based on discussions in Sections 4.1 and 4.2, we demonstrate that SCAFFOLD-M and SCAFFOLD-M-VR, in the context of partial client participation, can achieve state-of-the-art convergence rates without resorting to any stronger assumption, e.g., bounded data heterogeneity or impractical algorithmic structures such as large batchsize.
Experiments
We conducted an evaluation of our proposed methods using a three-layer fully connected neural network trained on the CIFAR-10 dataset. To generate non-iid data for the clients, we sample label ratios from the Dirichlet distribution (Hsu et al., 2019) with a parameter of for the full participation setting and for the partial participation setting. Our experimental setup involves clients and local updates. The weight decay is set as . The global learning rate is fixed as for all the algorithms, and we performe a grid search for the local learning rate in values . Similarly, we search for the momentum parameter in values .
Our experiments can be categorized into three parts.
Firstly, we compare the performance of FedAvg-M and SCAFFOLD-M with their momentumless counterparts, namely the vanilla FedAvg and SCAFFOLD, under full client participation. The results are presented in Figure 1(a), where it can be observed that incorporating momentum significantly accelerates the convergence of both FedAvg and SCAFFOLD.
Secondly, we compare three momentum-based variance-reduced methods: CE-LSGD, FedAvg-M-VR, and SCAFFOLD-M-VR, under the condition of full client participation. The comparison is illustrated in Figure 1(b). It is evident that our proposed methods outperform CE-LSGD with substantial margins.
Lastly, we investigate the partial participation setting with and compare the performance of SCAFFOLD-M and SCAFFOLD-M-VR with vanilla SCAFFOLD. The results are presented in Figure 1(c). Once again, we observe that the introduction of momentum leads to significant improvements even when only a few clients participate in each round of training.
Conclusion
This paper proposes momentum variants of FedAvg and SCAFFOLD under various client participation situations and smoothness properties. All of our momentum variants only make simple and practical modifications to FedAvg and SCAFFOLD yet obtain state-of-the-art performance among their peers, particularly when data heterogeneity is severe or gradient noise is trivial. In particular, FedAvg-M converges without relying on bounded data heterogeneity and can adopt constant local learning rates, giving the first neat convergence guarantee for FedAvg-type methods; SCAFFOLD-M is the first FL method that outperforms SCAFFOLD unconditionally. Experiments are conducted in the paper to support our theoretical findings.
References
Appendix A Preliminaries of proofs
Under Assumption 1, if , the following inequality holds for all :
Since is -smooth, we have
Since , using Young’s inequality, we further have
where the last inequality holds due to . Taking the global expectation completes the proof. ∎
To handle local updates and client sampling, we will also use the following technical lemmas.
for any ,
.
Letting be the indicator for the event , we prove this lemma as follows:
In the following subsections, we present complete proofs of our main results. For FedAvg-M and SCAFFOLD-M, our proofs only rely on Assumption 1 and 3, while for FedAvg-M-VR and SCAFFOLD-M-VR, our proofs rely on Assumption 2 and 3.
Appendix B FedAvg with momentm
If , the following holds for :
Additionally, it holds for that
Note that are sequentially correlated. Applying the AM-GM inequality and Lemma A.3, we have
Using the AM-GM inequality again and Assumption 1, we have
where we plug in and use in the last inequality. Similarly, for ,
If , the following holds for :
For any , using , we have
where the last inequality holds by unrolling the recursive bound and using . By Lemma A.3, it holds that for ,
This is also valid for . Summing up over and finishes the proof. ∎
If , then it holds for that
Note that ,
Using Young’s inequality, we have for any that
By letting , we have
Note that this inequality is valid for . Therefore, using equation B.5, we have
Rearranging the equation and applying upper bound of completes the proof. ∎
Under Assumption 1 and 3, if we take , , , and , then FedAvg-M converges as
Summing over from to and applying Lemma B.3, we have
Combining this inequality with Lemma A.1, we get
Finally, noticing that implies , we obtain
B.2 FedAvg-M-VR
When each local loss function is further assumed to be sample-wise smooth (i.e., Assumption 2), we can replace the local descent direction in Algorithm 1 with a variance-reduced momentum direction
to further enhance convergence, leading to FedAvg-M-VR, as presented in Algorithm 3. Here, the variable is the last-iterate global model maintained in the server. Same as FedAvg-M, turning off the variance-reduced momentum of FedAvg-M-VR, i.e., setting , recovers FedAvg. FedAvg-M-VR shares the same algorithmic structure and uplink communication workload as FedAvg.
B.2.2 Convergence analysis
If , the following holds for :
Also for , it holds that
By the AM-GM inequality and Assumption 2, we have
The last inequality is derived by and . Similarly, for , we can obtain
If , the following holds:
Note that . Then we have
Therefore for any ,
Here the second inequality is by . The last inequality is derived by unrolling the recursive bound and using . By Lemma A.3, it holds that
Summing up equation B.18 over , using equation B.16 and due to the condition on , we have
If and , then the following holds:
Recall that . Consequently, we have
Using the AM-GM inequality, we can obtain that for any ,
Taking in the inequality above, we have
This inequality holds as well trivially for . Therefore, we have
Here the last inequality is due to the upper bound of . ∎
Under Assumption 2 and 3, if we take with and set , , , and
Summing over from to and applying Lemma B.7, we get
Combining this inequality with Lemma A.1, we get
Finally, noticing that implies , we reach
Appendix C SCAFFOLD with momentum
If , the following holds for :
Note that holds for any . Using Lemma A.4, we have
For , similar to the proof of Lemma B.1, we have
Besides, by AM-GM inequality and Lemma A.3,
The case for is similar. ∎
If and , it holds for all that
Plug this inequality into the above bound completes the proof. ∎
Under the same conditions of Lemma C.2, if and , then we have
Using Young’s inequality repeatedly, we have
Here we apply Lemma A.3 to obtain the second inequality. Combining this with Lemma C.2, we get
where we apply the upper bound of . Therefore, we finish the proof by summing up over from to and rearranging the inequality. ∎
Under Assumption 1 and 3, if we take , with , and set , , , , then SCAFFOLD-M converges as
By Lemma C.1, we can get the following inequality by summing over from to and plugging in Lemma C.2 and Lemma C.3
Here the coefficients in the last inequality is derived by the following bounds:
Combining this inequality with Lemma A.1, we obtain
Finally, noticing that implies and implies , we reach
C.2 SCAFFOLD-M-VR
When each local loss function is further assumed to be sample-wise smooth (i.e., Assumption 2), we can replace the local descent direction in Algorithm 2 with a variance-reduced momentum direction
resulting in SCAFFOLD-M-VR, as presented in Algorithm 4. Here, the variable is the last-iterate global model maintained in the server. Same as SCAFFOLD-M, turning off the variance-reduced momentum of SCAFFOLD-M-VR, i.e., setting , recovers SCAFFOLD.
C.2.2 Convergence analysis
If , then the following holds for :
Using the same derivation as Lemma B.5, we can show that
The case for can be established similarly. ∎
If , , and , then it holds that
Note that , with the same procedures in Lemma B.6, we have
Hence, by applying , we obtain
Also, similar to Lemma C.3, it still holds that
Combine this with the upper bound of ,
where we apply the upper bound of in the last inequality. Iterating the above inequality completes the proof. ∎
Under Assumption 2 and 3, if we take with , and set , , , , then SCAFFOLD-M-VR converges as
By Lemma C.5, sum over from to and plug equation C.2, Lemma C.6 in,
Here the coefficients in the last inequality is derived by the following bounds:
Combining this inequality with Lemma A.1, we get
Finally, noticing that implies and implies , we reach