Robust Aggregation for Federated Learning

Krishna Pillutla, Sham M. Kakade, Zaid Harchaoui

Introduction

Federated learning is a key paradigm for machine learning and analytics on mobile, wearable and edge devices over wireless networks of 5G and beyond as well as edge networks and the internet of things. The paradigm has found widespread applications ranging from mobile apps deployed on millions of devices , to sensitive healthcare applications .

In federated learning, a number of devices with privacy-sensitive data collaboratively optimize a machine learning model under the orchestration of a central server, while keeping the data fully decentralized and private. Recent work has looked beyond supervised learning to domains such as data analytics but also semi-, self- and un-supervised learning, transfer learning, meta learning, and reinforcement learning .

We study a question relevant in all these areas: robustness to corrupted updates. Federated learning relies on aggregation of updates contributed by participating devices, where the aggregation is privacy-preserving. Sensitivity to corrupted updates, caused either by adversaries intending to attack the system or due to failures in low-cost hardware, is a vulnerability of the usual approach. The standard arithmetic mean aggregation in federated learning is not robust to corruptions, in the sense that even a single corrupted update in a round is sufficient to degrade the global model for all devices. In one dimension, the median is an attractive aggregate for its robustness to outliers. We adopt this approach to federated learning by considering a classical multidimensional generalization of the median, known variously as the geometric or spatial or L1L_{1} median .

Our robust approach preserves the privacy of the device updates by iteratively invoking the secure multi-party computation primitives used in typical non-robust federated learning . A device’s updates are information theoretically protected in that they are computationally indistinguishable from random noise and the sensitivity of the final aggregate to the contribution of each device is bounded. Our approach is scalable, since the underlying secure aggregation algorithms are implemented in production systems across millions of mobile users across the planet . The approach is communication-efficient, requiring a modest 1-3×\times the communication cost of the non-robust setting to compute the non-linear aggregate in a privacy-preserving manner.

The main take-away message of this work is: {quoting}[indentfirst=false,vskip=0.3em,leftmargin=1em,rightmargin=1em] Federated learning can be made robust to corrupted updates by replacing the weighted arithmetic mean aggregation with an approximate geometric median at 1-3 times the communication cost. To this end, we make the following concrete contributions.

Robust Aggregation: We design a novel robust aggregation oracle based on the classical geometric median. We analyze the convergence of the resulting federated learning algorithm, RFA, for least-squares estimation and show that the proposed method is robust to update corruption in up to half the devices in federated learning with bounded heterogeneity. We also describe an extension of the framework to handle arbitrary heterogeneity via personalization.

Algorithmic Implementation: We show how to implement this robust aggregation oracle in a practical and privacy-preserving manner. This relies on an alternating minimization algorithm which empirically exhibits rapid convergence. This algorithm can be interpreted as a numerically stable version of the classical algorithm of Weiszfeld , thus shedding new light on it.

Numerical Simulations: We demonstrate the effectiveness of our framework for data corruption and parameter update corruption, on federated learning tasks from computer vision and natural language processing, with linear models as well as convolutional and recurrent neural networks. In particular, our results show that the proposed RFA algorithm (i) outperforms the standard FedAvg , in high corruption and (ii) nearly matches the performance of the FedAvg in low corruption, both at 1-3 times the communication cost. Moreover, the proposed algorithm is agnostic to the actual level of corruption in the problem instance.

We open source an implementation of the proposed approach in TensorFlow Federated ; cf. Appendix B for a template implementation. The Python code and scripts used to reproduce experimental results are publicly available online .

Section 2 describes related work, and Section 3 describes the problem formulation and tradeoffs of robustness. Section 4 proposes a robust aggregation oracle and presents a convergence analysis of the resulting robust federated learning algorithm. Finally, Section 5 gives comprehensive numerical simulations demonstrating the robustness of the proposed federated learning algorithm compared to standard baselines.

Related Work

Federated Learning was introduced in as a distributed optimization approach to handle on-device machine learning, with secure multi-party averaging algorithms given in . Extensions were proposed in ; see also the recent surveys . We address robustness to corrupted updates, which is broadly applicable in these settings.

Distributed optimization has a long history . Recent work includes primal-dual frameworks and variants suited to decentralized , and asynchronous settings. From the lens of learning in networks , federated learning comprises a star network where agents (i.e., devices) with private data are connected to a server with no data, which orchestrates the cooperative learning. Further, for privacy, model updates from individual agents cannot be shared directly, but must be aggregated securely.

Robust estimation was pioneered by Huber . Robust median-of-means were introduced in , with follow ups in . Robust mean estimation, in particular, received much attention . Robust estimation in networks was considered in . These works consider the statistics of robust estimation in the i.i.d. case, while we focus on distributed optimization with privacy preservation.

Byzantine robustness, resilience to arbitrary behavior of some devices , was studied in distributed optimization with gradient aggregation . Byzantine robustness of federated learning is a priori not possible without additional assumptions because the secure multi-party computation protocols require faithful participation of the devices. Thus, we consider a more nuanced and less adversarial corruption model where devices participate faithfully in the aggregation loop; see Section 3 for practical examples. Further, it is unclear how to securely implement the nonlinear aggregation algorithms of these works. Lastly, the use of, e.g., secure enclaves in conjunction with our approach could guarantee Byzantine robustness in federated learning. We aggregate model parameters in a robust manner, which is more suited to the federated setting. We note that also aggregate model parameters rather than gradients by framing the problem in terms of consensus optimization. However, their algorithm requires devices to be always available and participate in multiple rounds, which is not practical in the federated setting .

Weiszfeld’s algorithm to compute the geometric median, has received much attention . The Weiszfeld algorithm is also known to exhibit asymptotic linear convergence . However, unlike these variants, ours is numerically stable. A theoretical proposal of a near-linear time algorithm for the geometric median was recently explored in .

Frameworks to guarantee privacy of user data include differential privacy and homomorphic encryption . These directions are orthogonal to ours, and could be used in conjunction. See for a broader discussion.

Problem Setup: Federated Learning with Corruptions

We begin this section by recalling the setup of federated learning (without corruption) and the standard FedAvg algorithm in Section 3.1. We then formally setup our corruption model and discuss the trade-offs introduced by requiring robustness to corrupted updates in Section 3.2.

Federated learning consists of nn client devices which collaboratively train a machine learning model under the orchestration of a central server or a fusion center . The data is local to the client devices while the job of the server is to orchestrate the training.

Federated learning aims to find a model w⋆w^{\star} that minimizes the average objective across all the devices,

where device ii is weighted by αi>0\alpha_{i}>0. In practice, the weight αi\alpha_{i} is chosen proportional to the amount of data on device ii. For instance, in an empirical risk minimization setting, each DiD_{i} is the uniform distribution over a finite set {zi,1,⋯ ,zi,Ni}\{z_{i,1},\cdots,z_{i,N_{i}}\} of size NiN_{i}. It is common practice to choose αi=Ni/N\alpha_{i}=N_{i}/N where N=∑i=1nNiN=\sum_{i=1}^{n}N_{i} so that the objective F(w)=(1/N)∑i=1n∑j=1Nif(w;zi,j)F(w)=(1/N)\sum_{i=1}^{n}\sum_{j=1}^{N_{i}}f(w;z_{i,j}) is simply the unweighted average over all samples from all nn devices.

Typical federated learning algorithms run in synchronized rounds of communication between the server and the devices with some local computation on the devices based on their local data, and aggregation of these updates to update the server model. The de facto standard training algorithm is FedAvg , which runs as follows.

The server samples a set StS_{t} of mm clients from [n][n] and broadcasts the current model w(t)w^{(t)} to these clients.

Staring from wi,0(t)=w(t)w_{i,0}^{(t)}=w^{(t)}, each client i∈Sti\in S_{t} makes τ\tau local gradient or stochastic gradient descent steps for k=0,⋯ ,τ−1k=0,\cdots,\tau-1 with a learning rate γ\gamma:

Each device i∈Sti\in S_{t} sends to the server a vector wi(t+1)w_{i}^{(t+1)} which is simply the final iterate, i.e., wi(t+1)=wi,τ(t)w_{i}^{(t+1)}=w_{i,\tau}^{(t)}. The server updates its global model using the weighted average

The federated learning algorithm, and in particular, the choice of aggregation, impacts the following three factors : communication efficiency, privacy, and robustness.

Besides the computation cost, the communication cost is an important parameter in distributed optimization. While communication is relatively fast in the datacenter, that is not the case of federated learning. The repeated exchange of massive models between the server and client devices over resource-limited wireless networks makes communication over the network more of a bottleneck in federated learning than local computation on the devices. Therefore, training algorithms should be able to trade-off more local computation for lower communication, similar to step (b) of FedAvg above. While the exact benefits (or lack thereof) of local steps is an active area of research, local steps have been found empirically to reduce the amount of communication required for a moderately accurate solution .

Accordingly, we set aside the local computation cost for a first order approximation, and compare algorithms in terms of their total communication cost . Since typical federated learning algorithms proceed in synchronized rounds of communication, we measure the complexity of the algorithms in terms of the number of communication rounds.

While the privacy-sensitive data z∼Diz\sim D_{i} is kept local to the device, the model updates wi(t+1)w_{i}^{(t+1)} might also leak privacy. To add a further layer of privacy protection, the server is not allowed to inspect individual updates wi(t+1)w_{i}^{(t+1)} in the aggregation step (c); it can only access the aggregate w(t+1)w^{(t+1)}.

As a result, we get the correct average (up to discretization) while not revealing any further information about a wiw_{i} or βi\beta_{i} to the server or other devices, beyond what can be inferred from the average. Hence, no further information about the underlying data distribution DiD_{i} is revealed either. In this work, we assume for simplicity that the secure average oracle returns the exact update, i.e., we ignore the effects of discretization on the integer ring and modular wraparound. This assumption is reasonable for a large enough value of MM.

We would like a federated learning algorithm to be robust to corrupted updates contributed by malicious devices or hardware/software failures. FedAvg uses an arithmetic mean to aggregate the device updates in (3), which is known to not be robust . This can be made precise by the notion of a breakdown point , which is the smallest fraction of the points which need to be changed to cause the aggregate to take on arbitrary values. The breakdown point of the mean is 0, since only one point needs to changed to arbitrarily change the aggregate . This means in federated learning that a single corrupted update, either due to an adversarial attack or a failure, can arbitrarily change the resulting aggregate in each round. We will give examples of adversarial corruptions in Section 3.2.

In the rest of this work, we aim to address the lack of robustness of FedAvg. A popular robust aggregation of scalars is the median rather than the mean. We investigate a multidimensional analogue of the median, while respecting the other two factors: communication efficiency and privacy. While the non-robust mean aggregation can be computed with secure multi-party computation via the secure average oracle, it is unclear if a robust aggregate can also satisfy this requirement. We discuss this as well as other tradeoffs involving robustness in the next section.

2 Corruption Model and Trade-offs of Robustness

This encompasses situations where the corrupted devices are individually or collectively trying to “attack” the global model, that is, reduce its predictive power over uncorrupted data. We define the corruption level ρ\rho as the total fraction of the weight of the corrupted devices:

Since the corrupted devices can only harm the global model through the updates they contribute in the aggregation step, we aim to robustify the aggregation in federated learning. However, it turns out that robustness is not directly compatible with the two other desiderata of federated learning, namely communication efficiency and privacy.

We first argue that any federated learning algorithm can only have two out of the three of robustness, communication and privacy under the existing techniques of secure multi-party computation. The standard approach of FedAvg is communication-efficient and privacy-preserving but not robust, as we discussed earlier. In fact, any aggregation scheme A(w1,⋯ ,wm)A(w_{1},\cdots,w_{m}) which is a linear function of w1,⋯ ,wmw_{1},\cdots,w_{m} is similarly non-robust. Therefore, any robust aggregate AA must be a non-linear function of the vectors it aggregates.

The approach of sending the updates to the server at a communication of O(md)O(md) and utilizing one of the many robust aggregates studied in the literature [e.g. 23, 88, 5] has robustness and communication efficiency but not privacy. If we try to make it privacy-preserving, however, we lose communication efficiency. Indeed, the secure multi-party computation primitives based on secret sharing, upon which privacy-preservation is built, are communication efficient only for linear functions of the inputs . The additional O(mlog⁡m)O(m\log m) overhead of secure averaging for linear functions becomes Ω(mdlog⁡m)\Omega(md\log m) for general non-linear functions required for robustness; this makes it impractical for large-scale systems . Therefore, one cannot have both communication efficiency and privacy preservation along with robustness.

In this work, we strike a compromise between robustness, communication and privacy. We will approximate a non-linear robust aggregate as an iterative secure aggregate, i.e., as a sequence of weighted averages, computed with a secure average oracle with weights being adaptively updated.

βi(r)\beta_{i}^{(r)} depends only on v(r)v^{(r)} and wiw_{i},

v(r+1)=∑i=1mβi(r)wi/∑i=1mβi(r)v^{(r+1)}=\sum_{i=1}^{m}\beta_{i}^{(r)}w_{i}/\sum_{i=1}^{m}\beta_{i}^{(r)}, and,

Further, the iterative secure aggregate is said to be ss-privacy preserving for some s∈(0,1)s\in(0,1) if

βi(r)/∑j=1mβj(r)≤s\beta_{i}^{(r)}/\sum_{j=1}^{m}\beta_{j}^{(r)}\leq s for all i∈[m]i\in[m] and r∈[R]r\in[R].

If we have an iterative secure aggregate with RR communication rounds which is also robust, we gain robustness at a RR-fold increase in communication cost. Condition (iv) ensures privacy preservation because it reveals only weighted averages with weights at most ss, so a user’s update is only available after being mixed with those from a large cohort of devices.

Heterogeneity is a key property of federated learning. The distribution DiD_{i} of device ii can be quite different from the distribution DjD_{j} of some other device jj, reflecting the heterogeneous data generated by a diverse set of users.

Next, we consider some examples of update corruption — see for a comprehensive treatment. Corrupted updates could be non-adversarial in nature, such as sensor malfunctions or hardware bugs in unreliable and heterogeneous devices (e.g., mobile phones) which are outside the control of the orchestrating server. On the other hand, we could also have adversarial corruptions of the following types:

Update Poisoning: The corrupted devices can send an arbitrary vector to the server for aggregation, as described by (4) in its full generality. This setting subsumes all previous examples as special cases.

The corruption model in (4) precludes the Byzantine setting [e.g., 46, Sec. 5.1], which refers to the worst-case model where a corrupted client device i∈Ci\in\mathcal{C} can behave arbitrarily, such as for instance, changing the weights βi(r)\beta_{i}^{(r)} or the vector wiw_{i} between each of the rounds of the iterative secure aggregate, as defined in Definition 1. It is provably impossible to design a Byzantine-robust iterative secure aggregate in this sense. The examples listed above highlight the importance of robustness to the corruption model under consideration.

Table 1 compares the various corruptions in terms of the capability of an adversary required to induce the corruption.

Robust Aggregation and the RFA Algorithm

In this section, we design a robust aggregation oracle and analyze the convergence of the resulting federated algorithm.

The GM has an optimal breakdown point of 1/2 . That is, to get the geometric median to equal an arbitrary point, at least half the points (in total weight) must be modified. We assume that w1,⋯wmw_{1},\cdots w_{m} are non-collinear, which is reasonable in the federated setting. Then, gg admits a unique minimizer v⋆v^{\star}. Further, we assume ∑iαi=1\sum_{i}\alpha_{i}=1 w.l.o.g. One could apply the results to g~(v):=g(v)/∑i=1mαi\widetilde{g}(v):=g(v)/\sum_{i=1}^{m}\alpha_{i}.

The RFA algorithm is obtained by replacing the mean aggregation of FedAvg with this GM-based robust aggregation oracle – the full algorithm is given in Algorithm 1. Similar to FedAvg, RFA also trades-off some communication for local computation by running multiple local steps in line 6. The communication efficiency and privacy preservation of RFA follow from computing the GM as an iterative secure aggregate, which we turn to next. Note that RFA is agnostic to the actual level of corruption in the problem and the aggregation is robust regardless of the convexity of the local objectives FiF_{i}.

While the GM is a natural robust aggregation oracle, the key challenge in the federated setting is to implement it as an iterative secure aggregate. Our approach, given in Algorithm 2, iteratively computes a new weight βi(r)∝1/∥v(r)−wi∥\beta_{i}^{(r)}\propto 1/\|v^{(r)}-w_{i}\|, up to a tolerance ν>0\nu>0, whose role is to prevent division by zero. This endows the algorithm with greater stability. We call it the smoothed Weiszfeld algorithm as it is a variation of Weiszfeld’s classical algorithm . The smoothed Weiszfeld algorithm satisfies the following convergence guarantee, proved in Appendix C.

The iterate v(R)v^{(R)} of Algorithm 2 with input v(0)∈conv⁡{w1,⋯ ,wm}v^{(0)}\in\operatorname*{conv}\{w_{1},\cdots,w_{m}\} and ν>0\nu>0 satisfies

where v⋆=arg min⁡gv^{\star}=\operatorname*{arg\,min}g and ν‾=min⁡r∈[R],i∈[m]ν∨∥v(r−1)−wi∥≥ν\overline{\nu}=\min_{r\in[R],i\in[m]}\nu\lor\|v^{(r-1)}-w_{i}\|\geq\nu. Furthermore, if 0<ν≤min⁡i=1,⋯ ,m∥v⋆−wi∥0<\nu\leq\min_{i=1,\cdots,m}\|v^{\star}-w_{i}\|, then it holds that g(v(R))−g(v⋆)≤2∥v(0)−v⋆∥2/ν‾R .g(v^{(R)})-g(v^{\star})\leq{2\|v^{(0)}-v^{\star}\|^{2}}/{\overline{\nu}R}\,.

For a ϵ\epsilon-approximate GM, we set ν=O(ϵ)\nu=O(\epsilon) to get a O(1/ϵ2)O(1/\epsilon^{2}) rate. However, if the GM v⋆v^{\star} is not too close to any wiw_{i}, then the same algorithm automatically enjoys a faster O(1/ϵ)O(1/\epsilon) rate. The algorithm enjoys plausibly an even faster convergence rate locally, and we leave this for future work.

Instead of minimizing g(v)g(v) directly using the equality g(v)=inf⁡η>0G(v,η)g(v)=\inf_{\eta>0}G(v,\eta), we impose the constraint ηi≥ν\eta_{i}\geq\nu instead to avoid division by small numbers. The following alternating minimization leads to Algorithm 2:

Numerically, we find in Figure 1 that Algorithm 2 is rapidly convergent, giving a high quality solution in 3 iterations. This ensures that the approximate GM as an iterative secure aggregate provides robustness at a modest 3×\times increase in communication cost over regular mean aggregation in FedAvg.

While we can compute the geometric median as an iterate secure aggregate, privacy preservation also requires that the effective weights βi(r)/∑jβj(r)\beta_{i}^{(r)}/\sum_{j}\beta_{j}^{(r)} are bounded away from 1 for each ii. We show this holds for mm large.

Since v(r)∈conv⁡{w1,⋯ ,wm}v^{(r)}\in\operatorname*{conv}\{w_{1},\cdots,w_{m}\}, we have νˉ≤∥v(r)−wi∥≤B\bar{\nu}\leq\|v^{(r)}-w_{i}\|\leq B. Hence, αi/B≤βi(r)≤αi/νˉ\alpha_{i}/B\leq\beta_{i}^{(r)}\leq\alpha_{i}/\bar{\nu} for each ii and rr and the proof follows. ∎

1 Convergence Analysis of RFA

We now present a convergence analysis of RFA under two simplifying assumptions. First, we focus on least-squares fitting of additive models, as it allows us to leverage sharp analyses of SGD and focus on the effect of the aggregation. Second, we assume w.l.o.g. that each device is weighted by αi=1/n\alpha_{i}=1/n to avoid technicalities of random sums ∑i∈Stαi\sum_{i\in S_{t}}\alpha_{i}. This assumption can be lifted with standard reductions; see Remark 5.

We now analyze RFA where the local SGD updates are equipped with “tail-averaging” so that wi(t+1)=(2/τ)∑k=τ/2τwi,k(t)w_{i}^{(t+1)}=(2/\tau)\sum_{k=\tau/2}^{\tau}w_{i,k}^{(t)} is averaged over the latter half of the trajectory of iterates instead of line 9 of Algorithm 1. We show that this variant of RFA converges up to the dissimilarity level Ω=ΩXΩY∣X\Omega=\Omega_{X}\Omega_{Y|X} when the corruption level ρ<1/2\rho<1/2.

Consider FF defined in (7) and suppose the corruption level satisfies ρ<1/2\rho<1/2. Consider Algorithm 1 run for TT outer iterations with a learning rate γ=1/(2R2)\gamma=1/(2R^{2}), and the local updates are run for τt\tau_{t} steps in outer iteration tt with tail averaging. Fix δ>0\delta>0 and θ∈(ρ,1/2)\theta\in(\rho,1/2), and set the number of devices per iteration, mm as

Define Cθ:=(1−2θ)−2C_{\theta}:=(1-2\theta)^{-2}, w⋆=arg min⁡Fw^{\star}=\operatorname*{arg\,min}F, F⋆=F(w∗)F^{\star}=F(w^{*}), κ:=R2/μ\kappa:=R^{2}/\mu and Δ0:=∥w(0)−w⋆∥2\Delta_{0}:=\|w^{(0)}-w^{\star}\|^{2}. Let τ≥4κlog⁡(128Cθκ)\tau\geq 4\kappa\log\left(128C_{\theta}\kappa\right). We have that the event E=⋂t=0T−1{∣St∩C∣≤θm}\mathcal{E}=\bigcap_{t=0}^{T-1}\{|S_{t}\cap\mathcal{C}|\leq\theta m\} holds with probability at least 1−δ1-\delta. Further, if τt=2tτ\tau_{t}=2^{t}\tau for each iteration tt, then the output w(T)w^{(T)} of Algorithm 1 satisfies,

where CC is a universal constant. If τt=τ\tau_{t}=\tau instead, then, the noise term above reads dσ2/μτd\sigma^{2}/{\mu\tau}.

Theorem 4 shows near-linear convergence O(T/2T)O(T/2^{T}) up to two error terms in the case that ρ\rho is bounded away from 1/21/2 (so that θ\theta and CθC_{\theta} can be taken to be constants). The increasing local computation τt=2tτ\tau_{t}=2^{t}\tau required by this rate is feasible since local computation is assumed to be cheaper than communication.

The first error term is ϵ2/m2\epsilon^{2}/m^{2} due to approximation ϵ\epsilon in the GM, which can be made arbitrarily small by increasing the number mm of devices sampled per round. The second error term Ω2\Omega^{2} is due to heterogeneity. Indeed, exact convergence as T→∞T\to\infty is not possible in the presence of corruption: lower bounds for robust mean estimation [e.g. 22, Theorem 2.2] imply that ∥w(T)−w⋆∥2≥Cρ2ΩY∣X2\|w^{(T)}-w^{\star}\|^{2}\geq C\rho^{2}\Omega_{Y|X}^{2} w.p. at least 1/21/2. Consistent with our theory, we find in real heterogeneous datasets in Section 5 that RFA can lead to marginally worse performance than FedAvg in the corruption-free regime (ρ=0\rho=0). Finally, while we focus on the setting of least squares, our results can be extended to the general convex case.

We use the following convergence result of SGD [43, Theorem 1], [44, Corollary 2].

Consider the local updates on an uncorrupted device i∈St∖Ci\in S_{t}\setminus\mathcal{C}, starting from w(t)w^{(t)}. Theorem 6 gives, upon using τt≥τ≥4κlog⁡(128Cθκ)\tau_{t}\geq\tau\geq 4\kappa\log(128C_{\theta}\kappa),

Note that w⋆=(1/n)∑j=1nH−1Hjwj⋆w^{\star}=(1/n)\sum_{j=1}^{n}H^{-1}H_{j}w_{j}^{\star}, so that

Using ∥a+b∥2≤2∥a∥2+2∥b∥2\|a+b\|^{2}\leq 2\|a\|^{2}+2\|b\|^{2}, we get,

We now apply the robustness property of the GM ([59, Thm. 2.2] or [86, Lem. 3]) to get,

where Γ=2Cθ(ϵ2/m2+16Ω2)\Gamma=2C_{\theta}(\epsilon^{2}/m^{2}+16\Omega^{2}). Taking an expectation conditioned on E\mathcal{E} and unrolling this inequality gives

When τt=2tτ\tau_{t}=2^{t}\tau, the series sums to 2−(T−1)T/τ2^{-(T-1)}T/\tau, while for τt=τ\tau_{t}=\tau, the series is upper bounded by 2/τ2/\tau. ∎

We now consider RFA in connection with the three factors mentioned in Section 3.1.

Communication Efficiency: Similar to FedAvg, RFA performs multiple local updates for each aggregation round, to save on the total communication. However, owing to the trade-off between communication, privacy and robustness, RFA requires a modest 3×\times more communication for robustness per aggregation. In the next section, we present a heuristic to reduce this communication cost to one secure average oracle call per aggregation.

Privacy Preservation: Algorithm 2 computes the aggregation as an iterative secure aggregate. This means that the server only learns the intermediate parameters after being averaged over all the devices, with effective weights bounded away from 11 (Proposition 3). The noisy parameter vectors sent by individual devices are uniformly uninformative in information theoretic sense with the use of secure multi-party computation.

Robustness: The geometric median has a breakdown point of 1/2 [59, Theorem 2.2], which is the highest possible [59, Theorem 2.1]. In the federated learning context, this means that convergence is still guaranteed by Theorem 4 when up to half the points in terms of total weight are corrupted. RFA is resistant to both data or update poisoning, while being privacy preserving. On the other hand, FedAvg has a breakdown point of 0, where a single corruption in each round can cause the model to become arbitrarily bad.

2 Extensions to RFA

We now discuss two extensions to RFA to reduce the communication cost (without sacrificing privacy) and better accommodate statistical heterogeneity in the data with model personalization.

Recall that RFA results in a 3-5×\times increase in the communication cost over FedAvg. Here, we give a heuristic variant of RFA in an extremely communication-constrained setting, where it is infeasible to run multiple iterations of Algorithm 2. We simply run Algorithm 2 with v(0)=0v^{(0)}=0 and a communication budget of R=1R=1; see Algorithm 3 for details. We find in Section 5.3 that one-step RFA retains most of the robustness of RFA.

We now show RFA can be extended to better handle heterogeneity in the devices with the use of personalization. The key idea is that predictions are made on device ii by summing the shared parameters ww maintained by the server with personalized parameters U={u1,⋯ ,un}U=\{u_{1},\cdots,u_{n}\} maintained individually on-device. In particular, the optimization problem we are interested in solving is

We outline the algorithm in Algorithm 4. We train the shared and personalized parameters on each other’s residuals, following the residual learning scheme of . Each selected device first updates its personalized parameters uiu_{i} while keeping the shared parameters ww fixed. Next, the updates to the shared parameter are computed on the residual of the personalized parameters. The updates to the shared parameter are aggregated with the geometric median, identical to RFA. Experiments in Section 5.3 show that personalization is effective in combating heterogeneity.

Numerical Simulations

We now conduct simulations to compare RFA with other federated learning algorithms. The simulations were run using TensorFlow and the data was preprocessed using LEAF . We first describe the experimental setup in Section 5.1, then study the robustness and convergence of RFA in Section 5.2. We study the effect of the extensions of RFA in Section 5.3. The full details from this section and more simulation results are given in Appendix D. The code and scripts to reproduce these experiments can be found online .

We consider three machine learning tasks. The datasets are described in Table 2. As described in Section 3.1, we take the weight αi\alpha_{i} of device ii to be proportional to the number of datapoints NiN_{i} on the device.

Character-Level Language Modeling: We learn a character-level language model over the Complete Works of Shakespeare . We formulate it as a multiclass classification problem, where the input xx is a window of 20 characters, the output yy is the next (i.e., 21st) character. Each device is a role from a play (e.g., Brutus from The Tragedy of Julius Caesar). We use a long-short term memory model (LSTM) together with the multinomial logistic loss. The performance is evaluated with the classification accuracy of next-character prediction.

Sentiment Analysis: We use the Sent140 dataset where the input xx is a tweet and the output y=±1y=\pm 1 is its sentiment. Each device is a distinct Twitter user. We use a linear model using average of the GloVe embeddings of the words of the tweet. It is trained with the binary logistic loss and evaluated with the classification accuracy.

We consider the following corruption models for corrupted devices C\mathcal{C}, cf. Section 3.2:

Update poisoning with Gaussian corruption: Each corrupted device i∈Ci\in\mathcal{C} returns wi(t+1)=wi,τ(t)+ζi(t)w_{i}^{(t+1)}=w_{i,\tau}^{(t)}+\zeta_{i}^{(t)}, where ζi(t)∼N(0,σ2I)\zeta_{i}^{(t)}\sim\mathcal{N}(0,\sigma^{2}I), where σ2\sigma^{2} is the variance across the components of wi,τ(t)−w(t)w_{i,\tau}^{(t)}-w^{(t)}. Model updates wi(t)−w(t)w_{i}^{(t)}-w^{(t)} are aggregated, not the models wi(t)w_{i}^{(t)} directly .

Update poisoning with omniscient corruption: The parameters wi(t+1)w_{i}^{(t+1)} returned by devices i∈Ci\in\mathcal{C} are modified so that the weighted arithmetic mean ∑i∈Stαiwi(t+1)\sum_{i\in S_{t}}\alpha_{i}w_{i}^{(t+1)} over the selected devices StS_{t} is set to −∑i∈Stαiwi,τ(t)-\sum_{i\in S_{t}}\alpha_{i}w_{i,\tau}^{(t)}, the negative of what it would to have been without the corruption. This is designed to hurt the weighted arithmetic mean aggregation.

The hyperparameters are chosen similar to the defaults of . A learning rate schedule was tuned on a validation set for FedAvg with no corruption. The same schedule was used for RFA. The aggregation in RFA is implemented using the smoothed Weiszfeld algorithm with a budget of R=3R=3 calls to the secure average oracle, thanks to its rapid empirical convergence (cf. Figure 1), and ν=10−6\nu=10^{-6} for numerical stability. Each simulation was repeated 5 times and the shaded area denotes the minimum and maximum over these runs. Appendix D gives details on hyperparameter, and a sensitivity analysis of the Weiszfeld communication budget.

2 Robustness and Convergence of RFA

First, we compare the robustness of RFA as opposed to vanilla FedAvg to different types of corruption across different datasets in Figure 2. We make the following observations.

RFA gives improved robustness to linear models with data corruption. For instance, consider the EMNIST linear model at ρ=1/4\rho=1/4. RFA achieves 52.8% accuracy, over 10% better than FedAvg at 41.2%.

RFA performs similarly to FedAvg in deep nets with data corruption. RFA and FedAvg are within one standard deviations of each other for the Shakespeare LSTM model, and nearly equal for the EMNIST ConvNet model. We note that the behavior of the training of a neural network when the data is corrupted is not well-understood in general [e.g., 90].

RFA gives improved robustness to omniscient corruptions for all models. For the omniscient corruption, the test accuracy of the FedAvg is close to 0% for the EMNIST linear model and ConvNet, while RFA still achieves over 40% at ρ=1/4\rho=1/4 for the former and well over 60% for the latter. A similar trend holds for the Shakespeare LSTM model.

RFA almost matches FedAvg in the absence of corruption. Recall from Section 3.2 that robustness comes at the cost of heterogeneity; this is also reflected in the theory of Section 4. Empirically, we find that the performance hit of RFA due to heterogeneity is quite small: 1.4% for the EMNIST linear model (64.3% vs. 62.9%), under 0.4% for the Shakespeare LSTM, and 0.3% for Sent140 (65.0% vs. 64.7%). Further, we demonstrate in Appendix D.5 that, consistent with the theory, this gap completely vanishes in the i.i.d. case.

Summary: robustness of RFA. Overall, we find that RFA is no worse than FedAvg in the presence of corruption and is often better, while being almost as good in the absence of corruption. Furthermore, RFA degrades more gracefully as the corruption level increases.

RFA requires only 3×3\times the communication of FedAvg. Next, we plot in Figure 4 the performance versus the number of rounds of communication as measured by the number of calls to the secure average oracle. We note that in the low corruption regime of ρ=0\rho=0 or ρ=10−2\rho=10^{-2} under data corruption, RFA requires 3×3\times the number of calls to the secure average oracle to reach the same performance. However, it matches the performance of FedAvg when measured in terms of the number of outer iterations, with the additional communication cost coming from multiple Weiszfeld iterations for computation of the average.

RFA exhibits more stable convergence under corruption. We also see from Figure 4 (ρ=1/4\rho=1/4, Data) that the variability of accuracy across random runs, denoted here by the shaded region, is much smaller for RFA. Indeed, by being robust to the corrupted updates sent by random sampling of corrupted clients, RFA exhibits a more stable convergence across iterations.

3 Extensions of RFA

We now study the proposed extensions: one-step RFA and personalization.

One-step RFA gives most of the robustness with no extra communication. From Figure 5, we observe that for one-step RFA is quite close in performance to RFA across different levels of corruption for both data corruption on an EMNIST linear model and omniscient corruption on an EMNIST ConvNet. For instance, in the former, one-step RFA gets 51.4% in accuracy, which is 10% better than FedAvg while being almost as good as full RFA (52.8%) at ρ=0.25\rho=0.25. Moreover, for the latter, we find that one-step RFA (67.9%) actually achieves higher test accuracy than full RFA (63.0%) at ρ=0.25\rho=0.25.

Personalization helps RFA offset effects of heterogeneity. Figure 6 plots the effect of RFA with personalization. First, we observe that personalization leads to an improvement with no corruption for both FedAvg and RFA. For the EMNIST linear model, we get 70.1% and 69.9% respectively from 64.3% and 62.9%. Second, we observe that RFA exhibits greater robustness to corruption with personalization. At ρ=1/4\rho=1/4 with the EMNIST linear model, RFA with personalization gives 66.4% (a reduction of 3.4%) while no personalization gives 52.8% (a reduction of 10.1%). The results for Sent140 are similar, with the exception that FedAvg with personalization is nearly identical to RFA with personalization.

Conclusion

We presented a robust aggregation approach, based on the geometric median and the smoothed Weiszfeld algorithm to efficiently compute it, to make federated learning more robust to settings where a fraction of the devices may be sending corrupted updates to the orchestrating server. The robust aggregation oracle preserves the privacy of participating devices, operating with calls to secure multi-party computation primitives enjoying privacy preservation theoretical guarantees. RFA is available in several variants, including a fast one with a single step of robust aggregation and a one adjusting to heterogeneity with on-device personalization. All variants are readily scalable while preserving privacy, building off secure multi-party computation primitives already used at planetary scale. The theoretical analysis of RFA with personalization is an interesting venue for future work. The further analysis of robustness under heterogeneity is also an interesting venue for future work.

The authors would like to thank Zachary Garrett, Peter Kairouz, Jakub Konečný, Brendan McMahan, Krzysztof Ostrowski and Keith Rush for fruitful discussions, as well as help with the implementation of RFA on Tensorflow Federated. This work was first presented at the Workshop on Federated Learning and Analytics in June 2019. This work was supported by NSF CCF-1740551, NSF CCF-1703574, NSF DMS-1839371, the Washington Research Foundation for innovation in Data-intensive Discovery, the program “Learning in Machines and Brains”, faculty research awards, and a JP Morgan PhD Fellowship.

References

Table of Contents

Appendix A Table of Notation

We summarize the notation used throughout the paper in Table 3.

Appendix B Template Implementation of RFA in TensorFlow Federated

We provide here a template implementation of RFA in Tensorflow Federated. The open source software is publicly available .

Appendix C The Smoothed Weiszfeld Algorithm: Convergence Analysis

In this section, we prove the rate of the smoothed Weiszfeld algorithm in Proposition 2. We start by a setup, prove a number of interesting properties, and finally prove Proposition 2 in Section C.5.

The points w1,⋯ ,wiw_{1},\cdots,w_{i} are not collinear.

The geometric median is defined as any minimizer of

Under Assumption 7, gg is known to have a unique minimizer - we denote it by z⋆z^{\star}.

Given a smoothing parameter ν>0\nu>0, its smoothed variant gνg_{\nu} is

In case ν=0\nu=0, we define g0≡gg_{0}\equiv g. It is known that ∥⋅∥(ν)\|\cdot\|_{(\nu)} is (1/ν)(1/\nu)-smooth and that

Under Assumption 7, gνg_{\nu} has a unique minimizer as well, denoted by vν⋆v_{\nu}^{\star}. We call vν⋆v_{\nu}^{\star} as the ν\nu-smoothed geometric median.

We let BB denote the diameter of the convex hull of {w1,⋯ ,wm}\{w_{1},\cdots,w_{m}\}, i.e.,

We also assume that ν<B\nu<B, since for all ν≥B\nu\geq B, the function gνg_{\nu} is simply a quadratic for all z∈conv⁡{w1,⋯ ,wm}z\in\operatorname*{conv}\{w_{1},\cdots,w_{m}\}.

C.2 Weiszfeld’s Algorithm: Review

The Weiszfeld algorithm performs the iterations

where βi(r)=αi/∥v(r)−wi∥\beta_{i}^{(r)}={\alpha_{i}}/{\|v^{(r)}-w_{i}\|}. It was shown in [49, Thm. 3.4] that the sequence (v(r))t=0∞\left(v^{(r)}\right)_{t=0}^{\infty} converges to the minimizer of gg from (10), provided no iterate coincides with one of the wiw_{i}’s. We modify Weiszfeld’s algorithm to find the smoothed geometric median by considering

This is also stated in Algorithm 5. Since each iteration of Weiszfeld’s algorithm or its smoothed variant consists in taking a weighted average of the wiw_{i}’s, the time complexity is O(md)\mathcal{O}(md) floating point operations per iteration.

C.3 Derivation

We now derive Weiszfeld’s algorithm with smoothing as as an alternating minimization algorithm or as an iterative minimization of a majorizing objective.

Note firstly that GG is jointly convex in z,ηz,\eta over its domain.

The first claim shows how to recover gg and gνg_{\nu} from GG.

Consider g,gνg,g_{\nu} and GG defined in Equations (10), (11) and (17), and fix ν>0\nu>0. Then we have the following:

so that G(z,η)=∑i=1mαiGi(z,ηi)G(z,\eta)=\sum_{i=1}^{m}\alpha_{i}G_{i}(z,\eta_{i}).

Since ηi>0\eta_{i}>0, the arithmetic-geometric mean inequality implies that Gi(z,ηi)≥∥z−wi∥G_{i}(z,\eta_{i})\geq\|z-w_{i}\| for each ii. When ∥z−wi∥>0\|z-w_{i}\|>0, the inequality above holds with equality when ∥z−wi∥2/ηi=ηi\|z-w_{i}\|^{2}/{\eta_{i}}=\eta_{i}, or equivalently, ηi=∥z−wi∥\eta_{i}=\|z-w_{i}\|. On the other hand, when ∥z−wi∥=0\|z-w_{i}\|=0, let ηi→0\eta_{i}\to 0 to conclude that

For the second part, we note that if ∥z−wi∥≥ν\|z-w_{i}\|\geq\nu, then ηi=∥z−wi∥≥ν\eta_{i}=\|z-w_{i}\|\geq\nu minimizes Gi(z,ηi)G_{i}(z,\eta_{i}), so that min⁡ηi≥νGi(z,ηi)=∥z−wi∥\min_{\eta_{i}\geq\nu}G_{i}(z,\eta_{i})=\|z-w_{i}\|. On the other hand, when ∥z−wi∥<ν\|z-w_{i}\|<\nu, we note that Gi(z,⋅)G_{i}(z,\cdot) is minimized over [ν,∞)[\nu,\infty) at ηi=ν\eta_{i}=\nu, in which case we get Gi(z,η)=∥z−wi∥2/(2ν)+ν/2G_{i}(z,\eta)=\|z-w_{i}\|^{2}/(2\nu)+\nu/2. From (12), we conclude that

The proof is complete since G(z,η)=∑i=1mαiGi(z,ηi)G(z,\eta)=\sum_{i=1}^{m}\alpha_{i}G_{i}(z,\eta_{i}). ∎

8 now allows us to consider the following problem in lieu of minimizing gνg_{\nu} from (11).

Application of this method to Problem (20) yields the updates

These updates can be written in closed form as

This gives the smoothed Weiszfeld algorithm, as pointed out by the following claim.

Follows from plugging in the expression from ηi(r)\eta_{i}^{(r)} in the update for v(r+1)v^{(r+1)} in (22). ∎

We now instantiate the smoothed Weiszfeld algorithm as a majorization-minimization scheme. In particular, it is the iterative minimization of a first-order surrogate in the sense of .

where η(r)\eta^{(r)} is as defined in (21). The zz-step of (21) simply sets v(r+1)v^{(r+1)} to be the minimizer of gν(r)g_{\nu}^{(r)}.

We note the following properties of gν(r)g_{\nu}^{(r)}.

For gν(r)g_{\nu}^{(r)} defined in (23), the following properties hold:

Moreover g(r)g^{(r)} can also be written as

For Eq. (25), note that the inequality above is an equality at v(r)v^{(r)} by the definition of η(r)\eta^{(r)} from (21). To see (26), note that

Then, by the definition of η(r)\eta^{(r)} from (22), we get that

The obtain the expansion (27), we write out the Taylor expansion of the quadratic g(r)(z)g^{(r)}(z) around v(r)v^{(r)} to get

and complete the proof by plugging in (25) and (26). ∎

The next claim rewrites the smoothed Weiszfeld algorithm as gradient descent on gνg_{\nu}.

C.4 Properties of Iterates

The first claim reasons about the iterates v(r),η(r)v^{(r)},\eta^{(r)}.

Starting from any v(0)∈conv⁡{w1,⋯ ,wm}v^{(0)}\in\operatorname*{conv}\{w_{1},\cdots,w_{m}\}, the sequences (η(r))(\eta^{(r)}) and (v(r))(v^{(r)}) produced by Algorithm 5 satisfy

v(r)∈conv⁡{w1,⋯ ,wm}v^{(r)}\in\operatorname*{conv}\{w_{1},\cdots,w_{m}\} for all t≥0t\geq 0, and,

ν≤ηi(r)≤B \nu\leq\eta_{i}^{(r)}\leq B\, for all i=1,⋯ ,mi=1,\cdots,m, and t≥1t\geq 1,

where B=diam⁡(conv⁡{w1,⋯ ,wm})B=\operatorname{diam}(\operatorname*{conv}\{w_{1},\cdots,w_{m}\}). Furthermore, L(r)L^{(r)} defined in (28) satisfies 1/B≤L(r)≤1/ν1/B\leq L^{(r)}\leq 1/\nu for all t≥0t\geq 0.

The first part follows for t≥1t\geq 1 from the update (16), where Claim 9 shows the equivalence of (16) and (21). Then case of t=0t=0 is assumed. The second part follows from (22) and the first part. The bound on L(r)L^{(r)} follows from the second part since ∑i=1mαi=1\sum_{i=1}^{m}\alpha_{i}=1. ∎

The next result shows that it is a descent algorithm. Note that the non-increasing nature of the sequence (gν(v(r)))\left(g_{\nu}(v^{(r)})\right) also follows from the majorization-minimization viewpoint . Here, we show that this sequence is strictly decreasing. Recall that vν⋆v_{\nu}^{\star} is the unique minimizer of gνg_{\nu}.

The sequence (v(r))(v^{(r)}) produced by Algorithm 5 satisfies gν(v(r+1))<gν(v(r))g_{\nu}(v^{(r+1)})<g_{\nu}(v^{(r)}) unless v(r)=vν⋆v^{(r)}=v_{\nu}^{\star}.

The next lemma shows that ∥v(r)−z⋆∥\|v^{(r)}-z^{\star}\| is non-increasing. This property was shown in [11, Corollary 5.1] for the case of Weiszfeld algorithm without smoothing.

The sequence (v(r))(v^{(r)}) produced by Algorithm 5 satisfies for all t≥0t\geq 0,

Furthermore, if gν(v(r+1))≥gν(z⋆)g_{\nu}(v^{(r+1)})\geq g_{\nu}(z^{\star}), then it holds that

where L(r)L^{(r)} is defined in (28). Starting from the results of Claim 10, we observe for any zz that,

Plugging in z=vν⋆z=v_{\nu}^{\star}, the fact that gν(v(r+1))≥gν(z⋆)g_{\nu}(v^{(r+1)})\geq g_{\nu}(z^{\star}) implies that ∥v(r+1)−vν⋆∥2≤∥v(r)−vν⋆∥2\|v^{(r+1)}-v_{\nu}^{\star}\|^{2}\leq\|v^{(r)}-v_{\nu}^{\star}\|^{2}, since L(r)≥1/BL^{(r)}\geq 1/B is strictly positive. Likewise, for z=z⋆z=z^{\star}, the claim holds under the condition that gν(v(r+1))≥gν(z⋆)g_{\nu}(v^{(r+1)})\geq g_{\nu}(z^{\star}). ∎

C.5 Rate of Convergence

We are now ready to prove the global sublinear rate of convergence of Algorithm 5.

The iterate v(R)v^{(R)} produced by Algorithm 5 with input v(0)∈conv⁡{w1,⋯ ,wm}v^{(0)}\in\operatorname*{conv}\{w_{1},\cdots,w_{m}\} and ν>0\nu>0 satisfies

where L(s)=∑i=1mαi/ηi(s)L^{(s)}=\sum_{i=1}^{m}{\alpha_{i}}/{\eta_{i}^{(s)}} is defined in (28), and

With the descent and contraction properties of Lemma 13 and Lemma 14 respectively, the proof now follows the classical proof technique of gradient descent [e.g., 71, Theorem 2.1.13]. Starting from the results of Claim 10, we observe for any zz that,

For ease of notation, we let Δ~r:=gν(v(r))−gν(vν⋆)\widetilde{\Delta}_{r}:=g_{\nu}(v^{(r)})-g_{\nu}(v_{\nu}^{\star}). We assume now that Δ~r+1\widetilde{\Delta}_{r+1} is nonzero, and hence, so is Δ~r\widetilde{\Delta}_{r} (Lemma 13). If Δ~r+1\widetilde{\Delta}_{r+1} were zero, then the theorem would hold trivially at t+1t+1.

Now, from convexity of gνg_{\nu} and the Cauchy-Schwartz inequality, we get that

Now, we divide by Δ~rΔ~r+1\widetilde{\Delta}_{r}\widetilde{\Delta}_{r+1}, which is nonzero by assumption, and use Δ~r/Δ~r+1≥1\widetilde{\Delta}_{r}/\widetilde{\Delta}_{r+1}\geq 1 (Lemma 13) to get

This proves the first inequality to be proved. The second inequality follows from the definition in Eq. (28) since ∑i=1mαi=1\sum_{i=1}^{m}\alpha_{i}=1.

The proof follows along the same ideas as the previous proof. Define Δr:=gν(v(r))−gν(z⋆)\Delta_{r}:=g_{\nu}(v^{(r)})-g_{\nu}(z^{\star}). Suppose Δr>0\Delta_{r}>0. Then, we proceed as previously for any s<ts<t to note by convexity and Cauchy-Schwartz inequality that

Again, plugging this into (32), using that Δs/Δs+1≥1\Delta_{s}/\Delta_{s+1}\geq 1 and invoking Lemma 14 gives (since Δs>0\Delta_{s}>0)

Telescoping and taking the reciprocal gives

Using (13) completes the proof for the case that Δr>0\Delta_{r}>0. Note that if Δr≤0\Delta_{r}\leq 0, it holds that Δt′≤0\Delta_{t^{\prime}}\leq 0 for all t′>tt^{\prime}>t. In this case, gν(v(r))−gν(z⋆)≤0g_{\nu}(v^{(r)})-g_{\nu}(z^{\star})\leq 0. Again, (13) implies that g(v(r))−g(z⋆)≤ν/2g(v^{(r)})-g(z^{\star})\leq\nu/2, which is trivially upper bounded by the quantity stated in the theorem statement. This completes the proof. ∎

The geometric median z⋆z^{\star} does not coincide with any of w1,⋯ ,wmw_{1},\cdots,w_{m}. In other words,

[11, Lemma 8.1] show a lower bound on ν~\widetilde{\nu} in terms of α1,⋯ ,αm\alpha_{1},\cdots,\alpha_{m} and w1,⋯ ,wmw_{1},\cdots,w_{m}.

Now, we analyze the condition under which the z⋆=vν⋆z^{\star}=v_{\nu}^{\star}.

Under Assumption 16, we have that z⋆=vν⋆z^{\star}=v_{\nu}^{\star} for all ν≤ν~\nu\leq\widetilde{\nu}, where ν~\widetilde{\nu} is defined in (33).

In this case, we get a better rate on the non-smooth objective gg.

Consider the setting of Theorem 15 where Assumption 16 holds and ν≤ν~\nu\leq\widetilde{\nu}. Then, the iterate v(R)v^{(R)} produced by Algorithm 5 satisfies,

where ν^\widehat{\nu} is defined in Eq. (31).

This follows from Theorem 15’s bound on gν(v(R))−gν(vν⋆)g_{\nu}(v^{(R)})-g_{\nu}(v_{\nu}^{\star}) with the observations that g(v(R))≤\eqrefeq:weiszfeld:norm:smooth:boundgν(v(R))g(v^{(R)})\stackrel{{\scriptstyle\eqref{eq:weiszfeld:norm:smooth:bound}}}{{\leq}}g_{\nu}(v^{(R)}) and g(z⋆)=gν(z⋆)g(z^{\star})=g_{\nu}(z^{\star}) (see the proof of Lemma 18). ∎

The previous corollary obtains the same rate as [11, Theorem 8.2], up to constants upon using the bound on ν~\widetilde{\nu} given by [11, Lemma 8.1].

We also get as a corollary a bound on the performance of Weiszfeld’s original algorithm without smoothing, although it could be numerically unstable in practice. This bound depends on the actual iterates, so it is not informative about the performance of the algorithm a priori.

Consider the setting of Theorem 15. Under Assumption 16, suppose the sequence (v(r))(v^{(r)}) produced by Weiszfeld’s algorithm in Eq. (15) satisfies ∥v(r)−wi∥>0\|v^{(r)}-w_{i}\|>0 for all rr and ii, then it also satisfies

Under these conditions, note that the sequence (v(s))s=0t(v^{(s)})_{s=0}^{t} produced by the Weiszfeld algorithm without smoothing coincides with the sequence (vν(r)(s))s=0t(v_{\nu^{(r)}}^{(s)})_{s=0}^{t} produced by the smoothed Weiszfeld algorithm at level ν=ν(r)\nu=\nu^{(r)}. Now apply Corollary 19. ∎

C.6 Comparison to Previous Work

We compare the results proved in the preceding section to prior work on the subject.

The authors present multiple different variants of the Weiszfeld algorithm. For a particular choice of initialization, they can guarantee that a rate of the order of 1/ν~R1/\widetilde{\nu}R. It is not clear how this choice of initialization can be implemented using a secure average oracle since, if at all. This is because it requires the computation of all pairwise distances ∥wi−wi′∥\|w_{i}-w_{i^{\prime}}\|. Moreover, a naive implementation of their algorithm could be numerically unstable since it would involve division by small numbers. Guarding against division by small numbers would lead to the smoothed variant considered here. Note that our algorithmic design choices are driven by the federated learning setting.

The author studies general alternating minimization algorithms, including the Weiszfeld algorithm as a special case, with a different smoothing than the one considered here. While their algorithm does not suffer from numerical issues arising from division by small numbers, it always suffers a bias from smoothing. On the other hand, the smoothing considered here is more natural in that it reduces to Weiszfeld’s original algorithm when ∥v(r)−wi∥>ν\|v^{(r)}-w_{i}\|>\nu, i.e., when we are not at a risk of dividing by small numbers. Furthermore, the bound in Theorem 15 exhibits a better dependence on the initialization v(0)v^{(0)}.

Appendix D Numerical Simulations: Full Details

The section contains a full description of the experimental setup as well as additional results.

We start with the dataset and task description in Section D.1, hyperparameter choices in Section D.2, and evaluation methodology in Section D.3. We provide some extra numerical results in Section D.5.

We experiment with three tasks, (1) handwritten-letter recognition, (2) character-level language modeling, and, (3) sentiment analysis. As discussed in Section 3.1, we take the weight αi∝Ni\alpha_{i}\propto N_{i}, which is the number of data points available on device ii.

The first dataset is the EMNIST dataset for handwritten letter recognition.

Each inpt xx is a gray-scale image resized to 28×2828\times 28. Each output yy is categorical variable which takes 62 different values, one per class of letter (0-9, a-z, A-Z).

The task of handwritten letter recognition is cast as a multi-class classification problem with 62 classes.

The handwritten characters in the images are annotated by the writer of the character as well. We use a non-i.i.d. split of the data grouped by a writer of a given image. We discard devices with less than 100 total input-output pairs (both train and test), leaving a total of 3461 devices. Of these, we sample 1000 devices to use for our simulations, corresponding to about 30%30\% of the data. This selection held constant throughout the simulations. The number of training examples across these devices summarized in the following statistics: median 160, mean 202, standard deviation 77, maximum 418 and minimum 92. This preprocessing was performed using LEAF .

For the model φ\varphi, we consider two options: a linear model and a convolutional neural network.

Convolutional Neural Network (ConvNet): The ConvNet we consider contains two convolutional layers with max-pooling, followed by a fully connected hidden layer, and another fully connected (F.C.) layer with 62 outputs. When given an input image xx, the output of this network is assigned as the scores of each of the classes. Probabilities are assigned similar to the linear model with a softmax operation on the scores. The schema of network is given below:

The model is evaluated based on the classification accuracy on the test set.

D.1.2 Character-Level Language Modeling

The second task is to learn a character-level language model over the Complete Works of Shakespeare . The goal is to read a few characters and predict the next character which appears.

The dataset consists of text from the Complete Works of William Shakespeare as raw text.

We formulate the task as a multi-class classification problem with 53 classes (a-z, A-Z, other) as follows. At each point, we consider the previous H=20H=20 characters, and build x∈{0,1}H×53x\in\{0,1\}^{H\times 53} as a one-hot encoding of these HH characters. The goal is then try to predict the next character, which can belong to 53 classes. In this manner, a text with ll total characters gives ll input-output pairs.

We use a non-i.i.d. split of the data. Each role in a given play (e.g., Brutus from The Tragedy of Julius Caesar) is assigned as a separate device. All devices with less than 100 total examples are discarded, leaving 628 devices. The training set is assigned a random 90% of the input-output pairs, and the other rest are held out for testing. This distribution of training examples is extremely skewed, with the following statistics: median 1170, mean 3579, standard deviation 6367, maximum 70600 and minimum 90. This preprocessing was performed using LEAF .

We use a long-short term memory model (LSTM) with 128128 hidden units for this purpose. This is followed by a fully connected layer with 53 outputs, the output of which is used as the score for each character. As previously, probabilities are obtained using the softmax operation.

The model is evaluated based on the accuracy of next-character prediction on the test set.

D.1.3 Sentiment Analysis

The third task is analyze the sentiment of tweets as positive or negative.

Sent140 is a text dataset of 1,600,498 tweets produced by 660,120 Twitter accounts. Each tweet is represented by a character string with emojis redacted. Each tweet is labeled with a binary sentiment reaction (i.e., positive or negative), which is inferred based on the emojis in the original tweet.

The task is a binary classification problem, with the output being a positive or negative sentiment, while the input is the raw text of the tweet.

We use a non-i.i.d. split of the data. Each client device represents a Twitter user and contains tweets from this user. We discarded all clients containing less that 50 tweets, leaving only 877 clients. The training set is assigned a random 80% of the input-output pairs, and the other rest are held out for testing. This distribution of training examples across client devices is skewed, with the following statistics: median 55, mean 65.3, standard deviation 32.4, maximum 439 and minimum 40. This preprocessing was performed using LEAF .

We use the binary classification accuracy.

D.2 Methods, Hyperparameters and Variants

We first describe the corruption model, followed by various methods tested.

Since the goal of this work to test the robustness of federated learning models in the setting of high corruption, we artificially corrupt updates while controlling the level of corruption. We use the following corruption models.

This is an example of update poisoning. The data is not modified here but the update of a client device is directly replaced by a Gaussian random variable, with standard deviation σ\sigma equal to the standard deviation of the original update across its components. Note that we corrupt the update to the model parameters transmitted by the device, which is typically much smaller in norm than the model parameters themselves.

This is an example of update poisoning. The data is not modified here but the parameters of a device are directly modified. In particular, wi(t+1)w_{i}^{(t+1)} for i∈Ci\in\mathcal{C} is set to be

In other words, the weighted arithmetic mean of the model parameter is set to be the negative of what it would have other been without the corruption. This corruption model requires full knowledge of the data and server state, and is adversarial in nature.

Given a corruption level ρ\rho, the set of devices which return corrupted updates are selected as follows:

Sample device ii uniformly without replacement and add to C\mathcal{C}. Stop when ∑i∈Cαi\sum_{i\in C}\alpha_{i} just exceeds ρ\rho.

D.2.2 Methods

the RFA algorithm proposed here in Algorithm 1,

the minibatch stochastic gradient descent (SGD) algorithm.

D.2.3 Hyperparameters

The hyperparameters for each of these algorithms are detailed below.

The FedAvg algorithm requires the following hyperparameters.

Devices per round mm: We use 100100 for EMNIST and 5050 for both the Shakespeare and Sent140 datasets.

Batch Size and Number of Local Epochs: Instead of running τ\tau local updates, we run for nen_{e} local epochs following with a batch size of bb. For the EMNIST dataset, we use b=50,ne=5b=50,n_{e}=5, and for Shakespeare and Sent140, we use b=10,ne=1b=10,n_{e}=1.

Learning Rate (γt)(\gamma_{t}): We use a learning a learning rate scheme γt=γ0C⌊t/t0⌋\gamma_{t}=\gamma_{0}C^{\lfloor t/t_{0}\rfloor}, where γ0\gamma_{0} and CC were tuned using grid search on validation set (20% held out from the training set) for a fixed time horizon on the uncorrupted data. The values which gave the highest validation accuracy were used for all settings - both corrupted and uncorrupted. The time horizon used was 2000 iterations for the EMNIST linear model, 1000 iterations for the EMNIST ConvNet 200 iterations for Shakespeare LSTM.

Initial Iterate w(0)w^{(0)}: Each element of w(0)w^{(0)} is initialized to a uniform random variable whose range is determined according to TensorFlow’s “glorot_uniform_initializer”.

RFA’s hyperparameters, in addition to those of FedAvg, are:

Algorithm: We use the smoothed Weiszfeld algorithm, as discussed in Sec. 4.

Smoothing parameter ν\nu: Based on the interpretation that ν\nu guards against division by small numbers, we simply use ν=10−6\nu=10^{-6} throughout.

Robust Aggregation Stopping Criterion: The concerns the stopping criterion used to terminate the smoothed Weiszfeld algorithm. We use two criteria: an iteration budget and a relative improvement condition - we terminate if a given iteration budget has been extinguished, or if the relative improvement in objective value ∣gν(v(r))−gν(v(r+1))∣/gν(v(r))≤10−6|g_{\nu}(v^{(r)})-g_{\nu}(v^{(r+1)})|/g_{\nu}(v^{(r)})\leq 10^{-6} is small.

D.3 Evaluation Methodology and Other Details

We specify here the quantities appearing on the xx and yy axes on the plots, as well as other details.

As mentioned in Section 3, the goal of federated learning is to learn the model with as few rounds of communication as possible. Therefore, we evaluate various methods against the number of rounds of communication, which we measure via the number of calls to a secure average oracle.

Note that FedAvg and SGD require one call to the secure average oracle per outer iteration, while RFA could require several. Hence, we also evaluate performance against the number of outer iterations.

We are primarily interested in the test accuracy, which measures the performance on unseen data. We also plot the function value FF, which is the quantity our optimization algorithm aims to minimize. We call this the train loss.

In simulations with data corruption, while the training is performed on corrupted data, we evaluate train and test progress using the corruption-free data.

We use the package LEAF to simulate the federated learning setting. The models used are implemented in TensorFlow.

Each simulation was run in a simulation as a single process. The EMNIST linear model simulations were run on two workstations with 126GB of memory, with one equipped with Intel i9 processor running at 2.80GHz, and the other with Intel Xeon processors running at 2.40GHz. Simulations involving neural networks were run either on a 1080Ti or a Titan Xp GPU.

Each simulation is repeated 5 times with different random seeds, and the solid lines in the plots here represents the mean over these runs, while the shaded areas show the maximum and minimum values obtained in these runs.

D.4 Simulation Results: Convergence of The Smoothed Weiszfeld Algorithm

For each of these models, we freeze FedAvg at a certain iteration and experiment with different robust aggregation algorithms.

We find that the smoothed Weiszfeld algorithm enjoys a fast convergence behavior, converging exactly to the smoothed geometric median in a few passes. In fact, the smoothed Weiszfeld algorithm displays (local) linear convergence, as evidenced by the straight line in log scale. Further, we also maintain a strict iteration budget of 3 iterations. This choice is also justified in hindsight by the results of Figure 10.

Next, we visualize the weights assigned by the geometric median to the corrupted updates. Note that the smoothed geometric median w1,⋯ ,wmw_{1},\cdots,w_{m} is some convex combination ∑i=1mβiwi\sum_{i=1}^{m}\beta_{i}w_{i}. This weight βi\beta_{i} of wiw_{i} is a measure of the influence of wiw_{i} on the aggregate. We plot in Figure 8(b) the ratio βi/αi\beta_{i}/\alpha_{i} for each device ii, where αi\alpha_{i} is its weight in the arithmetic mean and βi\beta_{i} is obtained by running the smoothed Weiszfeld algorithm to convergence. We expect this ratio to be smaller for worse corruptions and ideally zero for obvious corruptions. We find that the smoothed geometric median does indeed assign lower weights to the corruptions, while only accessing the points via a secure average oracle.

D.5 Additional Simulation Results

Here, we plot the analogue of Figure 2 for the Sent140 dataset with data corruption in the setting where the dataset was split in an i.i.d. manner across devices. Recall that we had a small gap of 0.3% between the performance of RFA and FedAvg in the setting of no corruption. Consistent with the theory, this gap completely vanishes in the i.i.d. case, as shown in Figure 9.

We study the effect of the iteration budget of the smoothed Weiszfeld algorithm in RFA. in Figure 10. We observe that a low communication budget is faster in the regime of low corruption, while more iterations work better in the high corruption regime. We used a budget of 3 calls to the secure average oracle throughout to trade-off between these two scenarios.

Figure 11 plots the performance of RFA against the number mm of devices chosen per round. We observe the following: in the regime of low corruption, good performance is achieved by selecting 50 devices per round (5%), where as 10 devices per round (1%) is not enough. On the other hand, in high corruption regimes, we see the benefit of choosing more devices per round, as a few runs with 10 or 50 devices per round with omniscient corruption at 25% diverged. This is consistent with Theorem 4, which requires the number of devices per round to increase with the level of corruption (cf. Eq. (9)).

Figure 12 plots the performance of FedAvg and RFA versus the amount of local computation. We see that the performance is always within one standard deviation of each other irrespective of the amount of local computation. However, we also note that RFA with a single local epoch is obtains a slightly lower test accuracy in the no-corruption regime than using more local computation.