Tackling System and Statistical Heterogeneity for Federated Learning with Adaptive Client Sampling
Bing Luo, Wenli Xiao, Shiqiang Wang, Jianwei Huang, Leandros Tassiulas
I Introduction
Federated learning (FL) enables many clientsWe use “device” and “client” interchangeably in this paper. to collaboratively train a model under the coordination of a central server while keeping the training data decentralized and private (e.g., ). Compared to traditional distributed machine learning techniques, FL has two unique features (e.g., ), as shown in Fig. 1. First, clients are massively distributed and with diverse and low communication rates (known as system heterogeneity), where stragglers can slow down the physical training time.As suggested , we consider mainstream synchronized FL in this paper due to its composability with other techniques (such as secure aggregation protocols and differential privacy). Second, the training data are distributed in a non-i.i.d. and unbalanced fashion across the clients (known as statistical heterogeneity), which negatively affects the convergence behavior.
Due to limited communication bandwidth and across geographically dispersed devices, FL algorithms (e.g., the de facto FedAvg algorithm in ) usually perform multiple local iterations on a fraction of randomly sampled clients (known as partial participation) and then aggregates their resulting local model updates via the central server periodically . Recent works have provided theoretical convergence analysis that demonstrates the effectiveness of FL with partial participation in various non-i.i.d. settings .
However, these prior works have focused on sampling schemes that select clients uniformly at random or proportional to the clients’ data sizes, which often suffer from slow convergence with respect to wall-clock (physical) timeWe use wall-clock time to distinguish from the number of training rounds. due to high degrees of the system and statistical heterogeneity. This is because the total FL time depends on both the number of training rounds for reaching the target precision and the physical time in each round . Although uniform sampling guarantees that the aggregated model update in each round is unbiased towards that with full client participation, the aggregated model may have a high variance due to data heterogeneity, thus, requiring more training rounds to converge to a target precision. Moreover, considering clients’ heterogeneous communication delay, uniform sampling also suffers from the straggling effect, as the probability of sampling a straggler within the sampled subset in each round can be relatively high,For example, suppose there are 100 clients with only 5 stragglers, then the probability of sampling at least a straggler for uniformly sampling 10 clients in each round is more than 40%. thus yielding a long per-round time.
One effective way of speeding up the convergence with respect to the number of training rounds is to choose clients according to some sampling distribution where “important” clients have high probabilities . For example, recent works adopted importance sampling approaches based on clients’ statistical property . However, their sampling schemes did not account for the heterogeneous physical time in each round, especially under straggling circumstances. Another line of works aims to minimize the learning time via optimizing client selection and scheduling based on their heterogeneous system resources . However, their optimization schemes did not consider how client selection schemes influence the convergence behavior due to data heterogeneity and thus may negatively affect the total learning time.
In a nutshell, the fundamental limitation of existing works is the lack of joint consideration of the impact of the inherent system heterogeneity and statistical heterogeneity on client sampling. In other words, clients with valuable data may have poor communication capabilities, whereas those who communicate fast may have low-quality data. This motivates us to study the following key question.
Key Question: How to design an optimal client sampling scheme that tackles both system and statistical heterogeneity to achieve fast convergence with respect to wall-clock time?
The challenge of this question is threefold: (1) It is difficult to obtain an analytical FL convergence result for arbitrary client sampling probabilities. (2) The total learning time minimization problem can be complex and non-convex due to the straggling effect. (3) The optimal client sampling solution contains unknown parameters from the convergence result, which we can only estimate during the learning process (known as the chicken-and-egg problem).
In light of the above discussion, we state the main results and key contributions of this paper as follows:
Optimal Client Sampling for Heterogeneous FL: We study how to design the optimal client sampling strategy to minimize FL wall-clock time with convergence guarantees. To the best of our knowledge, this is the first work that aims to optimize client sampling probabilities to address both system and statistical heterogeneity.
Convergence Bound for Arbitrary Sampling: Using an adaptive client sampling and model aggregation design, we obtain a new tractable convergence upper bound for FL algorithms with arbitrary client sampling probabilities. This enables us to establish the analytical relationship between the total learning time and client sampling probabilities and formulate a non-convex training time minimization problem.
Optimization Algorithm and Sampling Principle: We propose a low-cost substitute sampling approach to learn the convergence-related unknown parameters and develop an efficient algorithm to approximately solve the non-convex problem with low computational complexity. Our solution characterizes the impact of communication time (system heterogeneity) and data quantity and quality (statistical heterogeneity) on the optimal client sampling design.
Simulation and Prototype Experimentation: We evaluate the performance of our proposed algorithms through both a simulated environment and a hardware prototype. Experimental results demonstrate that for both convex and non-convex learning models, our proposed sampling scheme significantly reduces the convergence time compared to several baseline sampling schemes. For example, with our hardware prototype and the EMNIST dataset, our sampling scheme spends % less time than baseline uniform sampling for reaching the same target loss.
II Related Work
Active client sampling and selection play a crucial role in addressing the statistical and system heterogeneity challenges in cross-device FL. In the existing literature, the research efforts in speeding up the training process mainly focus on two aspects: importance sampling and resource-aware optimization-based approaches.
define The goal of importance sampling is to reduce the variance in traditional optimization algorithms based on stochastic gradient descent (SGD), where SGD draws data samples uniformly at random during the learning process (e.g., ). Recent works have adopted this idea in FL systems to improve communication efficiency via designing client sampling strategy. Specifically, clients with “important” data would have higher probabilities to be sampled in each round. For example, existing works use clients’ local gradient information (e.g., ) or local losses (e.g., ) to measure the importance of clients’ data. However, these schemes did not consider the speed of error convergence with respect to wall-clock time, especially the straggling effect due to heterogeneous transmission delays.
Another line of works aims to minimize wall-clock time via resource-aware optimization-based approaches, such as CPU frequency allocation (e.g., ), and communication bandwidth allocation (e.g., ), straggler-aware client scheduling (e.g., ), parameters control (e.g., ), and task offloading (e.g., ). While these papers provided some novel insights, their optimization approaches did not consider how client sampling affects the total wall-clock time and thus are orthogonal to our work.
Unlike all the above-mentioned works, our work focuses on how to design the optimal client sampling strategy that tackles both system and statistical heterogeneity to minimize the wall-clock time with convergence guarantees. In addition, most existing works on FL are based on computer simulations. In contrast, we implement our algorithm in an actual hardware prototype with resource-constrained devices, which allows us to capture real system operations.
The organization of the rest of the paper is as follows. Section III introduces the system model and problem formulation. Section IV presents our new error-convergence bound with arbitrary client sampling. Section V gives the optimal client sampling algorithm and solution insights. Section VI provides the simulation and prototype experimental results. We conclude this paper in Section VII.
III Preliminaries and System Model
We start by summarizing the basics of FL and its de facto algorithm FedAvg with unbiased client sampling. Then, we introduce the proposed adaptive client sampling for statistical and system heterogeneity based on FedAvg. Finally, we present our formulated optimization problem.
Consider a federated learning system involving a set of clients, coordinated by a central server. Each client has local training data samples (), and the total number of training data across devices is . Further, define as a loss function where indicates how the machine learning model parameter performs on the input data sample . Thus, the local loss function of client can be defined as
Denote as the weight of the -th device such that . Then, by denoting as the global loss function, the goal of FL is to solve the following optimization problem :
The most popular and de facto optimization algorithm to solve (2) is FedAvg . Here, denoting as the index of an FL round, we describe one round (e.g., the -th) of the FedAvg algorithm as follows:
The server uniformly at random samples a subset of clients (i.e., with ) and broadcasts the latest model to the selected clients.
Each sampled client chooses , and runs steps is originally defined as epochs of SGD in . In this paper, we denote as the number of local iterations for theoretical analysis. of local SGD on (1) to compute an updated model . Then, the sampled client lets and send it back to the server.
The server aggregates (with weight ) the clients’ updated model and computes a new global model .
The above process repeats for many rounds until the global loss converges.
Recent works have demonstrated the effectiveness of FedAvg with theoretical convergence guarantees in various settings . However, these works assume that the server samples clients either uniformly at random or proportional to data size, which may slow down the wall-clock time for convergence due to the straggling effect and non-i.i.d. data . Thus, a careful client sampling design should tackle both system and statistical heterogeneity for fast convergence.
III-B System Model of FL with Client Sampling 𝐪𝐪\mathbf{q}
We aim to sample clients according to a probability distribution , where and . Through optimizing , we want to address system and statistical heterogeneity so as to minimize the wall-clock time for convergence. We describe the system model as follows.
Following recent works , we assume that the server establishes the sampled client set by sampling times with replacement from the total clients, where is a multiset in which a client may appear more than once. The aggregation weight of each client is multiplied by the number of times it appears in .
) Statistical Heterogeneity Model
We consider the standard FL setting where the training data are distributed in an unbalanced and non-i.i.d. fashion among clients.
) System Heterogeneity Model
Following the same setup of and , we denote as the round time of client , which includes both local model computation time and global communication time. For simplicity, we assume that remains the same across different rounds for each client , while for different clients and , and can be different. The extension to time-varying is left for future work. Without loss of generality, as illustrated in Fig. 1, we sort all clients in the ascending order , such that
) Total Wall-clock Time Model
We consider the mainstream synchronized FL model where each sampled client performs multiple (e.g., ) steps of local SGD before sending back their model updates to the server (e.g., ). For synchronous FL, the per-round time is limited by the slowest client (known as straggler). Thus, the per-round time of the entire FL process is
Therefore, the total learning time after rounds is
III-C Problem Formulation
In Section IV and Section V, we address these two challenges, respectively, and propose approximate algorithms to find an approximate solution to Problem P1 efficiently.
IV Convergence Bound for Arbitrary Sampling
In this section, we address the first challenge by deriving a new tractable convergence bound for arbitrary client sampling probabilities.
To ensure a tractable convergence analysis, we first state several assumptions on the local objective functions .
L-smooth: For each client , is -smooth, i.e., for all and .
Strongly-convex: For each client , is -strongly convex, i.e., for all and .
IV-B Aggregation with Arbitrary Client Sampling Probabilities
This section shows how to aggregate clients’ model updates under sampling probabilities , such that the aggregated global model is unbiased compared to that with full client participation, which leads to our convergence result.
We first define the virtual weighted aggregated model with full client participation in round as
With this, we can derive the following result.
(Adaptive Client Sampling and Model Aggregation) When clients are sampled with probability and their local updates are aggregated as , we have
Proof Sketch. The basic idea is to take expectation over the aggregated global model of the sampled clients , and with some mathematical derivations, we have (8). ∎
Remark: The key insight of our sampling and aggregation is that since we sample different clients with different probabilities (e.g., for client ), we need to inversely re-weight their updated model in the aggregation step (e.g., for client ), such that the aggregated model is still unbiased towards that with full client participation. We summarize how the server performs client sampling and model aggregation in Algorithm 1, where the main differences compared to the de facto FedAvg in are the Sampling (Line 1) and Aggregation (Line 1) procedures. Notably, Algorithm 1 recovers FedAvg algorithm with uniform sampling when letting , or with weighted sampling when letting in .
IV-C Main Convergence Result for Arbitrary Client Sampling
Based on Lemma 1, we present the main convergence result for arbitrary client sampling in Theorem 1.
(Convergence Upper Bound) Let Assumptions 1 to 4 hold, , and decaying learning rate . For given client sampling probabilities and the corresponding aggregation described in Lemma 1, the optimality gap after rounds satisfies
where and , with and .
Proof Sketch. First, following the similar proof of convergence under full client participation in , we show that
For FL with homogeneous communication time i.e., , for all , the optimal client sampling probabilities for Problem P1 is
Thus, for a target precision , computing the optimal sampling for minimizing the upper bound of is equivalent to solving
The problem in (13) can be easily solved with the Lagrange multiplier method in closed form as shown in (11).
In the following, we show how to leverage the derived convergence bound in Theorem 1 to design the optimal client sampling for the general heterogeneous system of Problem P1.
V Optimal Adaptive Client Sampling Algorithm
We first show that the probability of client being the slowest one (e.g., straggler) amongst the sampled clients in each round is . Since we sample devices according to , taking the expectation of all clients over time gives (15), and for rounds we have (14). ∎
V-B Approximate Optimization Problem for Problem P1
Based on Theorem 2, and by letting the analytical convergence upper bound in (9) satisfy the convergence constraint,Optimization using upper bound as an approximation has also been adopted in . the original Problem P1 can be approximated as
Combining with (9), we can see that Problem P2 is more constrained than Problem P1, i.e., any feasible solution of Problem P2 is also feasible for Problem P1.
We further relax as a continuous variable to theoretically analyze Problem P2. For this relaxed problem, suppose (, ) is the optimal solution, then we must have
This is because if (17) holds with an inequality, we can always find an that satisfies (17) with equality, but the solution (, ) can further reduce the objective function value. Therefore, for the optimal , (17) always holds, and we can obtain from (17) and substitute it into the objective of Problem P2. Then, the objective of Problem P2 is
whichFor ease of analysis, we omit as it is a constant multiplied by the entire objective function. is only associated with client sampling probabilities .
Case 1: For homogeneous (), we have
Case 2: For heterogeneous with , we have
Remark: The objective function of Problem P3 is in a more straightforward form compared to Problem P2. However, to solve for the optimal sampling probabilities , we need to know the value of the parameters in (22), e.g., , , and .We assume that clients’ heterogeneous time and their dataset size can be measured offline.
In the following, we solve Problem P3 as an approximation of the original Problem P1. Our empirical results in Section VI demonstrate that the solution obtained from solving Problem P3 achieves superior total wall-clock time performances compared to baseline client sampling schemes.
V-C Solving Problem P3
Problem P3 is challenging to solve because we can only obtain the unknown parameters , and during the training process of FL. In this subsection, we first show how to estimate these unknown parameters. Then, we develop an efficient algorithm to solve Problem P3. We summarize the overall algorithm in Algorithm 2. Finally, we identify some insightful solution properties.
We first show how to estimate via a substitute sampling scheme.We only need to estimate the value of instead of and each, because we can divide parameter on the objective of Problem P3 without affecting the optimal sapling solution. Then, we show that we can indirectly acquire the knowledge of during the estimation process of .
The basic idea is to utilize the derived convergence upper bound in (9) to approximately solve as a single variable, via performing Algorithm 1 with two baseline sampling schemes: uniform sampling with and weighted sampling with , respectively.
Note that we only let sampling schemes and run until a pre-defined loss is reached (instead of running all the way until they converge to the precision ), because our goal is to find and run with the optimal sampling scheme so that we can achieve the target precision with the minimum wall-clock time.
Specifically, suppose and are the number of rounds for reaching the pre-defined loss for schemes and , respectively. Considering that the training loss decreases quickly at the beginning and slowly when approaching convergence, the values of and are normally very small compared to the required number of rounds for reaching the target loss (with precision ). According to (9), we have
Then, we can obtain from (24) once we know the value of . Notably, we can estimate during the procedure of estimating . The idea is to let the sampled clients send back the norm of their local SGD along with their returned local model updates, and then the server updates with the received norm values. This approach does not add much communication overhead, since we only need to additionally transmit the value of the gradient norm (e.g., only a few bits for quantization) instead of the full gradient information. In addition, instead of retraining the model using the calculated from the initial parameter , we can continue to train the global model after the estimation process, to avoid repeated training and reduce the overall training time.
In practice, due to the sampling variance, we may set several different to obtain an averaged estimation of . The overall estimation process corresponds to Lines 2–2 of Algorithm 2.
We first identify the property of Problem P3 and then show how to compute .
To solve Problem P3, we define a new control variable
where . Then, we rewrite Problem P3 as
For any fixed feasible value of , Problem P4 is convex with , because the objective function is strictly convex and the constraints are linear.
We will solve Problem P4 in two steps. First, for any fixed , we will solve for the optimal in Problem P4, via a convex optimization tool, e.g., CVX . This allows us to write the objective function of Problem P4 as . Then we will further solve the problem by using a linear search method with a fixed step-size over the interval , where we use the optimal and the corresponding in the search domain to approximate the optimal and in Problem P4. This optimization process corresponds to Lines 2–2 of Algorithm 2.
Remark: Our optimization algorithm is efficient in the sense that the linear search domain is independent of the scale of the problem, e.g., number of .
) Property of Optimal Client Sampling
Next we show some interesting properties of the optimal sampling strategy.
Suppose is the optimal solution of Problem P3. For two different clients and , if and , then .
Theorem 4 shows that the optimal client sampling strategy should allocate higher probabilities to those who have smaller and larger product value of , which characterizes the impact and interplay between system heterogeneity and statistical heterogeneity. Although it may be infeasible to derive an analytical relationship regarding the exact impact of and on due to non-convexity of Problem P3, we show by Corollary 2 that we can obtain the closed-form solution of with and when .
When , the global optimal solution of Problem P3 is
If , because is bounded between , we have . Then, the objective of Problem P3 can be written as
Hence, the minimum of (28) is , which is independent of . The equality of (29) holds if and only if (for an arbitrary scalar ). Noting yields (27) and concludes the proof.
Though Corollary 2 is valid only for the special case of , the global optimal sampling solution in (27) characterizes an analytical interplay between the system heterogeneity () and statistical heterogeneity (). Particularly, when , for each , in (27) recovers the optimal sampling solution in Corollary 1 for homogeneous systems.
VI Experimental Evaluation
In this section, we empirically evaluate the performance of our proposed client sampling scheme (Algorithm 2) and compare it with four other benchmarks in each round: 1) full participation, 2) uniform sampling, 3) weighted sampling, and 4) statistical sampling where we sample clients according to Corollary 1. Benchmarks 1–3 are widely adopted for convergence guarantees in . The fourth baseline is an offline variant of the proposed schemes in .The client sampling in is weighted by the norm of the local stochastic gradient in each round, which frequently requires the knowledge of stochastic gradient from all clients to calculate the sampling probabilities.
In the following, we first present the evaluation setup and then show the experimental results.
We conduct experiments both on a networked hardware prototype system and in a simulated environment.The prototype implementation allows us to capture real system operation time, and the simulation system allows us to simulate large-scale FL environments with manipulative parameters. As illustrated in Fig. 2, our prototype system consists of Raspberry Pis serving as clients and a laptop computer acting as the central server. All devices are interconnected via an enterprise-grade Wi-Fi router. We develop a TCP-based socket interface for the communication between the server and clients with bandwidth control. In the simulated system, we simulate virtual devices and a virtual central server.
) Datasets and Models
We evaluate our results on two real datasets and a synthetic dataset. For the real dataset, we adopted the widely used MNIST dataset and EMNIST dataset . For the synthetic dataset, we follow a similar setup to that in , which generates -dimensional random vectors as input data. We adopt both the convex multinomial logistic regression model and the non-convex convolutional neural network (CNN) model with LeNet-5 architecture .
) Implementation
Prototype Setup: We conduct the first experiment on the prototype system using logistic regression and the EMNIST dataset. To generate heterogeneous data partition, similar to , we randomly subsample lower case character samples from the EMNIST dataset and distribute among edge devices in an unbalanced (i.e., different devices have different numbers of data samples, following a power-law distribution) and non-i.i.d. fashion (i.e., each device has a randomly chosen number of classes, ranging from to ).The number of samples and the number of classes are randomly matched, such that clients with more data samples may not have more classes.
Simulation Setup 1: We conduct the second experiment in the simulated system using logistic regression and the Synthetic dataset. To simulate a heterogeneous setting, we use the non-i.i.d. setting. We generate data samples and distribute them among clients in an unbalanced power-law distribution.
Simulation Setup 2: We conduct the third experiment in the simulated system using CNN and the MNIST dataset, where we randomly subsample data samples from MNIST and distribute them among clients in an unbalanced (following the power-law distribution) and non-i.i.d. (i.e., each device has 1–6 classes) fashion.
) Training Parameters
For all experiments, we initialize our model with and use an SGD batch size of . We use an initial learning rate of with a decay rate of , where is the communication round index. We adopt the similar FedAvg settings as in , where we sample % of all clients in each round, i.e., for Prototype Setup and for Simulation Setups, with each client performing local iterations.We also conduct experiments both on Prototype and Simulation Setups with variant and , which show a similar performance as the experiments in this paper, and due to page limitations, we do not illustrate them all.
) Heterogeneous System Parameters
For the Prototype Setup, to enable a heterogeneous communication time, we control clients’ communication bandwidth and generate a uniform distribution seconds, with a mean of seconds and the standard deviation of seconds. For the simulation system, we generate the client transmission delays with an exponential distribution, i.e., seconds, with both mean and standard deviation as second.
VI-B Performance Results
We evaluate the wall-clock time performances of both the global training loss and test accuracy on the aggregated model in each round for all sampling schemes. We average each experiment over independent runs. For a fair comparison, we use the same random seed to compare sampling schemes in a single run and vary random seeds across different runs.
Fig. 3–5 show the results of Prototype Setup, Simulation Setup 1, and Simulation Setup 2, respectively. We summarize the key observations as follows.
As predicted by our theory, Figs. 3(a)–5(a) show that our proposed sampling scheme achieves the same target loss with significantly less time, compared to the baseline sampling schemes. Specifically, for Prototype Setup in Fig. 3(a), our proposed sampling scheme spends around % less time than full sampling and uniform sampling and around % less time than weighted sampling and statistical sampling for reaching the same target loss. Fig. 5(a) highlights the fact that our proposed sampling works well with the non-convex CNN model, under which the naive uniform sampling cannot reach the target loss within seconds, indicating the importance of a careful client sampling design. Table I summarizes the superior performances of our proposed sampling scheme in wall-clock time for reaching target loss in all three setups.
) Accuracy with Wall-clock Time
As shown in Fig. 3(b)–5(b), our proposed sampling scheme achieves the target test accuracyIn Fig. 3(b), Fig. 4(b), and Fig. 5(b), the target test accuracy corresponds to the test accuracy result when our proposed scheme reaches the target loss. much faster than the other benchmarks. Notably, for Simulation Setup 1 with the target test accuracy of % in Fig. 4(b), our proposed sampling scheme takes around % less time than full sampling and around % less time than the other sampling schemes. We can also observe the superior test accuracy performance of our proposed sampling schemes in Prototype Setup and non-convex Simulation Setup in Fig. 3(b) and Fig. 5(b), respectively.
) Loss with Number of Rounds
Fig. 3(c)–5(c) show that our proposed sampling scheme requires more training rounds for reaching the target loss compared to baseline statistical sampling and full participation schemes. This observation is expected since our proposed sampling scheme aims to minimize the wall-clock time instead of the number of rounds. Nevertheless, we notice that statistical sampling performs better than the other sampling schemes, which verifies Corollary 1 since the performance of loss with respect to the number of rounds is equivalent to that with respect to wall-clock time for homogeneous systems.
VII Conclusion and Future Work
In this work, we studied the optimal client sampling strategy that addresses the system and statistical heterogeneity in FL to minimize the wall-clock convergence time. We obtained a new tractable convergence bound for FL algorithms with arbitrary client sampling probabilities. Based on the bound, we formulated a non-convex wall-clock time minimization problem. We developed an efficient algorithm to learn the unknown parameters in the convergence bound and designed a low-complexity algorithm to approximately solve the non-convex problem. Our solution characterizes the interplay between clients’ communication delays (system heterogeneity) and data importance (statistical heterogeneity), and their impact on the optimal client sampling design. Experimental results validated the superiority of our proposed scheme compared to several baselines in speeding up wall-clock convergence time.