Convergence and Accuracy Trade-Offs in Federated Learning and Meta-Learning
Zachary Charles, Jakub Konečný
Introduction
Federated learning (McMahan et al., 2017) is a distributed framework for learning models without directly sharing data. In this framework, clients perform local updates (typically using first-order optimization) on their own data. In the popular FedAvg algorithm (McMahan et al., 2017), the client models are then averaged at a central server. Since the proposal of FedAvg, many new federated optimization algorithms have been developed (Li et al., 2020a; Reddi et al., 2020; Hsu et al., 2019; Xie et al., 2019; Basu et al., 2019; Li et al., 2020b; Karimireddy et al., 2019). These methods typically employ multiple local client epochs in order to improve communication-efficiency. We defer to Kairouz et al. (2019) and Li et al. (2019) for more detailed summaries of federated learning.
Local updates have also been used extensively in meta-learning. The celebrated MAML algorithm (Finn et al., 2017) employs multiple local model updates on a set of tasks in order to learn a model that quickly adapt to new tasks. MAML has inspired a number of model-agnostic meta-learning methods that also employ first-order local updates (Balcan et al., 2019; Fallah et al., 2020a; Nichol et al., 2018; Zhou et al., 2019). There are strong connections between federated learning and meta-learning, despite differences in practical concerns. Formal connections between the two were shown by Khodak et al. (2019) and have since been explored in many other works (Jiang et al., 2019; Fallah et al., 2020b).
We refer to methods that utilize multiple local updates across clients (or in the language of meta-learning, tasks) as local update methods (see Section 2.1 for a formal characterization). In practice, local update methods frequently outperform “centralized” methods such as SGD (McMahan et al., 2017; Finn et al., 2017; Hard et al., 2018; Yang et al., 2018; Hard et al., 2020). However, the empirical benefits of local update methods are not fully explained by existing theoretical analyses. For example, Woodworth et al. (2020) show that FedAvg often obtains convergence rates comparable to or worse than those of mini-batch SGD.
We focus on two difficulties that arise when analyzing local update methods. First, analyses must account for client drift (Karimireddy et al., 2019). As clients perform local updates on heterogeneous datasets, their local models drift apart. This hinders convergence to globally optimal models, and makes theoretical analyses more challenging. Similar phenomena were examined by Li et al. (2020a); Malinovsky et al. (2020); Pathak and Wainwright (2020) and Fallah et al. (2020b), who show that various local update methods do not converge to critical points of the empirical loss.
Second, local update methods are difficult to compare. Analyses of different methods may use different hyperparameters regimes, or make different assumptions. Even comparing seemingly similar methods can require significant theoretical insight (Karimireddy et al., 2019; Fallah et al., 2020b). Moreover, comparisons can be made in fundamentally different ways. One may wish to maximize the final accuracy, or minimize the number of communication rounds needed to attain a given accuracy. Thus, it is not even clear how local update methods should be compared.
In this work, we invert the conventional narrative that issues such as client drift harm convergence. Instead, we view such phenomena as improving convergence, but to sub-optimal points.
More generally, we show that local update methods face a fundamental trade-off between convergence and accuracy that is explicitly governed by algorithmic hyperparameters. Perceived failures of methods such as FedAvg actually correspond to operating points prioritizing convergence over accuracy. We use this trade-off to develop a novel framework for comparing local update methods. We compare methods based on their entire convergence-accuracy trade-off, not just their convergence to optimal points. In more detail:
We show that for quadratic models, local update methods are equivalent to optimizing a single surrogate loss function. The condition number of the surrogate is controlled by algorithmic choices. Popular local update methods, including FedAvg and MAML, reduce the surrogate’s condition number, but increase the discrepancy between the empirical and surrogate losses. Our results also encompass proximal local update methods (Li et al., 2020a; Zhou et al., 2019).
We derive novel convergence rates that showcasing this trade-off between convergence and accuracy. Our bounds demonstrate the benefit of local update methods over methods such as mini-batch SGD in communication-limited settings.
We use this theory to develop a framework for comparing local update methods through a novel Pareto frontier, which compares convergence-accuracy trade-offs of classes of algorithms. We use this to derive novel comparisons of many popular local update methods.
We use this technique to shed light on a broad range of phenomena, including the benefit of server momentum, the effect of proximal local updates, and differences between the dynamics of FedAvg and MAML.
While our theoretical results are restricted to quadratic models, we show that such convergence-accuracy trade-offs occur empirically in non-convex settings. We also validate our theoretical observations regarding server momentum and proximal updates on a non-convex task.
We view our work as a step towards holistic understandings of local update methods. Using the aforementioned Pareto frontiers, we highlight a number of new phenomena and open problems. One particularly intriguing observation is that the convergence-accuracy trade-off for FedAvg with heavy-ball server momentum appears to be completely symmetric. For more details, see Section 5. Our proof techniques may be of independent interest. We derive a novel analog of the Bhatia-Davis inequality (Bhatia and Davis, 2000) for mean absolute deviations, and use this to understand the accuracy of local update methods.
Notation
Accuracy and Meta-Learning
We study the accuracy of local update methods on the training population. However, meta-learning algorithms are designed to learn a model that adapts well to new tasks; The empirical loss is not necessarily indicative of the “post-adaptation” accuracy of such methods (Finn et al., 2017). Despite this our focus still yields novel insights into qualitative differences between the training dynamics of federated learning and meta-learning methods. Perhaps surprisingly, we show that in certain hyperparameter regions, these methods exhibit identical trade-offs between convergence and pre-adaptation accuracy (see Figures 4 and 5). While we believe our results can be adapted to post-adaptation accuracy via techniques developed by Fallah et al. (2020a), we leave the analysis to future work.
Problem Setup
For , we define the client loss function and the overall loss function as follows:
The joint distribution defines a distribution over , recovering standard risk minimization, as well as distributed risk minimization in which and all are uniform over finite sets. For , define:
We assume these expectations exist and are finite. One can show that up to some additive constant,
We make the following assumptions throughout.
There are such that for all , .
There is some such that for all , .
After receiving all client updates, the server treats their average as an estimate of the gradient of the loss function , and applies to a first-order optimization algorithm ServerOpt. For example, the server could perform a gradient descent step using the “pseudo-gradient” . We refer to this process (parameterized by and ServerOpt) as LocalUpdate and give pseudo-code in Algorithms 1 and 2.
LocalUpdate recovers many well-known algorithms for various choices . For convenience, define
Special cases of LocalUpdate when ServerOpt is gradient descent are given in Table 1. For details on the relation between FedAvg and LocalUpdate, see Appendix A. By changing ServerOpt, we can recover methods such as FedAvgM (Hsu et al., 2019) (server gradient descent with momentum), and FedAdam (Reddi et al., 2020) (server Adam (Kingma and Ba, 2014)).
Local Update Methods as First-Order Methods
LocalUpdate can vary drastically from first-order optimization methods on the empirical loss. Despite this, we will show that Algorithm 2 is equivalent in expectation to ServerOpt applied to a single surrogate loss. This surrogate loss is determined by the inputs and to Algorithm 2. For each client , we define its distortion matrix as
We define the surrogate loss function of client as
and the overall surrogate loss function as
When , , in which case there is no distortion. In general, can amplify the heterogeneity of the . We derive the following theorem linking the surrogate losses to Algorithm 2.
A version of Theorem 1 was shown for by Fallah et al. (2020b). We take this a step further and show that in certain settings, MAML is equivalent in expectation to ServerOpt on a surrogate loss.
MAML with local steps can be viewed as a modification of LocalUpdate. Algorithm 1 remains the same, and in Algorithm 2, each client executes mini-batch SGD steps. However, the client’s message to the server is different. Let be the function that runs steps of mini-batch SGD, starting from , for fixed mini-batches of size drawn independently from . Define
Each client sends a stochastic estimate of to the server. The rest is identical to LocalUpdate; The server averages the client outputs and uses this as a gradient estimate for ServerOpt. While MAML is not a special case of LocalUpdate, we show that if the clients use gradient descent, MAML is equivalent in expectation to LocalUpdate with .
If is the function that runs steps of gradient descent on with learning rate starting at , then
Convergence and Accuracy of Local Update Methods
We wish to better understand (9) in cases of interest. We first consider , as in FedAvg. Define
When , we recover the condition number of the empirical loss . We next consider , as in MAML-style algorithms. Define
As or , , which bounds the condition number of the empirical loss . If is not close to 0, we get an exponential reduction (in terms of ) of the condition number. While the analysis is not as clear for , one can show that , with equality if and only if , and either or . Moreover, decreases as or . For both and , increasing decreases .
Here we see the impact of local update methods on convergence: Popular methods such as FedAvg, FedProx, MAML, and Reptile reduce the condition number of the surrogate loss function they are actually optimizing. In the next section, we translate this into concrete convergence rates for LocalUpdate.
2 Convergence Rates
Suppose and ServerOpt is gradient descent with Nesterov, heavy-ball, or no momentum. Then for some hyperparameter setting of ServerOpt, and as in Table 2, the iterates of LocalUpdate satisfy
Thus, (properly tuned) server momentum improves the convergence of LocalUpdate, giving theoretical groundingYuan and Ma (2020) first showed that momentum can accelerate FedAvg, though they use a different momentum scheme with extra per-round communication. to the improved convergence of FedAvgM shown by Hsu et al. (2019) and Reddi et al. (2020). Since ServerOpt does not change the surrogate loss, this improvement in convergence does not degrade the accuracy of the learned model.
3 Distance Between Global Minimizers
We are interested in . While we focus on the setting where is a discrete distribution over some finite , our analysis can be generalized to arbitrary probability spaces . We derive the following bound.
Let and . Then
When , we can reduce the constant factor to , which we show is tight (see Appendix C.2). While we conjecture that this bound holds with a constant of for all , we leave this to future work.
Our proof technique for Lemma 5 may be of independent interest. We derive this result by first proving an analog of the Bhatia-Davis inequality (Bhatia and Davis, 2000) for mean absolute deviations of bounded random variables (Theorem 5 in Appendix C).
Let . Specializing to or , we derive a link between in (11) and (13) and the distance between optimizers.
Suppose that either (I) and or (II) and . Then for all , .
Applying Theorem 3, we bound the convergence of LocalUpdate to the empirical minimizer .
Under the same settings as Theorem 3, for some hyperparameter setting of ServerOpt, the iterates of LocalUpdate satisfy
where is given in (11) and (13), and is given in Table 2.
Here we see the benefit of local update methods in communication-limited settings. When is small and is large, we can achieve better convergence by decreasing and leaving the second term fixed. In such settings, FedAvg can arrive at a neighborhood of a critical point in fewer communication rounds than mini-batch SGD, but may not ever actually reach the critical point. If is small, we may be better served by using mini-batch SGD instead.
Comparing Local Update Methods
Comparing optimization algorithms is a fundamental theoretical effort. Many past works compare local update methods based on their convergence to critical points of the empirical loss. By Theorem 1, LocalUpdate is only guaranteed to converge to critical points of if or . Thus, existing analyses ignore many useful cases of LocalUpdate.
To remedy this, we compare local update algorithms on the basis of both convergence and accuracy. Instead of fixing and , we analyze LocalUpdate as and vary. To do so, we use our theory from Section 4. Given and , we define the convergence rate as the infimum over all such that for all , (14) holds. Values of when ServerOpt is gradient descent are given in Table 2. For or , we define the suboptimality by
By Theorem 4, this captures the asymptotic worst-case suboptimality of LocalUpdate.
Note that . Therefore, fixing and ServerOpt, we obtain a Pareto frontier in by plotting for various and . This curve represents the worst-case convergence/accuracy trade-off of a class of local update methods. We generally want the curve to be as close to as possible.
For example, in Figure 1 we let ServerOpt be gradient descent and set . We plot as we vary and fix , and vice-versa. When , we obtain nearly identical curves. The curves for are similar, except that when we fix and vary , we do not reach . While and have similar impacts on convergence-accuracy trade-offs, varying leads a larger set of attainable . Formally, this is because in (11), . Intuitively, recovers one-shot averaging while does not. Notably, the convergence-accuracy trade-off becomes closer to a linear trade-off as decreases.
The Pareto frontiers contain more information than just the convergence rate to a critical point (the curve’s intersection with the -axis). This information is useful in communication-limited regimes, where we wish to minimize the number of rounds needed to attain a given accuracy. The curves also help visualize various hyperparameter settings of an algorithms simultaneously. To illustrate this, we use the Pareto frontiers to derive novel findings regarding server momentum, proximal client updates, and qualitative differences between FedAvg and MAML. The results are all given below. For more results, see Appendix D.
As shown empirically by Hsu et al. (2019) and as reflected in Table 2, server momentum can improve convergence. To understand this, in Figure 2 we compare Pareto frontiers where and ServerOpt is gradient descent with various types of momentum (Nesterov, heavy-ball, or no momentum). We see a strict ordering of the server optimization methods. Heavy-ball momentum is better than Nesterov momentum, which is better than no momentum.
One important finding is that the benefit of momentum is more pronounced as increases. On the other hand, the benefit of server momentum diminishes for sufficiently large : In Figure 4, the various types of momentum lead to similar suboptimality when the convergence rate is close to 0. Intuitively, as , we recover one-shot averaging, which converges in a single communication round with or without momentum.
Another intriguing observation: The Pareto frontiers for heavy-ball momentum appear to be symmetric about the line . We conjecture this is true for any . While we believe that this may be provable by careful algebraic manipulation of our results above, ideally a proof would explain the root causes of this symmetry. Thus, we leave a proof to future work.
Proximal Client Updates
So far we have only considered . One might posit that as varies, the Pareto frontier moves closer to the origin. This appears to not be the case. In all settings we examined, changing did not bring the Pareto frontier closer to . Instead, the frontier for was simply a subset of the frontier for .
To illustrate this, we plot Pareto frontiers for varying in Figure 3. As increases, the frontier becomes a smaller subset of the frontier for . Thus, proximal client updates may not enable faster convergence. Rather, their benefit may be in guarding against setting too small or too large. Figure 3 shows that FedAvg can always attain the same as FedProx, but it may require different hyperparameters. The reverse is not true, as FedProx cannot recover one-shot averaging. Our findings are consistent with work by Wang et al. (2020), who show that FedProx can reduce the “objective inconsistency” of FedAvg, at the expense of increasing convergence time.
Comparing MAML to FedAvg
We now turn our attention to comparing FedAvg-style algorithms () to MAML-style algorithms (. We plot Pareto frontiers for the guaranteed by Theorems 3 and 4. The results are in Figure 4.
For each ServerOpt, the MAML frontier is a subset of the FedAvg frontier. Recall that in Theorem 3, we require for FedAvg, but for MAML. In Figure 4 this causes the frontier for to be more restrictive than for . However, it is still notable that these two fundamentally different methods, attain the same frontier when is large.
The Pareto frontiers are identical for small (mirroring Figure 4), but diverge when . The frontier for MAML then moves further from 0. Intuitively, FedAvg tries to learn a global model, while MAML tries to learn a model that adapts quickly to new tasks (Finn et al., 2017); MAML need not minimize . The MAML frontier is noisy for (as , depend on random eigenvalues of ), but stabilizes for . While we posit that this reflects a semi-circle law for eigenvalues of random matrices (Alon et al., 2002), we leave an analysis to future work.
One final observation that highlights the similarities and differences of FedAvg- and MAML-style methods: In Figure 4, the curve for MAML when has a clear cusp. This seems to occur at the same suboptimality (ie. -value) as the intersection of the FedAvg curve with the -axis. In other words, the behavior of MAML diverges substantially from FedAvg, but only after it reaches the same suboptimality as FedAvg for (which corresponds to one-shot averaging). We are unsure why the suboptimality of one-shot averaging corresponds to a cuspidal operating point of MAML, but this observation highlights significant nuance in the behavior of these methods.
Limitations and Discussion
Our convergence-accuracy framework and the resulting Pareto frontiers can be useful tools in understanding how algorithmic choices impact local update methods. The obvious limitation is that they only apply to quadratic models. While this is restrictive, we show empirically in Appendix F that even for non-convex functions, the client learning rate governs a convergence-accuracy trade-off for FedAvg.
Our framework may also be useful in identifying important phenomena underlying LocalUpdate, even in non-quadratic settings. To demonstrate this, we show that many of the observations in Section 5 hold in non-convex settings. We train a CNN on the FEMNIST dataset (Caldas et al., 2018) using LocalUpdate where . We tune client and server learning rates. See Appendix E for full details. In Figure 6, we illustrate how server momentum and change convergence. Our results match the Pareto frontiers in Figures 2 and 3: Server momentum improves convergence, while has little to no effect, provided we tune learning rates.
This brief example illustrates that our framework can identify crucial facets of local update methods. While our framework may not capture all relevant details of such methods, we believe it greatly simplifies their analysis, comparison, and design. In the future, we hope to extend this framework to more general loss functions. Other important extensions include stochastic settings with partial client participation, as well as trade-offs between convergence and post-adaptation accuracy of local update methods.
Appendix A Relations between FedAvg, FedProx, and LocalUpdate
We focus on the following (simplified) version of FedAvg, otherwise known as Local SGD (Zinkevich et al., 2010; Stich, 2019): At each iteration , we sample some set of clients of size from the client population . Each client receives the server’s model , and applies steps of mini-batch SGD to its local model, resulting in an updated local model . The server receives these models from the sampled clients, and 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
A similar analysis holds for FedProx, but with the usage of a proximal term with parameter . In both cases, this is exactly LocalUpdate (see Algorithms 1 and 2) with , , and where ServerOpt is gradient descent. However, by instead using a server learning rate of that is allowed 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”, 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 implied by Theorem 1, for we can only guarantee that the surrogate loss has the same critical points as the true loss by setting . More generally, we see that the effective learning rate used in the model-delta approach is the product . This can result in conflations between the effect of changing the server learning rate and changing the client learning rate . By disentangling these, we can better understand differences in the impact of these parameters on the underlying optimization dynamics.
Appendix B Omitted Proofs
B.1 Proofs of Theorems 1 and 2
Suppose we perform iterations of SGD on with learning rate . That is, starting at we generate a sequence of independent random vectors and corresponding SGD iterates satisfying, for ,
We then have the following lemma regarding the .
By (18), the law of total expectation, and the independence of the ,
Recall that for , and we define the matrix by
We can now prove Theorem 1. For convenience, we restate the theorem here.
Fix , , and . For , define
Note that , where is as in (3). Straightforward manipulation of (24) implies
Note that the stochastic gradients computed in Algorithm 2 are therefore independent stochastic gradients of . By applying Lemma 7 and noting that the constant term does not impact these stochastic gradients, we have that for ,
Expanding and using the fact that in Algorithm 2, , we have
This last step follows from (3). Taking a sum and using the linearity of expectation,
A similar analysis using Lemma 7 can be used to derive Theorem 2, which we also restate.
Fix , and for convenience of notation, let . Thus, for we have
Since for all , we have
The result follows from applying (23) for . ∎
B.2 Proofs of Lemmas 1, 2, 3, 4, and 6
These lemmas will follow from a spectral analysis of . We defer the proof of Lemma 5 to Appendix C due to its more elaborate nature. We first state a general result about eigenvalues of expected values of matrices.
We also compute the spectrum of and .
For each eigenvalue of , has an eigenvalue
and has an eigenvalue
both with the same multiplicity as .
Let be an eigenvector of with eigenvalue . Then is an eigenvector of with eigenvalue , implying (29) is an eigenvalue of with eigenvector . Similarly, we note that is an eigenvector of with eigenvalue , implying (30) is an eigenvalue of with eigenvector . The statement about multiplicities follows directly. ∎
Lemma 9 can be used in a straightforward manner to prove Lemma 1, which we restate and prove below.
By Assumption 1, we have that for every eigenvalue of , . Therefore, for any such ,
Since the are nonnegative and not all zero by assumption, we see that (29) is a sum of nonnegative terms, at least one of which must be positive. Therefore, all eigenvalues of are positive, and is therefore symmetric and positive definite.
By Assumption 1, is also symmetric and positive definite, hence the product is symmetric positive definite. However, since
Applying Lemma 8, we conclude the proof. ∎
By Lemma 10, it suffices to derive upper and lower bounds on the eigenvalues of . Fix . By Lemma 9, we see that the eigenvalues of are exactly of the form , where is an eigenvalue of . By Assumption 1, each such satisfies where .
Fix , and define . Since , basic properties of geometric sums imply that for ,
For a given , let . Simple but tedious computations show that if we take a derivative with respect to , we have
Note that since by assumption, for . Therefore, for , so any eigenvalue of must satisfy
We use a similar proof as that of Lemma 3 to prove Lemma 4, which we restate and prove below.
By Lemma 10, it suffices to bound the eigenvalues of . Fix . By Lemma 9, we see that the eigenvalues of are of the form where is an eigenvalue of . Note that by Assumption 1, any such satisfies .
Fix , and define . Let . Straightforward computations show
Since , we in particular have so for . Since , we also have the term for . Thus, for . Thus, any eigenvalue of must satisfy
Finally, we are now equipped to prove Lemma 6, which we restate here for posterity.
This will follow almost immediately from Lemmas 3, 4, and 9. First, consider the case . Then by Lemma 9 and Assumption 1, we see that
Here we used the fact that and Lemma 3. An almost identical proof gives the analogous result for . ∎
B.3 Tightness of Lemmas 3 and 4
In fact, Lemmas 3 and 4 are tight. Fix any and . Let be supported on a single client , and let this client’s dataset be supported on a single example where
Note that by (3), we then have . In fact, we will show that in this case, the bounds on the condition numbers given in Lemma 3 and 4 are tight. By direct computation,
If then similar reasoning to the proof of Lemma 3 implies that the condition number satisfies
By analogous reasoning to the proof of Lemma 4, if , we have
Appendix C Proof of Lemma 5
In order to prove the results in this section, we will use the following straightforward lemma regarding the structure of .
To prove Lemma 5, we will reduce it to a statement about mean absolute deviations of bounded random variables. We define the mean absolute deviation of a random variable below.
To derive our results, we bound the mean absolute deviation of bounded random variables. Our result is inspired by the bound by Bhatia and Davis (2000) on the variance of bounded random variables.
Moreover, this holds with equality iff is supported on .
Suppose takes on values with probabilities . We will first show that there is a random variable supported on such that .
Without loss of generality, suppose . Define
First note that . Simple analysis also shows
By iterating this procedure, (which is guaranteed to terminate after at most iterations), we obtain some random variable supported on such that , with equality if and only if is already supported on . Suppose takes on with probabilities . Straightforward calculation shows
Using the fact that for any real , (with equality iff ) we arrive at the following corollary, analogous to Popoviciu’s inequality on variances (Popoviciu, 1935).
If is a discrete random variable on , then
with equality iff takes on the values and , each with probability .
We can now prove a stronger version of Lemma 5 when . For simplicity, we assume is a discrete distribution on some finite , though the analysis can be generalized to arbitrary distributions.
Let , . Then for ,
Since each , we also have .By Lemma 11,
Maximizing the right-hand side for , we get
C.2 Tightness of Lemma 12
Let . Let . Then by (6), we have
Therefore, in this setting, . By applying L’Hopital’s rule, we find
Let be the distribution that selects with probability and with probability . Then
For any fixed , straightforward but tedious applications of L’Hopital’s rule yields the fact that
By (36) and (37), we see that by selecting sufficiently large and sufficiently close to 1, we can ensure that
C.3 General Case
The matrix-weighted mean of with respect to is given by
When , and the , this gives the standard mean of a discrete random variable taking values with probabilities . When the context is clear, we will simply denote this by . We first prove a simple lemma regarding the Loewner ordering and matrix-weighted means.
Suppose that for all , . Then .
Since the are commuting positive definite matrices and is positive definite, is similar to the matrix
By assumption, . Therefore,
An analogous argument shows that . By basic properties of matrix similarity, we therefore find . ∎
We can use this matrix-weighted mean to define a normalized, matrix-weighted version of the mean absolute deviation.
The normalized matrix-weighted discrepancy of with respect to is given by
where is the operator norm.
We will prove an analog of Theorem 5 for this normalized matrix-weighted discrepancy.
To prove this, we will require a straightforward lemma regarding eigenvalues of symmetric positive definite matrices.
Let be symmetric positive definite matrices and let . Then
By basic properties of the Loewner ordering,
Since is similar to , we have
We will proceed in a similar manner to the proof of Theorem 5. We will first show that we can always find a set of symmetric positive definite matrices such that:
For all , .
For all , commute.
For , we have .
By iterating this procedure, we can replace with matrices where the are all in the set . It will then suffice to show that satisfies the desired bound, which we do by a somewhat direct computation, though one that is made much easier due to the fact that are diagonal.
We now proceed in detail. Define the matrices
Note that since commute, and are products of symmetric, positive definite, commuting matrices. They are therefore symmetric, positive definite, commuting matrices as well. One can easily verify that
We will use to denote the matrices in (41). Note that by Lemma 13, we know that . We will show that this replacement of by and does not decrease the mean absolute deviation. By (40) and (41),
Note that since and the are positive definite, and are positive semi-definite matrices. Simple algebraic manipulation implies that . Since are positive definite matrices, we have
We therefore exhibit exactly the matrices satisfying properties (1)-(6) described above. By iterating this procedure, we obtain positive definite, symmetric matrices such that commute, each is equal to or , and such that
By consolidating that are equal, we can assume without loss of generality that we have matrices with associated symmetric positive definite matrices . Let , and let . We then have
Let . Then by direct computation,
After some straightforward but tedious algebraic manipulation, we find
Since are symmetric positive definite matrices, we have
Letting denote respectively, and noting that we therefore have , we have
Suppose and is the discrete distribution on with associated probabilities . For brevity, we will let denote . We will let and . By Lemma 11, we have
Here we used Lemma 9, which in particular shows that since , all the are positive definite. Moreover, Lemma 9 shows that share the same eigenvectors, and therefore commute with one another. Hence, the commute with the . Moreover, by Assumption 1, the are positive definite symmetric matrices, and by Assumption 2, . Applying Theorem 6, we have
The last inequality holds from simple algebraic manipulation.
Appendix D Additional Pareto Frontiers
To generate , we generate by sampling its entries independently from . We then set , where are the unique scalars such that
We plot the resulting Pareto frontiers for varying and fixed in Figure 7. We see that as increases with respect to , the discrepancy between the MAML curves and the FedAvg curves grows. In particular, for small , we see that the MAML curve recovers most of the FedAvg curve before diverging, while for large , the two diverge almost immediately. Again, we see that when , there is some noise in , which seems to approach some limiting behavior for .
We perform a similar experiment, but where we fix and vary ServerOpt over gradient descent with no momentum, with Nesterov momentum, and with heavy-ball momentum. The results are given in Figure 8. While the differences are not huge, we see that momentum helps convergence in all cases, FedAvg or MAML. Moreover, we see an interesting phenomenon where the type of momentum changes the concavity of the MAML Pareto frontier for . As we add momentum, the region to the right of the MAML curve becomes more convex, becoming more rounded for heavy-ball momentum than for Nesterov momentum.
D.2 Proximal MAML-style Pareto Frontiers
In Figure 9 we plot the analog of the Pareto frontiers in Figure 3, but for MAML-style algorithms where . We see a similar, though more subdued, version of the behavior in Figure 3. That is, adding a proximal term simply alters how much of the Pareto frontier is traversed; it does not change the fundamental shape. Note that here we only used satisfying , as required by Theorem 3. In particular, the only restriction on the shape of the curve seems to be coming from the fact that larger reduces the set of satisfying .
To see the effects of when , we use the same simulated approach as in Figure 5. We do this for varying in Figure 10. We see that while increasing shrinks the space of the Pareto curve of FedAvg, it does not seem to change the MAML curves by a meaningful amount.
Appendix E Experimental Setup
We use three datasets: the federated extended MNIST dataset (FEMNIST) (Caldas et al., 2018), CIFAR-100 (Krizhevsky and Hinton, 2009), and Shakespeare (Caldas et al., 2018). The first two are image datasets, and the third is a language dataset. All datasets are publicly available. We specifically use the versions available in TensorFlow Federated (Ingerman and Ostrowski, 2019), which gives a federated structure to all three. We keep the client partitioning when training, and create a test dataset by taking a union over all test client datasets. Statistics on the number of clients and examples in each dataset are given in Table 3.
The FEMNIST dataset consists of images hand-written alphanumeric characters. There are 62 total alphanumeric characters represented in the dataset. The images are partitioned among clients according to their author. The dataset has natural heterogeneity stemming from the writing style of each person. We train a convolutional network on the dataset (the same one used by Reddi et al. (2020)). The network has two convolutional layers. Each convolutional layer uses kernels, max pooling, and then dropout with probability . The model has a final dense softmax output layer.
CIFAR-100
The CIFAR-100 dataset is a computer vision dataset consisting of images with 100 possible labels. While this dataset does not have a natural partition among clients, a federated version was created by Reddi et al. (2020) using hierarchical latent Dirichlet allocation to enforce moderate amounts of heterogeneity among clients. 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.
Shakespeare
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.
E.2 Implementation and Hyperparameters
We implement LocalUpdate in TensorFlow Federated (Ingerman and Ostrowski, 2019). We use LocalUpdate with and client learning rate . In all experiments, ServerOpt is gradient descent with server learning rate , with either no momentum, Nesterov momentum, or heavy-ball momentum. We sample clients per round. We sample without replacement within a given round, and with replacement across rounds. In order to derive fair comparisons between different hyperparameter settings, we use a random seed to fix which clients are sampled at each round. We use a batch size of for FEMNIST and CIFAR-100, and for Shakespeare.
E.3 Details of Figure 6
For posterity’s sake, we re-plot Figure 6 in Figure 11. To generate these plots, we perform two distinct experiments. In the first experiment (Figures 6 and 11, left), we fix and vary ServerOpt. Specifically, we let ServerOpt be gradient descent with no momentum (gradient), gradient descent with Nesterov momentum (nesterov), and gradient descent with heavy-ball momentum (momentum). When ServerOpt uses Nesterov or heavy-ball momentum, we use a momentum parameter of . In the second experiment (Figures 6 and 11, right), we fix ServerOpt to be gradient descent with no momentum, and vary the proximal strength . In both cases, we fix , and tune over the range
We select the values of attaining the best average test accuracy over the last 100 rounds.
Appendix F Additional Experiments
We wish to showcase the convergence-accuracy trade-off discussed in Section 4 in non-convex settings. We train LocalUpdate with , , and let ServerOpt be gradient descent with learning rate . First, we fix and vary over
We plot the training loss over time on all three datasets in Figure 12, omitting results that diverge due to being too large.
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 described in Sections 4 and 5, 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. This reflects the idea presented in Section 5 that hyperparameters should be chosen according to the desired convergence-accuracy trade-off. In communication-limited settings, we should use larger (or ), while in cases where we can run many communication rounds, we should use smaller (or ).
In short, we see clear evidence that the choice of client learning rate leads to a trade-off between convergence and accuracy. However, as shown in Lemmas 3 and 4, the condition number of the surrogate loss changes depending on parameters such as . To derive asymptotically optimal rates for strongly convex functions (such as the ones in Table 2), one must generally set the learning rate according to the condition number. Thus, we repeat the experiments in Figure 12, but where we tune the server learning rate instead of fixing it. This helps account for how the optimization dynamics can change as a function of the client learning rate . We vary the server learning rate over
and select that leads to the smallest average training loss over the last 100 rounds. The result is given in Figure 13. Again, we see similar behavior, but see that when the server learning rate is tuned, larger client learning rates may do much better initially. This reflects the fact that in Table 2, the best convergence rates can only be obtained by setting parameters of ServerOpt correctly.
This points to another benefit of the Pareto frontiers proposed in Section 5. Many comparisons of different algorithms, especially empirical ones, can miss good hyperparameter settings. This is heightened by the fact that many FL and ML methods have hyperparameters for both client and server optimizers, comprehensive tuning extremely difficult. This may lead to unfair comparisons between methods. By contrast, the Pareto frontiers showcase convergence-accuracy trade-offs when the hyperparameters of ServerOpt are selected in an “optimal” way, helping derive fair comparisons between methods.