Distilled One-Shot Federated Learning

Yanlin Zhou, George Pu, Xiyao Ma, Xiaolin Li, Dapeng Wu

Introduction

Conventional supervised learning dictates that data be gathered into a central location where it can be used to train a model. However, this is intrusive and difficult, if data is spread across multiple devices or clients. For this reason, federated learning (FL) has garnered attention due to its ability to collectively train neural networks while keeping data private. The most popular FL algorithm is FedAvg . Each iteration, clients perform local training and forward the resulting model weights to a server. The server averages these to obtain a global model. Since learning processes happen at local level, neither the server nor other clients directly observe a client’s data.

Federated learning introduces distinct challenges not present in classical distributed machine learning . The main focus of this paper are expensive communication and statistical heterogeneity. Previous approaches try to learn faster when data is poorly distributed . They include modifying the training loss , using lifelong learning to prevent forgetting , and correcting local updates using control variates . These methods improve upon FedAvg, but can still take hundreds of communication rounds, while increasing the amount of information sent to the server per round.

Inspired by dataset distillation , we propose Distilled One-shot Federated Learning (DOSFL) (see Figure 1) to reduce the communication cost by up to 3 orders of magnitude. DOSFL requires only one round of communication between a server and its clients. Each client distills their data and uploads learned synthetic data, label and learning rate to the server, instead of transmitting bulky gradients or weights. Even large datasets containing thousands of examples can be compressed to only a few fabricated examples. The server then interleaves the clients’ distilled data together, using them to train a global model. To achieve good results even when client data is poorly distributed, we leverage soft labels and introduce two new techniques: soft reset and random masking.

Due to the reduction in rounds, we claim a communication reduction of up to 99.9% compared to FedAvg while achieving similar accuracy. DOSFL can also preserve between 93% to 99% of centralized training’s performance when data is independently and identically distributed (IID). Furthermore, under more challenging assumptions of low participation and asynchronous update with multiple rounds, DOSFL can achieve even better results and preserve higher centralized performance, i.e., almost 99% on IID Federated MNIST. Unlike dataset compression, fabricated inputs in dataset distillation (see Appendix B.1 and B.2)—along with the corresponding labels and learning rate—are generated by a specific model weight distribution and become useless after the model updates. We experimentally show that the final performance of a DOSFL is strongly dependent on the model parameters used to train on distilled data. Without knowing the the exact initial weights distributed by the server, an eavesdropper cannot reproduce the global model with leaked distilled data.

We believe DOSFL is part of a new paradigm in the field of FL. So far, nearly every FL algorithm communicates model weights or gradients. While effective, breaking from this pattern offers many benefits, such as low communication or private model architectures . We hope that DOSFL, along with related work, may inspire the machine learning community to explore possible techniques for weight-less and gradient-less FL.

Related Work

Since the introduction of FedAvg in 2016 , there has been an explosion of work directed towards the problem of statistical heterogeneity. When statistical hetergeniety is high, convergence for FedAvg slows and becomes unstable . The issue is that, the difference between the local losses and the global objective—their weighted sum—may be large. As such, minimizing a particular local loss does not ensure that the global loss is also minimized. This is problematic, even when the losses are convex and smooth . In applications where privacy loss can be tolerated, Zhao et al. demonstrate massive gains in performance by making as little as 5% of data public .

Numerous successors to FedAvg have been suggested. Server momentum introduces a global momentum parameter that improves convergence theoretically and experimentally . In Shoham et al., the loss is modified with Elastic Weight Consolidation to prevent forgetting as clients perform local training . SCAFFOLD uses control variates, namely the gradient of the global model, to address drifting among clients during local training . These schemes, while effective, at least double the per round communication cost. More efforts have not had this drawback. The work in proposes several federated version of optimizers, FedAdaGrad, FedAdam and FedYogi, to provide better convergence for non-IID data. The issue of objective inconsistency was discussed in and solved by the proposed FedNova, which averages the normalized local gradients.

While faster learning decreases the total number of communication rounds, strategies have been devised to explicitly reduce communication costs. FedPAQ quantizes local updates before transmission, with averaging happens at both client and server sides . Sparsifying the weights may perform better than FedAvg alone . Asynchronous model updates have also been explored, using adaptive weighted averaging to combat staleness, combined with a proximal loss, and updating deeper layer less frequently than shallower layers . SAFA takes a semi-synchronous approach; only up-to-date and deprecated clients synchronize with the server .

A few papers have made first steps towards one-shot FL. Guha et al. try different heuristics for selecting the client models which would form the best ensemble . By swapping weight averaging for ensembling, only one round of communication is necessary. Upper and lower bounds have been proven for one-shot distributed optimization, along with an order optimal algorithm . The extent to which these results apply to FL of neural networks is unknown as the local losses must be convex and drawn from the same probability distribution. Recently, knowledge distillation has been introduced to the field of FL to reduce the computation costs for edge devices . In addition, group knowledge transfer reduces communication costs by accelerating convergence.

2 Distillation

There is a wealth of literature studying dataset compression, while maintaining the most crucial features for training models. These methods include dataset pruning and core set construction , which keep the examples that are measured to be more useful for training and remove the rest. The drawback of drawing distilled images from the original dataset is that the level of compression achieved is much lower than that of dataset distillation, which is exempt from the requirement that distilled data be real . Dataset distillation was introduced by Wang et al., to compress a large dataset with thousands to millions of images down to only a few synthetic training images. The key idea is to use gradient descent to learn the features most helpful for rapidly training a neural network. Given some model parameters θ0\theta_{0}, dataset distillation minimizes the loss of adapted parameters θ1\theta_{1}, obtained by performing gradient descent on θ0\theta_{0} and the distilled data. This procedure resembles meta-learning, which performs task-specific adaption followed by a meta-update . With dataset distillation, 10 synthetic digits can train a neural network from 13% to 94% test accuracy in 3 iterations, near the 98%98\% test accuracy reached by training on MNIST.

Dataset distillation originally was limited to only image classification tasks, because the distilled labels were predetermined and fixed. Learnable or soft labels not only decrease the number of required labels, but also expand dataset distillation to language tasks such as sentiment classification . Soft labels have a long history, being proposed for model distillation by Hinton et al. and for k-nearest neighbors by El Gayar et al. . Using soft label dataset distillation, Sucholutsky et al. were able to train LeNet to 96%96\% accuracy with only 10 images . Examples of distilled data (i.e., text, grey image and RGB image) are shown in Figure 1.

Distilled One Shot Federated Learning

Suppose we have numbered clients 1,…,N1,\dots,N each with their own local models fθ1,…,fθNf_{\theta_{1}},\dots,f_{\theta_{N}} with parameters θ1,…,θN\theta_{1},\dots,\theta_{N} and loss functions L1,…,LNL_{1},\dots,L_{N}. Given some probability vector p=(p1,…,pN)p=(p_{1},\dots,p_{N}) (each 0≤pk≤10\leq p_{k}\leq 1 and ∑kpk=1\sum_{k}p_{k}=1), our goal is to find some parameters θ∗\theta^{*} that minimize the weighted sum L=∑k=1NpkLk(θ)L=\sum_{k=1}^{N}p_{k}L_{k}(\theta).

Our solution consists of 3 steps. These steps are summarized in Algorithm 1.

A central server randomly initializes model parameters θ0\theta_{0}. This can be distributed to the clients as a random seed.

Afterwards, minimize the loss of θ1\theta_{1} evaluated on a batch of real data B\mathcal{B}.

The clients upload the distilled data to the server. If Sd>1S_{d}>1, the server sorts the distilled data by index, e.g. {x1,x2,x3},{y1,y2,y3}\{x_{1},x_{2},x_{3}\},\{y_{1},y_{2},y_{3}\} from clients 1 and 2 become {x1,y1,x2,y2,x3,y3}\{x_{1},y_{1},x_{2},y_{2},x_{3},y_{3}\} where xj,yjx_{j},y_{j} are 3-tuples. The server then trains its own model on the combined sequence.

The last step can cause issues when the data is non-IID. Consider two clients 11 and 22 with distilled examples x1x_{1} and y1y_{1} respectively with Ed=1E_{d}=1. The server first trains θ0\theta_{0} on x1x_{1}, arriving at θ1\theta_{1}, which is then trained on y1y_{1}. But y1y_{1} has been distilled to train θ0\theta_{0}. To combat this interference, we introduce two new techniques for improving performance on non-IID data.

Soft resets sample the starting parameters, θ0\theta_{0} from a Gaussian distribution, around the server’s parameters N(θ0,σsr2)\mathcal{N}(\theta_{0},\sigma_{sr}^{2}). By sampling θ0\theta_{0} between distillation iterations, dataset distillation learns more robust examples capable of training any model with weights θ∼N(θ0,σsr2)\theta\sim\mathcal{N}(\theta_{0},\sigma_{sr}^{2}). This technique is based off of the “hard resets” introduced in , which completely re-initializes θ0\theta_{0}. Data distilled with “hard resets” can be used on any randomly initialized model, but cannot train models to the same level of accuracy as models trained on data distilled without resets.

Random masking randomly selects a fraction prmp_{rm} of the distilled data at each training iteration and replaces it with a random tensor. The random tensors randomly adjusts the model during training, while also reducing the amount of distilled data to actually train the starting parameters. After the training iteration, the original distilled data are restored. Now, sequences of distilled data can still train a model even when there is interference from other distilled steps. However, resetting and storing distilled data is compute and memory intensive, which slows down distillation.

Experiments

We evaluate DOSFL on several federated classification tasks. Because of this, cross entropy loss is used for all experiments. To train the distilled data, we use ADAM with a learning rate α\alpha that is halved every τ\tau epochs. We have α=0.01,τ=40\alpha=0.01,\tau=40, α=0.01,τ=10\alpha=0.01,\tau=10, and α=0.1,τ=30\alpha=0.1,\tau=30 for federated MNIST, IMDB, and TREC-6 respectively. These hyperparameters are mirrored from and have been found to be near-optimal. For federated Sent140, we replicated the hyperparmeters from federated IMDB. Clients distill the data for E=30E=30 epochs for image datasets and E=50E=50 epochs for text datasets with a batch size of B=512B=512. All experiments were run on an Nvidia M40 GPU and Intel Xeon E5-2695 CPU, taking 2-3 minutes per client for 100 client federated MNIST, <1<1 minute per client for 29 client federated IMDB and 100 client TREC-6, and 14-15 minutes for 100 client federated Sent140. We use the default train and test splits associated with each dataset.

The client-server architecture is simulated by partitioning a dataset into subsets, and then distilling these subsets. The server models have their weights Xavier initialized . The weights are then replicated across each client. Following the methodology of McMahan et al., IID partitions are created by randomly dividing the dataset into subsets . For non-IID partitions, we first sort the entire dataset by label and then divide it into NsNs shards of equal length. Starting from NN empty subsets, the shards are randomly assigned to the subsets until each has ss shards. As ss increases, the partition becomes more IID, with subsets more likely to contain examples from each class.

We first test DOSFL on 10 and 100 client Federated MNIST with different combinations of soft labels, soft resets, and random masking. All federated MNIST experiments use LeNet as the model architecture . Our distilled data are not single examples but batches with size Bd=10B_{d}=10. Within each batch, the labels are initialized to one of each class, e.g. one label is ‘1’, one label is ‘2’, etc. Using soft labels, BdB_{d} could be made much smaller without loss of performance. After the server model has trained on distilled data from the clients, its accuracy is measured on a test set. Each experiment is run 55 times, and the best result is reported in Table 1.

The distill steps are Sd=30S_{d}=30, the distill epochs are Ed=3E_{d}=3, and the initial distill learning rate is η0=0.02\eta_{0}=0.02. Of the proposed additions to dataset distillation, soft resets provide the largest jump in non-IID performance, followed by random masking and soft labels. The reset variance σsr2=0.2\sigma^{2}_{sr}=0.2 and masking probability prm=0.3p_{rm}=0.3 were chosen from the search spaces {0.1,0.2,…,0.6}\{0.1,0.2,\dots,0.6\} and {0.1,0.2,0.3,0.5,1}\{0.1,0.2,0.3,0.5,1\}. However, soft labels also boost accuracy when data is IID, where as the other two methods cause dips in the final accuracy. The distillation additions are not additive; even with all add-ons, non-IID DOSFL caps at ∼79%\sim 79\% test accuracy. Surprisingly, the behavior of these additions changes depending on the number of clients. While accuracies in the 100 client case are lower in general, soft resets and no additions work better.

2 Text Classification

To show that DOSFL is not limited to image-based tasks, we test DOSFL on federated IMDB (sentiment analysis) , federated TREC-6 (question classification) , and federated Sent140 (sentiment analysis) . Directly applying dataset distillation for language tasks is challenging as text data is discrete. Each token is a one-hot vector with dimension equal to the number of tokens, which can be in the thousands. To overcome this issue, we use pre-trained GloVe embeddings with a look up table to convert one-hot token ids to word vectors in 100D Euclidean space . Distilled sentences now are fixed-size real-valued matrices. Real sentences are also padded or truncated to the same fixed length: 200 for federated IMDB and Sent140, 30 for federated TREC-6.

The results, best out of 5 runs, are provided in Table 2. We also provide the accuracy of non-IID FedAvg after an equivalent amount of communication. This allows DOSFL and FedAvg to be compared in an equal communication cost setting. We tuned the FedAvg hyperparameters to maximize initial learning: local epochs E=5E=5, batch size B=10B=10, and learning rate η=0.01\eta=0.01. The batch size was the lowest in the considered range {10,20,50,100,200,500}\{10,20,50,100,200,500\} and the learning rate the largest in {10−2,3×10−3,10−3,3×10−4,10−4}\{10^{-2},3\times 10^{-3},10^{-3},3\times 10^{-4},10^{-4}\}. Note that, the preserved accuracy, defined as the DOSFL accuracy over baseline performance, of all tasks is at least 93%93\%. This shows that DOSFL is capable of handling different tasks with small or large datasets in one round.

For federated IMDB, we use a simple CNN model called TextCNN. We test DOSFL with soft labels for 10 and 100 clients, IID and non-IID federated IMDB. Here the distill steps Sd=5S_{d}=5, the distill epochs Ed=10E_{d}=10, the distill batch size Bd=1B_{d}=1, and the starting distill learning rate is η0=0.01\eta_{0}=0.01. Since there exists only 2 classes in IMDB dataset (positive or negative sentiment), non-IID performance is within 2% of IID. Approximately \nicefrac34\nicefrac{{3}}{{4}} clients contain labels from all classes, whereas in federated MNIST no client can have more than 4 classes.

For federated TREC-6, we adopt a Bi-LSTM model to show that DOSFL can be used with non-CNN models. We use 2 and 29 for the number of clients, since the size of the dataset is 5452 and the client dataset sizes must be divisible by the shard count s=2s=2. The amount of training data for the 2 client federated TREC-6 and 10 client federated IMDB are almost equal (2726∼25002726\sim 2500). Similarly, 29 client federated TREC-6 is comparable with 100 client federated IMDB (188∼250188\sim 250). We have Sd=2S_{d}=2, Ed=1E_{d}=1, Bd=1B_{d}=1, and η0=1.5\eta_{0}=1.5. Due to the low number of clients, we were able to reduce the amount of distilled data needed compared to the previous two tasks. Unlike federated IMDB, there is a larger ∼6%\sim 6\% gap in accuracy between the IID and non-IID settings. Furthermore, we extend DOSFL to a larger dataset, Sent140, using TextCNN. The hyperparameters are distill steps Sd=5S_{d}=5, distill epochs Ed=15E_{d}=15, Bd=1B_{d}=1, and initial learning rate η0=0.3\eta_{0}=0.3.

3 DOSFL with Stragglers and Low Participation

So far The above mentioned settings assumes synchronous full participation which is highly unlikely in the real-world. We design an alternate version of DOSFL for when the participation rate is extremely low such that client communication is almost one-by-one. First, the server selects only one client to distill its data. Afterwards, the server updates the global model by training on the distilled data. A different client then performs dataset distillation targeting the updated parameters. This process repeats until each client has distilled their data. Hence, the global model is updated NN times, once for each client. The communication is highly serial; the next client can only begin distillation after the current client finishes.

We name this setting LP-DOSFL and evaluate its performance on MNIST in Figure 2, using the best out of 5 trials again. Surprisingly, LP-DOSFL achieves almost 99% accuracy when MNIST is IID, 7%7\% more than vanilla DOSFL. This advantage disappears when the data becomes non-IID. The final accuracy of LP-DOSFL is only 1%1\% larger than plain DOSFL. The reason for this is simple: when the clients’ datasets are IID, a model trained on one client’s dataset will transfer to the others. Encouragingly, the server model achieves its final accuracy after as few as 15% of the clients finish dataset distillation. This is true even when the data is non-IID, although it takes longer (around 40% of the clients). Thus, the total amount of communication needed for LP-DOSFL is less than DOSFL.

Discussion

We now compare the total communication cost (TCC) of DOSFL with that of FedAvg measured in the amount of scalar values sent between the clients and the server. Since the server model’s initialization can be distributed as a random seed, we ignore the cost of the first server-to-client transmission. Let C=0.1C=0.1 be the fraction of the NN clients that participate each round. FedAvg sends Θ\Theta server model parameters to each client, who responds with locally trained parameters. Let TT be the number of communication rounds.

For DOSFL, we only need to consider the single expense of sending distilled data to the server. Let ndatan_{data} be the number of elements in each data point and BdB_{d} be the batch size of the distilled data.

Both formulas can be used to calculate the number of communication rounds—the break even round—needed for lifetime cost of FedAvg to equal DOSFL for the tasks in Section 4.

Note that this value is independent of the number of clients NN. Break even rounds for federated MNIST, IMDB, TREC-6, and Sent140 are provided in Table 3 along with the data and model size. We also investigate potential communication savings when using larger models, such as Transformers, for Sent140. The higher break even round for MNIST, compared to the text tasks, is due to LeNet having significantly fewer parameters than either TextCNN or Bi-LSTM. The best accuracy at the break even round Tbreak evenT_{break\ even} is reported in Table 2.

Finally, we conclude our discussion of communication efficiency by comparing DOSFL with FedAvg under an iso-accuracy setting. Let Tiso accuracyT_{iso\ accuracy} be the number of communication rounds required for FedAvg to reach the accuracy achieved in Table 2. Define the DOSFL to FedAvg ratio as

We can also calculate the percent communication reduction using the above ratio.

We choose the smallest iso-accuracy round Tiso accuracyT_{iso\ accuracy} out of 5 trials. The values for both quantities is shown in Table 4. For some tasks, FedAvg fails to converge due to having only 1 client update the gradient each round. In general, DOSFL saves more communication when the size of the model increases or when the dataset becomes more challenging. Besides MNIST, communication savings are up to 3 orders of magnitude.

DOSFL provides an efficient way to trade computation and a few-to-none percentage points of accuracy for a great amount of communication cost reduction. Next, clients can choose to either adopt another Federated Learning algorithm to continue improving the global model or personalize by training on local data. As such, DOSFL is best suited for cross-silo FL, where 2-100 organizations seek to learn a shared model without sharing data . In cross-silo learning, participants likely would be able to dedicate hardware for the sole purpose of FL. Big models are also probable, and the communication savings of DOSFL increase as the models get larger. Computation resources are also cheaper to acquire compared to communication resources. A company can purchase dedicated GPUs and use them for years at reasonable cost without losing much of their intrinsic value.

2 Privacy and security

Suppose the server or a client attempted to learn private information about other clients. To the human eye, the distilled images or text appear random (see Appendix B.1, B.2). So the only known way to extract information from the distilled data is train the targeted model. Because distilled data are targeted towards a specific initialization, training a differently initialized model on distilled data should result in a low accuracy model. Therefore, we examine the security of DOSFL against an eavesdropping attack and show that an attacker cannot reproduce the global model without the server’s initial weights. Referring to Figure 1, assume that the attacker intercepts the distilled data, labels and learning rates from each client between Step 2 and 3.

As such, the privacy of DOSFL should be no worse than plain FedAvg. There are some security risks from other types of attacks. In , the authors used dataset distillation for data poisoning attacks. Images were distilled such that, after one gradient descent step, the final model would misclassify the attack category. However, this security risk is also present with FedAvg and most other FL algorithms: clients can deliberately upload poisoned weights to the server . For DOSFL, bounding the gradient value or using momentum when training the global model on distilled data could mitigate this threat. While defenses against other types of attacks are beyond the scope of this paper, we believe that existing differential privacy and secure multi-party computation tools will prove sufficient.

References

Appendix A Hyperparameters

Appendix B Distilled Data

B.2 IMDB

We provide a distilled sentence from one of 100 clients for federated IMDB with IID distribution. The logit is 1.63 for the positive class and 0 for the negative class. The corresponding distill learning rate is 0.0272.

shaw malone assembled shelly pendleton tha insanity vietnam finishes morton leather watts respectable mastery funky idle watched peripheral ely glossy 1934 honed periods suppress setting eden arises resides moses aura succumb prc missing dyer angela emulate showcased meredith embraces bonnie translates replicate potts segment affects enhances stein juliet bumping mystic resistance token alienate hays unnamed mira rewarded fateful aspire uniformly bliss mermaid burnt joins unforgettable martino namely marshal ivan morse segment pleads boasting victorian closeness rafael reid saddle boot hawks lingered landon … Further, we exhibit a distilled sentence for non-IID federated IMDB. The logit is 1.68 for the positive class and 0 for the negative class. The corresponding distill learning rate is 0.0284.

B.3 TREC6

In addition, we show a distilled sentence from 1 out of 29 clients for federated TREC6 with IID distribution. The logit is 1.96 for class 1, and 0 for the remaining classes. The corresponding distill learning rate is 2.25.

conversion loop monster manufactured causing besides stealing yankee 1932 igor supplier nicholas lloyd sees businessman alternate alternate photograph portrayed tale trials 49 principal sequel authors topped donation fictional bull philip At last, we illustrate a distilled sentence from 1 out of 29 clients for federated TREC6 with non-IID distribution. The logit is 1.58 for class 2, and 0 for the remaining classes. The corresponding distill learning rate is 1.87.

Appendix C Additional Results

We provide results for LP-DOSFL on non-federated MNIST tasks in Table 6: federated IMDB, TREC-6, and Sent140. Hyperparameters and methodology are identical to those used for regular DOSFL in Section 4, other than the change in distillation order from parallel to serial.

C.2 Moderate Non-IID

We perform an analysis of the impact non-IIDness has on DOSFL. We ran vanilla and LP-DOSFL on 10 client Federated MNIST for shard counts 22 through 3030. The results are given in Figure 7. Importantly, plain DOSFL and LP-DOSFL maintain their IID performance even as the shard count ss drops to 1010. This is a moderately non-IID setting; each client on average still contains examples of all 10 digits. However, LP-DOSFL slightly curves downward as ss decreases while vanilla DOSFL is flat. Beyond this point, test accuracy degrades quickly until both vanilla and LP-DOSFL have similar test accuracies once s=2s=2.