On the Outsized Importance of Learning Rates in Local Update Methods
Zachary Charles, Jakub Konečný
Introduction
Historically, machine learning was analyzed from a “centralized” perspective, in which a model is trained on a single central source of data. In recent years, there has been a shift away from centralized machine learning, due in part to the increase of user data and the increasing awareness of the risks to privacy that can accompany centralized data collection.
Federated learning (FL) (Kairouz et al., 2019) is a distributed framework for learning models without directly sharing user data. In this framework, heterogeneous clients all use their own data to perform local training. In the popular FedAvg algorithm (McMahan et al., 2017), the client models are then averaged at a central server, broadcast to a (possibly different) sample of clients, and the process is repeated. The core tenet is that instead of having clients share data, we instead share the results of local updates the clients perform on their own datasets using an optimization algorithm.
While there has been growing interest in FL in research communities (see (Kairouz et al., 2019) and (Li et al., 2019) for surveys of many recent works and open problems), this general paradigm of performing local updates on heterogeneous datasets has a storied history in machine learning. In particular, much of the work on meta-learning has focused on trying to learn models that perform well (or can quickly learn how to perform well) on a large number of heterogeneous tasks. This similarity with FL is even more clear in work on model-agnostic meta-learning (MAML) (Finn et al., 2017), in which local client gradient updates are used to learn a global model. Connections between these two areas were noted by Jiang et al. (2019) and have since been explored in many other works (Khodak et al., 2019; Fallah et al., 2020).
While there is a wide variety of theoretical and empirical analyses of the aforementioned methods, it is generally difficult to understand their behavior in heterogeneous settings. There is enough evidence that these methods are useful in practice in complex scenarios (Hard et al., 2018; Yang et al., 2018; Hard et al., 2020), yet on a theoretical level, many works derive results comparable to, or worse than, that of mini-batch SGD in heterogeneous or even homogeneous settings; See (Kairouz et al., 2019) for a discussion of homogeneity and heterogeneity, and see (Woodworth et al., 2020) for a detailed discussion of comparisons to mini-batch SGD. Unfortunately, these results shed little light onto how methods such as FedAvg improve (or degrade) convergence.
In this work, we analyze a generalized local update paradigm that encompasses many FL and MAML methods, as well as other popular optimization methods such as mini-batch SGD. In order to better understand the structure of these methods in heterogeneous settings without an abundance of assumptions, we focus on the special case of quadratic loss functions. We are generally concerned with understanding the following questions that bridge both theory and practice.
How do local update methods improve or hinder convergence?
Why, despite a relative paucity of theoretical evidence, do these methods often perform better in practice than theoretically established methods such as mini-batch SGD?
What obstacles are there to the performance of local update methods, and how do we mitigate these issues?
As a partial answer to these questions, we highlight the main findings of our work.
We show that in the quadratic case, local update methods are equivalent to the stochastic gradient method on a surrogate loss function which we exactly characterize. Thus, we can view local update methods that use multiple heterogeneous datasets as instead performing SGD on a single “central” loss function.
We show that methods such as FedAvg and many incarnations of MAML implicitly regularize the condition number of this surrogate loss function, allowing for improved convergence of the surrogate loss. On the other hand, we show that this condition number reduction comes at the cost of increasing the discrepancy between minimizers of the surrogate and the true loss function. Notably, this trade-off is controlled by fundamental algorithmic choices, especially the choice of learning rate.
We give explicit convergence rates for FedAvg that exhibit the trade-off between the condition number and the discrepancy between the surrogate and true loss functions above. Our results are similar in scope to work by Woodworth et al. (2020) (showing that local SGD can outperform mini-batch SGD), but work under heterogeneous data settings.
We use our theoretical insights to design practical improvements to federated learning methods. First, we show that decoupling client and server learning rates has significant implications for improving convergence to better models. We show that despite the non-optimality of critical points of FedAvg, combining this learning rate decoupling with proper tuning can result in near-optimal performance in settings with limited communication. Finally, we detail a simple, practical method for automatic learning rate decay in federated learning that helps reduce the burden of learning rate tuning. We show empirically that this method improves the convergence of FedAvg, without requiring manually crafted learning rate schedules, across a suite of realistic and challenging non-convex tasks.
Federated learning is a distributed machine learning paradigm in which training is done locally on clients, without any centralized data aggregation. Federated learning has enabled privacy-aware learning in a variety of applications (Hard et al., 2018; Chen et al., 2019; Brisimi et al., 2018; Samarakoon et al., 2018; Hard et al., 2020), and has seen a large volume of work on the intersection of federated learning with topics including differential privacy (McMahan et al., 2018; Augenstein et al., 2020), fairness (Mohri et al., 2019; Li et al., 2020b), robustness (Ghosh et al., 2019; Bagdasaryan et al., 2018; Sun et al., 2019), and communication-efficiency (Konečný et al., 2016; Sattler et al., 2019; Basu et al., 2019; Reisizadeh et al., 2020). For a more detailed discussion of federated learning, we defer to surveys by Kairouz et al. (2019) and Li et al. (2019).
In meta-learning (aka learning to learn), the objective is to use a collection of tasks to learn how to learn a new task efficiently (Vanschoren, 2019). A particularly influential recent approach is model-agnostic meta-learning (MAML) proposed by Finn et al. (2017). The core idea has inspired a number of extensions (Antoniou et al., 2019; Nichol et al., 2018; Rusu et al., 2019; Grant et al., 2018; Rajeswaran et al., 2019; Raghu et al., 2020), which broadly use a two-level optimization structure to perform meta-learning. Convergence properties of some of these optimization algorithms were recently studied by Fallah et al. (2019), who also highlight differences in convergence of MAML and first-order approximations to MAML.
One of the most common approaches to optimization in the setting of federated learning is the FedAvg method (McMahan et al., 2017). While designed for heterogeneous sources of data, the study of FedAvg has roots in that of Local SGD (Zinkevich et al., 2010; Stich, 2019; Wang and Joshi, 2018; Stich and Karimireddy, 2019; Yu et al., 2019; Khaled et al., 2020), a communication-efficient optimization method for homogeneous clients. As interest in federated learning has grown, so too has the number of proposed federated optimization methods. These can often be seen as variants of FedAvg, that incorporate techniques such as momentum (Hsu et al., 2019), adaptive optimization (Reddi et al., 2020; Xie et al., 2019), proximal updates (Li et al., 2020a; Pathak and Wainwright, 2020) and control variates (Karimireddy et al., 2019). We again defer to Kairouz et al. (2019) and Li et al. (2019) for more detailed references.
While we defer to Kairouz et al. (2019, Section 3.2) for a complete discussion of federated optimization, we discuss a few important connections. First, while there has been huge progress in theoretical understandings of FedAvg, existing works generally have not been able to show that these methods consistently improve upon mini-batch SGD (Woodworth et al., 2020). Even theoretically and empirically successful techniques such as SCAFFOLD (Karimireddy et al., 2019) have only been shown to converge faster than mini-batch SGD on quadratic objectives.
This failure of convergence was noted by Li et al. (2020c), who showed that without learning rate decay, FedAvg is not guaranteed to converge. Later, Karimireddy et al. (2019) and Woodworth et al. (2020) showed that there are settings where FedAvg converges provably slower than mini-batch SGD. Similarly, Malinovsky et al. (2020) and Pathak and Wainwright (2020) showed that in heterogeneous settings, FedAvg can converge to sub-optimal points, even in non-stochastic, strongly convex settings. Pathak and Wainwright (2020) further give a proximal version of federated gradient descent that converges to the empirical risk minimizer in convex settings.
Our work is most closely related to that of Malinovsky et al. (2020), Pathak and Wainwright (2020). We also evince the non-convergence of FedAvg. However, we extend the analysis to stochastic settings, and to a more general class of algorithms that encompasses many meta-learning algorithms. As such, our work is also closely related to that of Fallah et al. (2019), who demonstrated differences (and non-convergence issues) of various MAML algorithms. Our work takes this a step further, where we give a unified view of both MAML and federated learning methods, and give a broader characterization of the sub-optimal convergence of these methods in the case of quadratic losses. Our work is also novel in its focus on the interplay between convergence, suboptimality, and algorithmic choices, especially learning rates.
Preliminaries
We let denote the gradient of the function with respect to . For , we define the client loss function and the overall loss function as follows:
One common objective in our setup is to minimize , though this is often not the direct goal of MAML methods. Note that the joint distribution over implicitly defines a (marginal) distribution over , recovering standard risk minimization frameworks. This framework also encompasses distributed risk minimization in which is a uniform distribution over a finite set of nodes and is the uniform distribution over the (finite) dataset stored at node . However, we take a more general approach and do not assume or to be finite throughout. We also focus on the heterogeneous setting, where the client distributions are not all identical, as opposed to the homogeneous setting, where all are identical.
As FL has matured, it has become more evident that there are two varieties, with distinct system-imposed constraints, recently termed by Kairouz et al. (2019) as cross-device federated learning and cross-silo federated learning.A different categorization, vertical and horizontal, was proposed by Yang et al. (2019), which is based on modelling constraints, rather than on system constraints. The setup in this work applies primarily to horizontal FL, though we expect that much of our framework carries over to the vertical setting. The primary distinction between these two frameworks that is relevant to our work is that in cross-silo FL, there are relatively few participating clients. Moreover, these clients are typically reliable and almost always available. By contrast, in cross-device FL there are potentially very large numbers of clients, only a small fraction of which are available at any given point in time. Furthermore, the clients cannot be addressed directly or re-identified if participating multiple times. For a more detailed summary, see (Kairouz et al., 2019, Table 1).
In cross-device FL, a client sampled from corresponds to a single device, and corresponds to the data available on that device. In many practical cross-device FL systems (see Bonawitz et al. (2019); Hard et al. (2018)), the server does not control the selection of clients from the global population . Instead, participation is initiated by the clients, based on pre-defined eligibility criteria, such as whether the device is charging and on unmetered wifi. Thus, the client distribution can be considered as fixed, with only minor possibilities for it to be shaped by the server (e.g. whether to enforce sampling without replacement).
On the other hand, in many examples of cross-silo FL, participating clients correspond to various medical or financial organizations, or different geographical regions of the same organization (Wen et al., 2019; Yang et al., 2019). The participating clients are typically fixed in advance, and often all of them participate in every communication round. Thus, while cross-silo FL may be accurately described by a finite-sum optimization problem, this framework is less useful for cross-device FL.
To see this, consider the task of next word prediction on mobile devices. The FL training described by Hard et al. (2018) runs for communication rounds, with up to clients participating in each round. That is at most million distinct clients, a small fraction of the total number of possible clientsAs of May 26, 2020, the Google Play Store reports “1,000,000,000+ installs” for the GBoard application.. This also implies that it is nearly impossible to compute exact values of the loss . Instead, evaluation of a model’s quality is done using the same mechanism as the training – by using a subset of the clients eligible at a given time – which has significant implications for algorithm design (as we discuss in Sections 8 and 9). These issues are exacerbated by heterogeneity; Under extreme heterogeneity, finite-sum modelling approaches may lead to theory that does not accurately represent practical FL systems. Thus, our modelling assumptions are designed to encompass both cross-silo and cross-device setting.
Our setup is also relevant to that of model-agnostic meta-learning (MAML), first proposed by Finn et al. (2017). In MAML, the main objective is to find a gradient-based mechanism, which given a task sampled from , adapts to have good performance on the distribution . Unlike cross-device FL, where we generally cannot quantify directly because of data restrictions, the distribution is the primary object of interest in MAML. However, much like cross-device FL this distribution is generally not known a priori, but instead is problem-dependent.
1 LocalUpdate algorithms
Throughout our work, we will omit the trailing zeros in any with finite support. The server averages the available updates, and, treating this average as a stochastic gradient of the loss function , performs a gradient step with a server (outer) learning rate . Algorithms 1 and 2 give pseudo-code for LocalUpdate.
This method recovers some well-known algorithms for specific choice of , and . For convenience of notation, we define
so in particular and , and similarly
Many existing training algorithms can be expressed as special cases of LocalUpdate. We give a non-exhaustive list below.
The simplest setting is mini-batch SGD. This can be recovered in multiple ways. For example, suppose each client corresponds to a single example . Then, LocalUpdate with is equivalent to mini-batch SGD with batch size and learning rate .
Alternatively, if there is only a single client (), then LocalUpdate with becomes mini-batch SGD with batch size and learning rate . As expected, the choice of has no impact in either instance of mini-batch SGD.
More generally, setting recovers distributed mini-batch SGD, with total batch size and learning rate . Again, has no impact on the global model.
When there is a single client and , then LocalUpdate recovers the Lookahead optimizer (Zhang et al., 2019) with “fast weights”.
In the homogeneous setting, if , and , then LocalUpdate is equivalent to Local SGD with local steps. For further details, see Appendix A.
In the heterogeneous setting, if we set and , then LocalUpdate is equivalent to FedAvg with local steps (see Appendix A for details). When and are not necessarily equal, we actually recover Reptile (Nichol et al., 2018), as well as the Generalized FedAvg algorithm in (Reddi et al., 2020). For convenience of notation, we will refer to this algorithm as FedAvg/Reptile throughout. This equivalence between FedAvg and Reptile was first noted by Jiang et al. (2019). In fact, we show in Section 8 that this decoupling of client and server learning rates is critical to understanding and improving the convergence of FedAvg.
When , we recover the first-order MAML (FOMAML) algorithm of Finn et al. (2017). A similar functional relation between FOMAML and Reptile was previously described by Nichol et al. (2018).
In the MAML algorithm, Finn et al. (2017) use local update steps for each “task” (in our vocabulary, client). We will refer to this as -MAML throughout. As we show in Section 3.1, when the underlying loss functions are quadratic and the clients perform gradient descent updates, -MAML is recovered by setting . This gives a previously unknown connection between FL and MAML algorithms. As we discuss in Section 3.1, this does not hold when the clients use SGD due to potential biases in estimating Hessian-gradient products via stochastic gradients.
As written, both clients and server use SGD as their optimizer in LocalUpdate. However, one could use techniques such as momentum or adaptive learning rates on either the server (as explored by Reddi et al. (2020)) or the client (as explored by Xie et al. (2019)). While our results can be extended to these settings, we leave this to future work. Our goal is not to derive convergence results for as broad a class of algorithms as possible. Rather, we wish to understand how the choice of and impact the dynamics of optimization, especially in heterogeneous settings.
We note that in Algorithm 2, each clients performs a designated number of steps of mini-batch SGD, with samples taken from some underlying client distribution . When is the uniform distribution over some finite set , we could instead write Algorithm 2 in terms of performing some number of epochs of mini-batch SGD over , as is done in (McMahan et al., 2017) and many other works on federated learning. In this case, the batch size dictates the number of client gradient steps (as the client roughly take steps). Thus, in such settings, the choice of has an analogous impact as the choice of the number of local steps in Algorithm 2. For simplicity of analysis, we will analyze the latter throughout, but our results can be easily extended to the former.
2 Outline
The rest of this paper is organized as follows. In Section 3, we show that a round of LocalUpdate method is equivalent to performing a single (stochastic) gradient step with respect to a surrogate objective, which we exactly characterize.
In Section 4, we use simple examples to show that the surrogate loss and the original loss can vary substantially. Moreover, we show how choices of and affect the discrepancy between the two losses. In particular, we highlight how the choice of is crucial to the performance of LocalUpdate. In Section 5, we analyze spectral properties of the surrogate loss, and show that LocalUpdate can be viewed as implicit regularization on the condition number, where the amount of regularization is controlled by and .
The next sections present to the best of our knowledge a novel proof technique, characterizing the convergence of FedAvg/Reptile in heterogeneous settings.While we focus on FedAvg and Reptile, we note that a similar analysis can be performed for any of the special cases listed above, using a similar proof strategy. In Section 6, we bound the distance between the minimizers of the surrogate and the true loss function in terms of the client learning rate . We use these results in Section 7 to derive convergence rates for FedAvg/Reptile that highlight how the choice of client learning rate gives rise to a trade-off between local and global optimization. In particular, we show that learning rate decay is both sufficient and necessary for convergence to the true risk minimizer.
While our theoretical results are valid only for quadratic loss functions, in Section 8 we show empirically that our conclusions carry over to more general settings, including non-convex objectives. Our empirical results highlight the importance of learning rate tuning in federated learning. In Section 9, we combine our theoretical insights with important systems-level constraints to design a method for automatic learning rate decay methods for local update methods. In particular, we present a simple, easy to implement method for automatic learning rate decay, and show its efficacy in improving accuracy and reducing the need for client learning rate tuning.
LocalUpdate as SGD
When and , the dynamics of LocalUpdate may be very different than those of mini-batch SGD. We will show that for quadratic functions, these dynamics are related but distinct. In particular, we will show that any local update method on a quadratic function can be viewed as SGD on some appropriately defined surrogate loss function. Moreover, the discrepancy between the true loss function and the surrogate loss function is dictated by the choice of client learning rate and .
We assume throughout that is finite and invertible. We also define
Again, we assume this is finite. We then have the following lemma.
For all , there is some constant such that
In the sequel, we will omit the constant term , and let
as this does not change the gradients of the loss . Since each is symmetric and positive definite, so is . We define the following:
We will assume that these expectations exist and are finite throughout. We will also utilize the following mild assumptions at different times.
There are such that for all ,
There are finite and such that
Assumption 3 assumes that the matrices and optimal points for each loss function have bounded variance. We do not assume that the gradients computed by the clients have bounded norm. Intuitively, as , local update methods should provide more benefit, as the clients are taking more steps towards a shared optimum. While in the case of homogeneous data distributions (i.e. are the same for all ), these two conditions are not equivalent. There are heterogeneous data distributions which still yield . Also, note that is in general not the minimizer of the objective .
Fix , and consider Algorithm 2. We initialize , and then at each iteration we sample a set uniformly at random (with replacement) from , then update via
We first prove a basic recurrence relation concerning the local gradients for task .
For , suppose that is invertible and as in Algorithm 2, for all ,
Using Lemma 2, we will show that Algorithm 2 can be viewed as performing SGD on a surrogate loss. This surrogate loss will be parameterized by the inputs and to Algorithm 2. To define the surrogate loss, we first define, for each client , a distortion matrix as follows:
We can then define, for each , the client’s surrogate loss function:
The overall surrogate loss function is then given by
Informally, the matrix can be viewed as causing a distortion to the matrix . When , one can see that , in which case there is no distortion. For other , may significantly distort , and can amplify heterogeneity of the . Using Lemma 2, we derive the following property of the output of the Algorithm 2.
This proves the first equality. The second follows from noting that .
We note that a version of Theorem was first shown for the case by Fallah et al. (2020), and was used to compare the behavior of FOMAML and MAML. We will take this comparison a step further, by showing in Section 3.1 that in the non-stochastic client setting, MAML can also be viewed as performing SGD on a similarly-defined surrogate loss.
1 MAML
As previously discussed, in the setting above, one can actually view MAML as a special case of LocalUpdate. In this section we elaborate on the claim, using a similar presentation of MAML as in (Nichol et al., 2018). MAML with local steps (which we refer to as -MAML) can be viewed as a simple modification of LocalUpdate. Algorithm 1 proceeds in the same manner. In Algorithm 2, each client still executes mini-batch SGD steps. However, what each client sends to the server differs from Algorithm 1.
For simplicity, we define as the function that runs steps of mini-batch SGD, starting from , for some fixed mini-batches of size drawn independently from . For convenience, we let . We then define
Note that these are implicitly functions of the mini-batches sampled from . The output of client (as a function of its initial model ) is a stochastic estimate of , so that
The remainder of the MAML algorithm proceeds in the same way as LocalUpdate. Namely, the server averages the client outputs, and uses this as a gradient estimate with learning rate . That is,
We now show that when the clients use gradient descent to perform their local update, -MAML is in expectation equivalent to performing LocalUpdate with .
If is the function that runs steps of gradient descent, starting from , on the client dataset , then
It is fruitful to reflect on what this means. Informally, this result shows that for quadratic functions, the gradient of the loss after steps of gradient descent, taken with respect to the initial point , is in expectation the gradient of the loss function after additional SGD steps. In particular, given as in (14), we have
We note that this result relied on the clients using gradient descent to compute . For computational efficiency, this is often instead done using mini-batch SGD. However, computing then involves computing unbiased estimates of the Hessian and gradient using the same batches of data. By the chain rule, estimating involves multiplying these Hessian and gradient estimates. However, the product of these unbiased estimators need not be unbiased since they were computed with respect to the same batch of data. Thus, this correspondence between -MAML and LocalUpdate may break down in computationally-efficient (but biased) MAML implementations. For more detailed discussion on this bias, see (Fallah et al., 2019). We also note that Fallah et al. (2019) analyze the convergence properties of MAML and FOMAML, and independently observe that MAML and FOMAML need not share stationary points for quadratic objectives.
Local update methods tend towards different global minima
In order to enhance our understanding of the surrogate loss function, we first give both analytic and empirical examples.
Let have support , with each option equally likely. Suppose that have support only on the points and respectively, and suppose , so that . We see that is minimized at . On the other hand, let . For , we can compute the -th distortion matrix by
Note that this is positive definite for as long as . By (10),
For , this is a positive definite quadratic function, with minimum given by
Therefore, even if we run LocalUpdate until convergence, it would not converge to the true risk minimizer. This holds even though there are only two clients, each with a single data point. In other words, some form of learning rate decay is necessary for convergence to the risk minimizer. The necessity of learning rate decay was first shown by Li et al. (2020c), and later shown in (Malinovsky et al., 2020) and (Pathak and Wainwright, 2020). We take this analysis further, by showing how this sub-optimal behavior is explicitly governed by algorithmic choices, especially learning rate and the number of local steps taken.
The client learning rate should be set sufficiently small so that all client loss functions are well-behaved, even if the overall loss function is well-behaved.
A similar analysis shows that if we instead fix , the distance between the surrogate risk and true risk minimizers depends on . Let . By (8),
This is a positive definite quadratic with minima given by
When , this gap is 0, while as , the distance increases monotonically to . In fact, as , the minimizer of the surrogate loss function converges to the expected value of the client loss minimizers. In Lemma 15 we prove an even stronger statement, and show that it holds for all positive definite quadratic loss functions.
Next, we give an empirical generalization of the above example for further illustration. Let for (ie. and ). We let have support and density function . For each , we let the client distribution be supported on a single point , so that . Again, deterministic will still be sufficient observe discrepancies between the true and surrogate loss functions.
If or , we converge to . As or increases, we converge to a point further from . As we converge to the average minimizer. We also see that decreasing the local stepsize increases the variance. This is to be expected: When , LocalUpdate with reduces to mini-batch SGD with batch size , but the gradients in the batch are summed rather than averaged. For , the magnitude of the gradients being summed decreases as the client converges to its minimizer. If we set to be larger than , we see an even greater gap between the surrogate minimizer and the true minimizer (due to the presence of negative-definite clients).
In Section 6, we derive general bounds on the distance between surrogate and true minimizers. To do so, we will use spectral properties of the matrix , which we derive in the next section.
Surrogate loss properties
Let and suppose that . Then
is symmetric and positive definite.
For each eigenvalue of , has an eigenvalue
For convenience, we note some special cases of Lemma 6 for .
Let and suppose that . Then
If Assumption 2 holds and , then
Let and suppose that . Then
If Assumption 2 holds and , then
As the eigenvalue bounds above suggest, actually has a relatively simple form, as we show in the next lemma.
Note that when , (8) implies .
Let .
For each eigenvector and eigenvalue pair of , is an eigenvector of with eigenvalue
If , is symmetric and positive definite with eigenvalues satisfying
The bounds in Lemma 10 can be refined for specific . We first consider FedAvg/Reptile, when . By Lemma 9, when , we have
We will therefore be able to compute the eigenvalues of in terms of the function
In fact, is actually continuous at , with its value being given by . One way to see this is by noting that by basic properties of geometric sums, for ,
We can now give strong bounds on the spectrum of . We get the following:
Let .
For each eigenvector, eigenvalue pair of , is an eigenvector of with eigenvalue .
If , the maximum and minimum eigenvalues of are given by
If Assumption 2 holds and , then
We can also tighten the bounds in Lemma 10 for the MAML-style algorithms, where , as long as the learning rate is set appropriately.
Let .
For each eigenvector, eigenvalue pair of , is an eigenvector of with eigenvalue .
If , the maximum and minimum eigenvalues of are given by
If Assumption 2 holds and , then
Lemmas 11 and 12 imply the following results regarding the condition number of the surrogate loss.
It is well known that the condition number measures how quickly methods such as gradient descent can find a minimizer of a strongly convex function (see Chapter 3 of (Bubeck, 2017) for reference). By performing more local computations on the clients, and with larger learning rate, we actually reduce the condition number of the surrogate loss. We see that intuitive notions about methods such as FedAvg (e.g. that more local computation improves convergence) can be made formal by analyzing properties of the surrogate loss. Thus, we have the following important takeaway:
Methods such as MAML, FedAvg, and Reptile perform implicit regularization on the condition number of the surrogate loss function they are actually optimizing.
When or , we see that LocalUpdate may be able to optimize the surrogate loss more quickly (due to the condition number reduction). However, as shown in Section 4, the surrogate loss may differ drastically from the true loss. In the next section, we use the spectral properties of the surrogate loss derived above to quantify the distance between the minimizers of these two functions.
Bounding the distance between global minima
We are first interested in how far apart the two minimizers can possibly be. We first consider the asymptotic affect of when setting . In fact, varying can only change the distance between the two by a fixed amount, as shown in the following.
Suppose . Then for all ,
Intuitively, we see that as long as the client learning rate is not too high, the worst possible surrogate loss is the one defined by the average distance to the client optimizers. Intuitively, as , LocalUpdate will take steps oriented more and more towards the average of the client minimizers (one-shot averaging), which is reflected in the experiment in Figure 1.
We now wish to understand the non-asymptotic regime, especially the distance between the surrogate risk minimizer and the true risk minimizer, as this will help inform us how to set in LocalUpdate. We have the following result.
where is as in (17). We will see in the following theorem that the distance between minimizers is controlled by and .
Informally, the term measures the discrepancy between the surrogate loss function and the true loss function; as , .
In order to get better control on Theorem 17 for , we will show that when is sufficiently small, is close to .
Let and suppose
Plugging these into Theorem 17, we get the following.
We can also use a similar analysis to bound the distance between for different values of . We will focus on the FedAvg/Reptile case. We will show that this distance depends on the discrepancy between the eigenvalues of the matrices .
Let . Then
The presence of the terms makes the dependence on somewhat opaque. In fact, we have the following simpler (though looser) bound.
Let . Then
One particularly useful consequence is that if the client learning rate satisfies for some constant , then
We will use this later to show that by decaying the client learning rate in this manner, successive model updates in LocalUpdate will be closely aligned.
Convergence of FedAvg/Reptile
Our goal in this section is two-fold. First, we wish to understand the trade-offs incurred by performing local computation instead of mini-batch SGD. Second, we wish to show that by Theorem 3, we can analyze federated learning and meta-learning algorithms using classical optimization techniques. There are a large number of important works on federated optimization, that consider more general cases than ours. Unfortunately, the proof techniques behind many of these are relatively opaque, and require careful accounting of the bias incurred by performing local computation. This sometimes leads to either proof errors, or else omitted critical assumptions (Woodworth et al., 2020, Appendix A). Both of these can hinder understanding or make comparisons between convergence rates difficult. By contrast, while limited to a much narrower range of loss functions, our analysis uses essentially standard convex optimization analyses (such as by Rakhlin et al. (2012) and Bottou et al. (2018)), combined with the results from Section 6. We also emphasize that while we focus on FedAvg/Reptile, our results can be easily extended to more general instances of LocalUpdate.
We first wish to understand the setting where the client learning rate is fixed. Fix a client step-size and , and for notational convenience, define
Fix for all in LocalUpdate. Then, at each iteration the server starts at a point which it broadcasts to some number of clients. The clients compute local updates via (Algorithm 2), and send these values to the server. The server then computes the average of the and updates its model via .
Given a starting point , we let denote the random vector computed by averaging vectors of the form where . Thus, in Algorithm 1 with a constant client learning rate , . Recall that by Theorem 3,
Throughout this section, we will assume Assumptions 2 and 3, as well as the following “bounded variance” condition.
We can then translate this into a bound on the variance of .
Suppose Assumption 4 holds. Then for all ,
We will also use the following bound on the strong convexity parameter of our surrogate loss functions.
Note that these results follow directly from Lemma 11 and the fact that the strong convexity and smoothness parameters are governed by the maximum and minimum eigenvalues of .
Using techniques similar to those in (Rakhlin et al., 2012), we arrive at the following descent lemma.
Suppose that Assumptions 2 and 4 hold, and that we have step sizes satisfying
Note that in Lemma 25, we assume a slightly stronger condition than is often assumed in optimization literature, namely that where is the condition number. Typically, works on optimization would only require to be at most the inverse of the Lipschitz constant. While it is an open question as to whether this condition is necessary, there are a few relevant factors. First, we note that we can relax (22) to if we strengthen Assumption 4 to a bounded gradient assumption instead of a bounded variance assumption. Second, when is moderately large with respect to , Lemma 13 implies that , so (22) gives a similar condition to assuming . Finally, we note that a similar bound on the learning rate was used by Reisizadeh et al. (2020) in conjunction with a bounded variance assumption. While we conjecture that this condition can be relaxed, we leave this for future work.
Applying Lemma 25 repeatedly, we derive at the following.
Suppose that . Then the outputs of satisfy
Many prior results for fixed client learning rate provide a bound of the same general form as Theorem 26 (ie. a sum of a decaying term and a constant error term), but bound the distance from , rather than from . This makes the error term’s significance more opaque. While non-federated optimization results often have constant error terms due to stochasticity, in federated convergence results (eg. (Khaled et al., 2020, Theorem 5)), the constant term often does not disappear in deterministic setting (). To the reader, it may not be immediately clear why this is the case. This could be due to actual convergence properties, or due to the analysis not being tight. By contrast, our result shows that this error term is an inherent property of the algorithm, as in general, .
As is the case in general stochastic optimization, a constant learning rate is only sufficient to arrive in a neighborhood of the critical point of the underlying loss. However, this critical point is not the true risk minimizer . As in Section 4, we see that in heterogeneous settings, client learning rate decay is necessary for convergence to the true risk minimizer.
To make the suboptimality gap tend towards zero, we must decay the server learning rate over time, as in the following theorem.
Then the outputs of satisfy
We can now derive a convergence rate towards the true risk minimizer, .
Fix satisfying . Suppose
and that is as in Theorem 27. Then the outputs of satisfy
While notationally complex, this result has a few important facets. Define
As , , LocalUpdate becomes roughly equivalent to mini-batch SGD with batches of size on clients. In fact, is the convergence term we would derive from performing mini-batch SGD with batches of size on clients per round. In particular, we get a variance reduction of . This variance reduction is similar in nature to work by Woodworth et al. (2020) in the homogeneous setting (), which shows an analogous improvement in convergence rates. Note that gets larger as gets larger. Thus, the variance reduction is reduced in heterogeneous settings.
The last term measures the discrepancy between the surrogate loss function and the true loss function. If , for instance, when the data is completely homogeneous, this discrepancy is 0. We then recover similar results to that of Woodworth et al. (2020), which shows a improvement in convergence for Local SGD with steps, but in the heterogeneous setting. If , we can still remove the effect of heterogeneity by setting (which means setting ). In this case, FedAvg reduces to mini-batch SGD with batches of size on each client.
2 Decaying client and server learning rates
In this section, we will show that by decaying the client and server learning rates appropriately, we can derive a bound on that does not require setting the learning rates in terms of the desired optimality gap .
Throughout this section, we again focus on the FedAvg/Reptile setting. At every iteration , we will use client learning rate and server learning rate that decay at a rate. Let denote the condition number of . We have the following theorem.
Suppose we run LocalUpdate with as above and to produce iterates . Then for all ,
No attempt was made to optimize constants. Rather, the point was to show that we can derive bounds on the distance to the true risk minimizer that hold for all and that help illustrate the implicit trade-offs in LocalUpdate. In particular, we do not require that , , or be defined in terms of the desired suboptimality gap . Rather, all that was required was learning rate decay on the order of at both the server and the client.
Recall that by Theorem 26, some form of learning rate decay is necessary for convergence to the true risk minimizer. We therefore see that in some settings, learning rate decay is sufficient for LocalUpdate to converge to the true risk minimizer. While this was previously shown in the finite-sum setting for strongly convex functions by Li et al. (2020c), this result required bounded gradients, and did not illustrate the trade-offs in using FedAvg/Reptile over mini-batch SGD.
Similar to Theorem 1 of (Woodworth et al., 2020), we see that performing LocalUpdate incurs a kind of variance reduction of when compared to performing vanilla SGD over a shuffled version of the entire dataset. We also see that using LocalUpdate with does incur a potential benefit: Rather than having the convergence depend on the suboptimality gap (as is the case for mini-batch SGD), it depends on . This may be much smaller depending on the initialization. For example, recall that by Lemma 15, if , then as , tends to the “one-shot average” of the client minimizers, which typically requires many fewer communication rounds to estimate than the true risk minimizer. However, this reduction only benefits the convergence up to a point, in which case it becomes beneficial to use smaller .
As suggested by lower bounds on FedAvg in (Karimireddy et al., 2019) on FedAvg, our bounds do not show that FedAvg always converges faster than mini-batch SGD, and in fact, doing so may not be possible without further assumptions (such as a bound on the heterogeneity among clients) or more sophisticated optimization techniques (such as the use of control variates in SCAFFOLD (Karimireddy et al., 2019)).
While our theoretical results are only valid for the case of quadratic loss functions, we conjecture that even in much broader settings, the choice of learning rate still dictates a trade-off between accuracy and initial convergence. we will show in the next section that this holds empirically, even in non-convex settings.
Experimental results
In this section, we analyze LocalUpdate empirically, in order to understand how the choice of client and server learning rates impact convergence in more realistic machine learning tasks. In particular, we focus on (not necessarily convex) tasks and datasets that reflect federated learning in practice. We will show that both the choice of client learning rate and tuning of the corresponding server learning rate can be vital to attain the best performance of LocalUpdate, especially in limited communication settings.
We use four different datasets: the federated extended MNIST dataset (FEMNIST) (Caldas et al., 2018), the federated version of CIFAR-100 created by Reddi et al. (2020), the Shakespeare dataset (Caldas et al., 2018) and the Stack Overflow dataset (Authors, 2019). The first two are image datasets, the second two are text datasets. All datasets are publicly available. We specifically use the versions available in TensorFlow Federated (Ingerman and Ostrowski, 2019). All four datasets contain training and test clients. For the purposes of our experiments, we only use the training clients, as our work only concerns the loss of clients in the training population. Notably, our work does not broach the subject of generalization, which we leave to future work. The number of clients and examples in each dataset is presented in Table 1.
For FEMNIST, we train a moderately-sized CNN (the same as used in (McMahan et al., 2017)) to perform character recognition. For CIFAR-100, we train a ResNet-18 (where we replace the batch norm layers with group norm, as suggested by (Hsieh et al., 2019) and used by (Reddi et al., 2020)). For Shakespeare, we train an RNN with 2 LSTM layers to perform next-character-prediction. For Stack Overflow, we perform two distinct task: tag prediction (TP) and next-word-prediction (NWP). For Stack Overflow TP, we use a logistic regression classifier with one-versus-all classification. Note that this implies that Stack Overflow TP is a convex task. For Stack Overflow NWP, we train an RNN with 1 LSTM layer to perform next-word-prediction. These five tasks were previously analyzed by Reddi et al. (2020) for the purposes of comparing adaptive and non-adaptive federated optimization methods. For further details on datasets and models, see Appendix C.
We implement LocalUpdate in TensorFlow Federated (Ingerman and Ostrowski, 2019). In all experiments, the set of clients is finite and we let be the uniform distribution over these clients. Each is the uniform distribution over some finite set of client examples. We analyze the performance of LocalUpdate with (ie. FedAvg/Reptile with local steps) across the tasks discussed above. For FEMNIST, CIFAR-100, and Shakespeare, we sample clients per round, while for Stack Overflow, we sample clients (due to its much larger number of clients). In order to derive fair comparisons for different hyperparameter settings, we use a random seed to determine which clients are sampled at each round from . All plots are made using the same seed. We sample clients without replacement within each round, but with replacement across rounds. We use a batch size of for FEMNIST and CIFAR-100, for Shakespeare, and for Stack Overflow.
1 Fixed server learning rates
We first perform a comparable analysis to that in Section 4 above for FEMNIST, CIFAR-100, and Shakespeare. Namely, we fix the server learning rate , and see how the training loss varies as a function of the client learning rate . We vary over
and omit results that diverged due to the client learning rate being set too large. We plot the true loss function (defined in (2)) in Figure 2.
We see that on all three tasks, especially CIFAR-100, the choice of client learning rate can impact not just the speed of convergence, but what point the algorithm converges to. In general, we see very similar behavior to that in Section 4, despite the non-convex loss functions involved in all three tasks. For both FEMNIST and CIFAR-100, smaller client learning rates eventually reach lower training losses than higher learning rates. This is particularly evident in the results for CIFAR-100. While initially performs better than all other methods, it is eventually surpassed by , and ends up obtaining a comparable accuracy. Conversely, setting results in a sub-optimal training loss.
We also see that while suboptimality gaps exist for FEMNIST and Shakespeare, they are much smaller than for CIFAR-100. Thus, our results also suggest that the theoretical suboptimality of larger may not be as important a facet in practice. We see here that despite the asymptotic suboptimality of FedAvg with , in realistic settings where the number of communication rounds is limited, there may be little to no disadvantage to using . In fact, as we see for CIFAR-100, there may be advantages to using larger in settings with relatively few communication rounds.
2 Deriving fair comparisons between different client learning rates
In Figure 3, we plot for different choices of in the CIFAR-100 task. Specifically, we plot the mean value of for each consecutive 1000 rounds, as well as the standard deviation within those rounds.
One important caveat to this observation is that it is not enough to simply directly tie the client and server learning rate together. This is in fact what is done in the original incarnation of FedAvg in (McMahan et al., 2017). As discussed in Section 2.1, the original “vanilla” FedAvg algorithm corresponds to LocalUpdate with and . This does allow for the use of larger server learning rates with smaller client learning rates. However, this is not enough to necessarily derive optimal performance of LocalUpdate.
To demonstrate this, we plot the performance of LocalUpdate with and for the CIFAR-100 task in Figure 4. We see results in stark contrast to Figure 2(b). In particular, actually results in an increased training loss, and outperforms for most rounds. Moreover, by setting , we do not allow , even though necessarily results in a surrogate loss that does not match the true training loss. While this version of FedAvg is convenient from an implementation and hyperparameter tuning point of view, it does not result in the best performance of LocalUpdate. As we show in the next section, we can greatly improve performance by tuning client and server learning rates separately.
3 On the importance of tuned server learning rates
Based on the discussion in the section above, to give the most fair comparisons between different values of in LocalUpdate, we must also tune the server learning rate . We perform the same experiments as in Figure 2, but where we tune the server learning rate for each choice of . We vary over
and select that results in the smallest average training loss over the last 100 rounds. We note that we use this averaging method as a single round of federated learning only samples a small number of clients. We plot the loss of LocalUpdate under these settings in Figures 5. For a list of all client learning rates and the corresponding best server learning rate for each task, see Appendix D.
Notably, we still see the same general trend discussed in our convergence rates in Section 7. While large values of may obtain a smaller training loss initially, eventually smaller values of perform comparably, if not better. However, it is instructive to note that this may take many thousands of communication rounds. In particular, for all three tasks, we see that does not perform comparably larger until near the end of our training procedure. We thus see the following:
By performing appropriate server and client learning rate tuning, we can mitigate the suboptimality of FedAvg/Reptile for , especially when the number of communication rounds is limited.
In particular, our empirical results suggest that in realistic federated learning training tasks, the suboptimality of FedAvg/Reptile may only be an asymptotic concern. If we can only perform a limited number of training rounds, it often does benefit us to use larger values of . We note that this aligns with our discussion of how leads to condition number regularization (see Corollary 13).
4 Adaptive optimization and LocalUpdate
In order to understand the behavior of LocalUpdate more generally, we perform similar experiments on the Stack Overflow dataset. However, as shown by Reddi et al. (2020), performance of FedAvg on next-word-prediction tasks can be greatly improved by the use of adaptive optimization. In particular, Reddi et al. (2020) found that the use of the Yogi optimizer (Zaheer et al., 2018) improved performance in a wide variety of settings, including on the same tag-prediction and next-word-prediction tasks. Note that the Yogi optimizer is similar to the Adam optimizer (Kingma and Ba, 2014), except that it uses a kind of additive adaptive update that can improve convergence by making more controlled progress. For more details, see (Zaheer et al., 2018).
Thus, for these tasks, we use a modified version of LocalUpdate in which the server uses the client update as an estimate of the gradient of the loss function, and applies Yogi to this gradient. That is, we use Algorithm 1, but update the model via
We use the version of Yogi proposed by Zaheer et al. (2018), with first momentum term of , second momentum parameter of , an initial accumulator value of , and an value of . We then perform analogous experiments to those above, where we first fix and vary . However, due to the size of the Stack Oveflow dataset (see Table 1), we did not compute the total loss over all clients. Instead, we plot the average loss of the clients that participated in a given round, before local training occurs. This “loss at current round” can be viewed as a stochastic estimate of the true loss . This loss is plotted in Figure 6. We also perform experiments where we vary and tune . We select the value of with the smallest average training loss over the last 100 communication rounds. The result is given in Figure 7. For a list of all client learning rates and the corresponding best server learning rate for each task, see Appendix D.
For fixed , large leads to an initially smaller loss that is eventually beaten or matched by smaller . However, we see similar behavior among all but the largest in both tasks, potentially due to the adaptivity in Yogi. Looking at Figure 7, we see that tuning the server learning rate has slightly different effects on the two tasks: While Stack Overflow TP with tuned leads to the smaller performing better throughout the training process, Stack Overflow NWP with tuned actually allows larger to achieve lower loss throughout. Notably, we see that the gap between large and small in Stack Overflow NWP winnows as the number of rounds increases. Our results reinforce the notion that in communication-limited settings, server learning rate tuning is critical to ensure the best performance possible. This holds even when not using SGD on the server, but instead using an adaptive optimizer.
Automatic learning rate decay
In short, our results in the section above suggest that depending on the desired number of communication rounds, we may wish to use different client learning rates. Unfortunately, this requires a large degree of hyperparameter tuning. Both the client and server learning rate must be tuned. However, if the number of communication rounds in a federated learning system is limited, it may not be feasible to conduct extensive hyperparameter tuning, as the communication rounds required to do so may be better utilized by training your model for more rounds.
To help reduce the amount of learning rate tuning required, we propose a method for automatic learning rate decay that helps mitigate the need for client learning rate tuning. Our method will utilize our theoretical and empirical observations above showing that large client learning rates should be used initially, while smaller learning rates eventually reach a lower training loss. By decaying the client learning rate automatically over time, we mitigate the need to tune it. Instead, it can be set to any moderate value that does not result in divergent behavior on clients. We will be particularly concerned with systems-level constraints (such as those encountered in federated learning) when describing our method.
One particularly important restriction in many local update settings is the ability to compute the training loss . Recall that we defined
where is a distribution over clients, and is the client’s data distribution. In settings with limited communication, it may not be feasible to sample most or even a moderate fraction of the clients from . Moreover, sampling a client purely for the purposes of estimating the loss may be a waste of resources, as that client could be used for training purposes.
where is as in (24). This is the average loss of all clients in the last rounds. Note that Algorithm 3 corresponds to . We would then decay the client and server learning rates if
Second, even using moving windows to estimate , the heterogeneity of clients can still cause problematic variance. Thus, one can instead decay the learning rates if there has been no progress (relative to ) for consecutive rounds. That is, we keep a counter for how many consecutive rounds the condition (25) holds. If this counter ever reaches , we then decay and . If
then we reset the counter to 0. Note that Algorithm 3 corresponds to .
Last, it is often useful to have a cooldown period , where after decaying the learning rates and , we do not decay the learning rate for the next rounds. This is beneficial both for recovering a new estimate of the loss function, and for ensuring that the learning rate does not decay too frequently. Additionally, we recommend using this cooldown period for the first rounds as well, as this allows one to develop a better estimate of before decaying the learning rate.
While in practice, setting , and may seem difficult, we found that setting them all to be the same value led to good behavior across datasets and tasks. Moreover, the value can be estimated using simple heuristics. For example, suppose that we think that randomly sampled clients are sufficient to give a good representation of . Then, should serve as a default value for these parameters. In practice, we found that fixing all three values to some moderate constant (ex. ) was sufficient, even across datasets with widely varying numbers of clients.
2 Empirical evaluation
The primary motivation for LocalUpdateDecay is removing the need for client learning rate tuning. Intuitively, we can use any moderately large value of that results in non-divergent client behavior, and this will be gradually scaled back over time as we reach suboptimal critical points of the corresponding surrogate loss. To validate this, we compare LocalUpdateDecay to LocalUpdate on FEMNIST, CIFAR-100, Shakespeare, and Stack Overflow. For Stack Overflow, we again use a modified version where we apply Yogi on the server. In particular, we compare tuned but constant and , to LocalUpdateDecay.
We tune no other hyperparameters. In LocalUpdateDecay, specifically Algorithm 3, we set . We also use the practical refinements discussed in Section 9.1, setting . All other parameters are identical to that of standard LocalUpdate.
For FEMNIST, CIFAR-100, and Shakespeare, we plot the results in Figure 8, where we also compare to the tuned results in Section 8.3. We find that in all three tasks, LocalUpdateDecay eventually does as well as LocalUpdate with tuned , without the need for client learning rate tuning. For FEMNIST and Shakespeare, we find that LocalUpdateDecay almost immediately does better than LocalUpdate and continues to do at least as well throughout the course of training, often better. While this is not true for CIFAR-100, the results are still instructive. We see that while the client learning rate is initially set to a suboptimal value (), the automatic learning rate decay enables us to move away from this suboptimal basin and towards something comparable to the best tuned client learning rate after enough rounds. We see that despite initializing with a bad , the non-divergence in earlier rounds is sufficient to allow eventually near-optimal performance in the later rounds.
For both Stack Overflow tasks we plot an analogous results in Figure 9. As discussed in Section 8.4, due to the size of the Stack Overflow dataset, we do not plot the loss over all clients. Instead, we plot the average loss among all clients in each round before training. For Stack Overflow NWP, we see that LocalUpdateDecay performs comparably to the best tuned client learning rate. However, for Stack Overflow TP, we see that the decay actually helps significantly. While we initialize with a suboptimal client learning rate , (which clearly achieves higher loss throughout), by decaying the client learning rate we are able to obtain comparable or lower loss than all other fixed .
Open questions
Our work above opens up a number of possible follow-up directions in the area of federated optimization. First, we expect the same kind of analysis obtained above to apply to methods similar to LocalUpdate. For example, while the FedProx algorithm (Li et al., 2019) does not fit the format of LocalUpdate as presented in this work, we believe that it can be analyzed in a similar way. As shown by Pathak and Wainwright (2020), even in the non-stochastic setting, FedProx is not optimizing the true loss function. Thus, a natural question is to understand exactly what loss function is being optimized, and how the structure of FedProx encourages convergence over FedAvg. Another natural algorithm for analysis is a more general version of LocalUpdate in which the client learning rate is decayed during a client’s local computation. This may help combat adverse effects incurred by setting to be too large.
More generally, we would also like to understand the behavior of LocalUpdate on non-quadratic functions. Even generalizing the analysis above to the strongly convex case would be substantial progress towards understanding federated learning in heterogeneous settings. While there may be no surrogate loss that LocalUpdate is directly optimizing through SGD, we believe that the intuition behind our work can still be carried forward for more general loss functions. In particular, one might expect that the optimization dynamics of LocalUpdate can be parameterized in terms of and , even in non-quadratic settings, and that the selection of these parameters governs a trade-off between the speed of convergence and the accuracy of the resulting critical point.
Another interesting open direction is determining how to best set for a given problem. As seen above, the choice of drastically alters optimization dynamics. While it is often chosen in an ad hoc manner (based in part on the cost of communication), one could imagine attempting to minimizing the number of rounds need to obtain a given accuracy level with respect to . Even for FedAvg/Reptile, it is not clear how to set the number of local steps . Insights into this could greatly improve the performance of federated learning algorithms.
Finally, we note that our work is fundamentally concerned with the training dynamics of local update methods. In practice, we are often instead interested in the generalization ability of a model. We suspect that the choice of parameters in LocalUpdate can have large implications for generalization ability, the study of which we leave to future work.
We would like to thank Keith Rush and Sai Praneeth Karimireddy for the remarks that accidentally helped spark this work. We would also like to thank H. Brendan McMahan and Zachary Garrett for fruitful discussions about decoupling client and server learning rates in federated learning. Finally, we gratefully acknowledge Shanshan Wu for insights on personalization in federated learning.
A Relation between FedAvg, Local SGD, and LocalUpdate
In this section, we formalize the connection between FedAvg, Local SGD, and LocalUpdate. First, we note that in common descriptions of FedAvg and Local SGD algorithms (see (McMahan et al., 2017) and (Stich, 2019) for example), these two are effectively the same algorithm. The difference in nomenclature often reflects the distributed setting: Local SGD is often referred to in settings with homogeneous data, while FedAvg is often used in heterogeneous settings, especially for the purposes of federated learning.
We will use the following (simplified) version of the algorithms: At each iteration of FedAvg/Local SGD, we have some set of clients of size . Each client receives the server’s model , and applies steps of mini-batch SGD updates to its local model to create an updated local model . The server then updates its model via
Fix , and let denote the -th mini-batch gradient of client . Suppose we use a learning rate of on each client when performing mini-batch SGD. Then we have
This is exactly LocalUpdate with and . However, by allowing to vary independently of in LocalUpdate, we can obtain markedly different convergence behavior. We note that a form of this decoupling has previously been explored by Karimireddy et al. (2019) and Reddi et al. (2020). However, these versions instead perform averaging on the so-called “model delta” (see (Reddi et al., 2020)), in which the server model is updated via
Thus, while this does decouple and to some degree, it does not fully do so. In particular, if we set , then , in which case we can make no progress overall. This is particularly important because, as we show above, for many , the only way for the surrogate loss to have the same critical point as the true loss is by setting . More generally, we see that the effective learning rate used in such an update is actually the product , which can result in conflating the effect of changes in with changes in .
B Proof of results
Proof Using the fact that is symmetric for all , we have
B.1.2 Lemma 2
Proof For simplicity of notation, let and . Let . By (5), is independent of given . Therefore, we have that for any ,
By linearity of expectation, we then find that
By the assumption that is positive definite, this implies
B.1.3 Theorem 4
Proof For convenience of notation, we will fix and let denote , and let . Since we assume this is computed using gradient descent, we have
Note that for all . Therefore, for ,
This last step follows from (29) and (30), and the fact that is symmetric for all . The remainder of the result follows directly from Theorem 3.
B.2 Results from Section 5
In either event, we see that where . Therefore,
B.2.2 Lemma 6
Proof of Property 1: Since and is symmetric and positive definite we know that the matrix and therefore is symmetric and positive definite for all . By (8) and Assumption 1, is a nonnegative linear combination of positive definite, symmetric matrices, with some coefficient . It is therefore positive definite and symmetric.
Proof of Property 2: Since and are symmetric and positive definite, so too is their product. Therefore, the -th surrogate loss defined in (9) is a positive definite quadratic function, and therefore has a unique minima. Solving for this minima explicitly by setting the gradient equal to 0, we get
Here, we again used the fact that is symmetric and positive definite.
Proof of Property 3: Let be an eigenvector of with eigenvalue . Then note that is an eigenvector of with eigenvalue . Therefore, is an eigenvector of with eigenvalue
Proof of Property 4: This follows from Property 3, noting that since , we have for all eigenvalues of . By Property 3, the eigenvalue is maximized when and minimized when .
B.2.3 Lemma 10
Proof of Property 1: This follows by analogous reasoning to Property 3 in Lemma 6. Note that every eigenvector of with associated eigenvalue is also an eigenvector of with eigenvalue . By basic properties of eigenvectors, this implies the desired result.
Proof of Property 2: By Lemma 6, and because , for all ,
Since and are positive definite, we have
Repeatedly using the fact that for Hermitian matrices, , we derive the upper bound on . The lower bound on follows from a similar argument, using the fact that for Hermitian matrices, .
B.2.4 Lemma 11
Proof of Property 1: This follows immediately from Property 1 of Lemma 10 and (18).
Proof of Property 2: Fix and define
Since , for , we have . Hence, for we know
Therefore, are the minimum and maximum eigenvalues of .
Proof of Property 3: As in the proof of Property 2, for a fixed define
Note that by assumption, we have . Therefore, for , , so for all ,
B.2.5 Lemma 9
Proof Because , we have
is a partial geometric sum of matrices with eigenvalues satisfying . This implies that, much like a geometric series of scalars,
Note that here we used the fact that commutes with to interchange their order.
B.2.6 Lemma 12
Proof of Property 1: This follows immediately from Property 1 of Lemma 10 and (3).
Proof of Property 2: Fix and define
Since , for , we have
In particular, this implies , so . Hence, for we know
Therefore, are the minimum and maximum eigenvalues of .
Proof of Property 3: As in the proof of Property 2, for a fixed define
By assumption, . Therefore, for , , so for such ,
B.3 Results from Section 6
In order to prove the results in this section, we will use the following straightforward lemma regarding the structure of .
Scaling by in both parts of the right-hand side, we derive the result.
Let be a random matrix such that is always real, symmetric, and positive definite. Then
In particular, this implies that for such a random matrix , we have
Proof We will first prove analogous statements for any . Fix . By Lemma 11, we have
Since , by (17), we have , implying that
Since , this implies
By the dominating convergence theorem, we have that for any ,
B.3.2 Theorem 16
Proof Throughout this proof, we will use the fact that for a positive semi-definite, symmetric matrix , and that if is further positive definite, .
For , we again use the Cauchy-Schwarz inequality, as
Note that by Assumption 2, we have that . By Lemma 6 and the definition of , this implies that for all ,
By Lemma 6 and Assumption 2, we find that for all ,
B.3.3 Theorem 17
Note that for , the quantity defined in is given by . We will require one further auxiliary lemma.
If , then for all we have
where is as defined in (17).
Proof First note that since ,
As discussed in the proof of Lemma 10 and Lemma 11, has eigenvalues of the form where is an eigenvalue of . Therefore, has eigenvalues of the form
where is an eigenvalue of . As noted in (18), for ,
Therefore, is symmetric and positive semi-definite. Moreover, a simple computation shows
which is nonnegative for . In particular it is nonnegative for (as in Assumption 2), implying that
Proof [Proof of Theorem 17] We begin our proof in a similar manner to the proof of Theorem 16. Note that for FedAvg with local steps, . By Lemma 30,
Note that the last step follows by (17) and Lemma 11.
This last step holds by Lemma 32 and Assumption 3.
For , we again use the Cauchy-Schwarz inequality, as
Putting this all together, we have derive the result.
B.3.4 Lemma 18
Proof By properties of geometric sums, we have
Note that this implies that the function is differentiable everywhere. Taking a derivative, we have
Note that all terms in this sum are nonnegative when . By assumption on , we have
It therefore suffices sto show the desired lower bound on when
By definition of , we have
Let . Note that by assumption on , . It then suffices to show that
Letting , we then use the fact that for ,
B.3.5 Theorem 21
We will first require an auxiliary lemma.
Suppose and Assumption 2 holds. Then for all , the matrix
Proof Recall that by Lemma 11, the eigenvalues of are of the form
Since and share the same eigenvectors as , the eigenvalues of are of the form where is an eigenvalue of . Since , we clearly have for , implying that is positive semidefinite.
For the maximum eigenvalue of , we consider . A simple calculation shows
Since , we find that for . Therefore, the maximum eigenvalue of satisfies
We can now use this to prove the desired result, in a manner similar to the proof of Theorem 17.
Proof [Proof of Theorem 21] For notational convenience, we define the following quantities (where
As in the proof of Theorem 16, we can use the fact that and are symmetric and positive definite to bound , as we have
The penultimate step follows by Lemma 31. By definition of ,
This last step follows directly from Lemma 11.
Here, this last step follows from Lemma 33.
Note that can be bounded in the same manner as above. For the remaining term, we have
B.3.6 Corollary 22
To prove this, we will need to first bound the term . To do so, we first require a simple lemma regarding ratios of sums.
For , let be positive real numbers such that for ,
Proof We will prove this inductively. Note that when , the result immediately follows by assumption.
For , applying the inductive hypothesis to , we have
Let . Note that , so applying the inductive hypothesis we have
We can now derive a bound on the condition number .
For , and ,
Since , . Therefore,
With this in hand, we can prove Corollary 22.
Proof [Proof of Corollary 22] By definition of , we have
Here the second line follows from the fact that , and the third line follows from the fact that . Therefore,
The result follows by applying Lemma 35 and noting that .
B.4 Results in Section 7
where (Algorithm 2) and is a set of size sampled independently and uniformly at random from . Since the are independent, it suffices to show that for any , the vector satisfies
where is a mini-batch stochastic gradient of batch size taken at , and the are updated via
By Assumption 4, we have that for any and ,
where and each is identically and independently distributed, we have
B.4.2 Lemma 25
Proof We will proceed using a similar analysis to Lemma 1 in Rakhlin et al. (2012).
Let . Recall that we have
Using the equations (37) and (38), we have
Applying Lemma 23, (39) and (22), we have
B.4.3 Theorem 27
Proof We will proceed using similar techniques to those in Theorem 4.7 of Bottou et al. (2018). Note that by construction of , we have that for all ,
We then proceed by induction. For , we have
For , let . Therefore, . Using (40) and the inductive hypothesis,
Note that by assumption on , and that simple analysis shows
B.4.4 Corollary 28
Using the bounds in Theorems 27 and Corollary 20, we have
By Lemma 18 and assumption on , we have
Using this to derive an upper bound on the first term in , we conclude the proof.
B.4.5 Theorem 29
For convenience of notation (and in a slight abuse of previous notation), we define
Therefore, we can apply Lemma 25 (with , ) to find
For any , we therefore have
Using (41), we derive the following recursion on the .
Let and let . . We will use (43) inductively to show that
Similar analysis can be done in the case that . When , using the inductive hypothesis, we have
However, since , by Lemma 18, we know
Therefore, multiplying (46) by , we have
This bounds the first part of (43). For the second part, we will use Corollary 22. In particular, since
Here we used the fact that . Multiplying by ,
This proves (44). To get a bound on the distance to the minima , we then have
We can again use Corollary 22, letting in the statement of that result. We then get
This again uses the fact that . Combining, this implies
Substituting in , this proves the desired result.
C Datasets and Models
Below, we provide detailed description of the datasets and models used in the paper. We use federated versions of vision datasets FEMNIST (Caldas et al., 2018) and CIFAR-100 (Krizhevsky and Hinton, 2009), and language modeling datasets Shakespeare (McMahan et al., 2017) and StackOverflow (Authors, 2019). We give descriptions of the datasets, models, and tasks below.
The CIFAR-100 dataset is a popular computer vision dataset consisting of images with 100 possible labels. While this dataset is not a federated dataset, a federated version was created by Reddi et al. (2020), using hierarchical latent Dirichlet allocation to enforce moderate amounts of heterogeneity among clients. The resulting dataset has 500 clients, each with 100 unique examples. We train a ResNet-18 on this dataset, where we replace all batch normalization layers with group normalization layers (Wu and He, 2018). The use of group norm over batch norm in federated learning was first advocated by Hsieh et al. (2019).
We perform small amounts of data augmentation and preprocessing, as is standard with CIFAR-100. We first perform a random crop to shape , followed by a random horizontal flip. We then normalize the pixel values according to their mean and standard deviation. Thus, given an image , we compute where is the average of the pixel values in , and is the standard deviation.
FEMNIST consists of gray-scale images of both numbers and upper- and lower-case English characters, with 62 possible labels in total. The digits are partitioned according to their author, resulting in a naturally heterogeneous federated dataset. We do not use any preprocessing on the images. We train a moderately-sized CNN, with identical architecture to the CNN used by McMahan et al. (2017). The CNN contains two convolutional layer, each with kernels. The convolutional layers have 32 and 64 filters, respectively, and are each followed by a max pooling layer. Finally, the model has a dense layer with 512 units and ReLU activation, followed by a softmax activation.
The Shakespeare dataset is derived from the benchmark designed by Caldas et al. (2018). The dataset corpus is the collected works of William Shakespeare, and the clients correspond to roles in Shakespeare’s plays with at least two lines of dialogue. To eliminate confusion, character here will refer to alphanumeric and other such symbols, while we will use client to denote the various roles in plays. We split each client’s lines into sequences of 80 characters, padding if necessary. We use a vocabulary size of 90: 86 characters contained in Shakespeare’s work, beginning and end of line tokens, padding tokens, and out-of-vocabulary tokens. We perform next-character prediction on the clients’ dialogue using an RNN. The RNN takes as input a sequence of 80 characters, embeds it into a learned 8-dimensional space, and passes the embedding through 2 LSTM layers, each with 256 units. Finally, we use a softmax output layer with 80 units, where we try to predict a sequence of 80 characters formed by shifting the input sequence over by one. Therefore, our output dimension is . We compute loss using cross-entropy loss.
Stack Overflow is a text datasets consisting of questions and answers posted to the Stack Overflow website. Each user is a client, and their datasets consist of questions and answers posted by this user. Each post has associated meta-data, including a list of associated tags (e.g. a post could have the tag javascript if it concerns the javascript language). We perform two tasks on this dataset: tag prediction, and next word prediction. In both cases, we restrict to the 10,000 most frequently used words in the total dataset, as well as the 500 most frequently used tags for the tag prediction task.
For Stack Overflow tag prediction, we use a multi-class logistic regression classifier with 500 output units (one for each of the 500 most frequently used tags), and adopt a one-versus-rest classification strategy. Note that the corresponding multi-class logistic loss is convex. The inputs to our model are 10,000-dimensional vectors forming bag-of-words vectors for each post. Each vector is normalized to have sum 1.
For Stack Overflow next word prediction, we restrict each client to the first 128 posts in their history (for computational efficiency reasons, as some clients have tens of thousands of posts). We perform truncation and padding so that each post has 21 words (including word tokens for beginning of sentence, end of sentence, padding, and out-of-vocabulary words). The sequence is split into input and output length-20 sequences, corresponding to the first and the last 20 characters (ie. one is the other sequence, shifted by one). The first of these sequences is embedded into a learned 96-dimensional space, and then fed into an LSTM with 670 units. Finally, the output is fed into a densely connected softmax layer with 10,004 units (corresponding to the 10,000 in-vocabulary words, and the extra tokens mentioned above). We attempt to predict the shifted-by-one sequence, and compute the loss via cross-entropy.
D Tuned Server Learning Rates
In this section, we detail the best server learning rate found for each corresponding client learning rate and task.