Overcoming Forgetting in Federated Learning on Non-IID Data

Neta Shoham, Tomer Avidor, Aviv Keren, Nadav Israel, Daniel Benditkis, Liron Mor-Yosef, Itai Zeitak

Introduction

Recent years have seen the advent of smart devices and sensors gathering data at the edge and being able to act on that data. The desire to keep data private and other considerations have led the machine learning community to study algorithms for distributed training that do not require sending the data out of the edge devices. Edge devices most often have low networking availability and capacity, which could prohibit training through standard SGD. The Federated Averaging (FedAvg) algorithm of McMahan et al. lets the devices train on their local data for several epochs (using local SGD) before sending the trained model to a central server. The server then aggregates the models and sends the aggregated model back to the devices. This is done iteratively until convergence is achieved.

Federated Learning poses three challenges that make it different from traditional distributed learning. The first one is the number of computing stations, which can be in the hundreds of millions. In order to cope with this, it is common practice to select only a subset of devices at every training iteration . For simplicity of presentation, we will ignore this method, for which our suggested algorithm can be easily adapted. The second is much slower communication compared to the inter cluster communication found in data centers. The third difference, on which we focus in this work, is the highly non i.i.d. manner in which the data may be distributed among the devices.

In some real-life cases, Federated Learning has shown robustness to non i.i.d. distribution . There are also recent theoretical results proving the convergence of Federated Learning algorithms on non i.i.d. data. It is evident, however, that even in very simple scenarios, Federated Learning on non i.i.d. distributions has trouble achieving good results (in terms of accuracy and the number of communication rounds, as compared to the i.i.d. case) .

There is a deep parallel between the Federated Learning problem and another fundamental machine learning problem called Lifelong Learning (and the related Multi-Task Learning). In Lifelong Learning, the challenge is to learn task AA, and continue on to learn task BB using the same model, but without "forgetting", without severely hurting the performance on, task AA; or in general, learning tasks A1,A2…A_{1},A_{2}\dots in sequence without forgetting previously-learnt tasks for which samples are not presented anymore. Besides learning tasks serially rather than in parallel, in Lifelong Learning each task is thus seen only once, whereas in Federated Learning there is no such limitation. But these differences aside, the paradigms share a common main challenge - how to learn a task without disturbing different ones learnt on the same model.

It is not surprising, then, that similar approaches are being applied to solve the Federated Learning and the Lifelong Learning problems. One such example is data distillation, in which representative data samples are shared between tasks . However, Federated Learning is frequently used in order to achieve data privacy, which would be broken by sending a piece of data from one device to another, or from one device to a central point. We therefore seek for some other type of information to be shared between the tasks.

The answer to what kind of information to use may be found in Kirkpatrick et al. . In this work, the authors present a new algorithm for Lifelong Learning - Elastic Weight Consolidation (EWC). EWC aims to prevent catastrophic forgetting when moving from learning task AA to learning task BB. The idea is to identify the coordinates in the network parameters θ\theta that are the most informative for task AA, and then, while task BB is being learned, penalize the learner for changing these parameters. The basic assumption is that deep neural networks are over-parameterized enough, so that there are good chances of finding an optimal solution θB∗\theta^{*}_{B} to task BB in the neighborhood of previously learned θA∗\theta^{*}_{A}.

In order to control the stiffness of θ\theta per coordinate while learning task BB, the authors suggest to use the diagonal of the Fisher information matrix IA∗=IA(θA∗)\mathcal{I}^{*}_{A}=\mathcal{I}_{A}(\theta_{A}^{*}) to selectively penalize parts of the parameters vector θ\theta that are getting too far from θA∗\theta^{*}_{A}. This is done using the following objective

The formal justification they provide for (1) is Bayesian: Let DAD_{A} and DBD_{B} be independent datasets used for tasks AA and BB. We have that

It is also well known that under some regularity conditions, the information matrix approximates the Hessian HLH_{L} of L(θ)L(\theta), at θ=θ∗\theta=\theta^{*} . By this we get a non Bayesian interpretation of (1),

where L(θ)=LB(θ)+LA(θ)L(\theta)=L_{B}(\theta)+L_{A}(\theta) is exactly the loss we want to minimize. In general, one can learn a sequence of tasks A1…ATA_{1}\dots A_{T}. In section 3 we rely on the above interpretation as a second order approximation in order to construct an algorithm for Federated Learning. We will further show how to implement such an algorithm in a way that preserves the privacy benefits of the standard FedAvg algorithm.

Related Work

There are only a handful of works that directly try to cope with the challenge of Federated Learning with non i.i.d. distribution. One approach is to just give up the periodical averaging, and reduce the communication by sparsification and quantization of the updates sent to the central point after each local mini batch . In Zhao et al. it was shown that by sharing only a small portion of the data between different nodes, one can achieve a great improvement in model accuracy. However, sharing data is not acceptable in many Federated Learning scenarios.

Perhaps the closest work to ours is Sahu et al. , where the authors present their FedProx algorithm, which, like our algorithm, also uses parameter stiffness. However, unlike our algorithm, in FedProx the penalty term is isotropic, 12μ∥θ−θt∥\frac{1}{2}\mu\|\theta-\theta_{t}\|. DANE augments FedProx by adding a gradient correction term −(∇Li(θt−1)−η∇L(θt−1))Tθ-(\nabla L_{i}(\theta_{t-1})-\eta\nabla L(\theta_{t-1}))^{T}\theta to accelerate convergence, but is not robust to non i.i.d. data . AIDE improves the ability of DANE to deal with non i.i.d. data. However, it does so by using an inexact version of DANE, through a limitation on the amount of local computations.

A recent work proves convergence of FedAvg for the non i.i.d. case. It also provides a theoretical explanation for a phenomenon known in practice, of performance degradation when the number of local iterations is too high. This is exactly the problem that we tackle in this work.

Federated Curvature

In this section we present our adaptation of the EWC algorithm to the Federated Learning scenario. We call it FedCurv (for Federated Curvature, motivated by (2)). We mark by S={1…N}S={\{1\dots N\}} the NN nodes, with the tasks’ local datasets {A1,…AN}{\{A_{1},\dots A_{N}\}}. We diverge from the FedAvg algorithm and in each round tt we use all the nodes in SS instead of randomly selecting a subset on them. (Our algorithm can easily be extended to select a subset.) At round tt each node s∈Ss\in S optimizes the following loss:

At first glance, maintaining all the historical data required by FedCurv might look cumbersome and expensive to store and transmit. It also looks like a sensitive information is passed between nodes. However by careful implementation we can avoid these potential drawbacks. We note that (3) can also be rearranged as

The central point needs only to maintain and transmit to the edge node two additional elements, besides θ\theta, of the same size as θ\theta,

Privacy

It should be noted that we only need to send local gradient-related aggregated information (aggregated per local data sample) from the devices to the central point. In terms of privacy, it is not significantly different from the classical FedAvg algorithm. The central point itself, like in FedAvg, needs only to keep globally aggregated information from the devices. We see no reason why secure aggregation methods which were successfully applied to FedAvg could not be applied to FedCurv.

Further potential bandwidth reduction

The diagonal of the Fisher information has been used successfully for parameter pruning in neural networks . This gives us a straightforward way to save bandwidth by using sparse versions of \diag(I^)\diag(\hat{\mathcal{I}}) \diag(I^)θ^\diag({\mathcal{\hat{I}}})\hat{\theta} and even Δθ\Delta\theta, as \diag(I^)\diag(\hat{\mathcal{I}}) provides a natural evaluation for the importance measure of the parameters of θ^\hat{\theta}. The sparse versions are achieved by keeping only a fraction 0<q≤10<q\leq 1 of indices that are related to the qq largest elements of the diagonal of the Fisher information matrix. We have not explored this idea in practice.

Experiments

We conducted our experiments on a group of 96 simulated devices. We divided the MNIST dataset into 96×296\times 2 blocks of homogeneous labels (discarding a small amount of data). We randomly assigned two blocks to each device. We used the CNN architecture from the MNIST PyTorch example .

We explored two factors: (1) Learning method - we considered three algorithms, FedAvg, FedProx, and FedCurv (our algorithm); (2) EE, the number of epochs in each round, which is of special interest in this work, as our algorithm is designed for large values of EE. CC, the fraction of devices that participate in each iteration, and BB, the local batch size, were kept fixed at C=1.0,B=256C=1.0,B=256. For all the experiments, we have also used a constant learning rate of η=0.01\eta=0.01.

FedProx’s μ\mu and FedCurv’s λ\lambda values were chosen in the following way: We looked for values that reached 90% test-accuracy in the smallest number of rounds. We did it by searching on a multiplicative grid using a factor of 10 and then a factor of 2 in order to ensure a minimum. Table 1 shows the number of rounds required in order to achieve 95%, 90% and 85% test-accuracy with these chosen parameters. We see that for E=50E=50, FedCurv achieved 90% test-accuracy three times as fast as the vanila FedAvg algorithm. FedProx also reached 90% faster than FedAvg. However, while our algorithm achieved 95% twice as fast as FedAvg, FedProx achieved it two times slower. For E=10E=10, the improvement of both FedCurv and FedProx is less significant, with FedCurv still outperforming FedProx and FedAvg.

In Figure 2 and Figure 2 we can see that both FedProx and FedCurv are doing well at the beginning of the training process. However, while FedCurv provides enough flexibility with θ\theta that allows for reaching high accuracy at the end of the process, the stiffness of the parameters in FedProx comes at the expense of accuracy. FedCurv gives more significant improvements for higher values of EE (as does FedProx), as expected by the theory.

Conclusion

This work has provided a novel approach to the problem of Federated Learning on non i.i.d. data. It built on a solution from Lifelong Learning, which uses the diagonal of the Fisher information matrix in order to protect the parameters that are important to each task. The adaptation required modifying that sequential solution (from Lifelong Learning) into a parallel form (of Federated Learning), which a priori involves excessive sharing of data. We showed that this can be done efficiently, without substantially increasing bandwidth usage and compromising privacy. As our experiments have demonstrated, our FedCurv algorithm guards the parameters important to each task, improving convergence.

References