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 , dataset distillation minimizes the loss of adapted parameters , obtained by performing gradient descent on 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 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 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 each with their own local models with parameters and loss functions . Given some probability vector (each and ), our goal is to find some parameters that minimize the weighted sum .
Our solution consists of 3 steps. These steps are summarized in Algorithm 1.
A central server randomly initializes model parameters . This can be distributed to the clients as a random seed.
Afterwards, minimize the loss of evaluated on a batch of real data .
The clients upload the distilled data to the server. If , the server sorts the distilled data by index, e.g. from clients 1 and 2 become where 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 and with distilled examples and respectively with . The server first trains on , arriving at , which is then trained on . But has been distilled to train . To combat this interference, we introduce two new techniques for improving performance on non-IID data.
Soft resets sample the starting parameters, from a Gaussian distribution, around the server’s parameters . By sampling between distillation iterations, dataset distillation learns more robust examples capable of training any model with weights . This technique is based off of the “hard resets” introduced in , which completely re-initializes . 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 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 that is halved every epochs. We have , , and 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 epochs for image datasets and epochs for text datasets with a batch size of . 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, 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 shards of equal length. Starting from empty subsets, the shards are randomly assigned to the subsets until each has shards. As 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 . 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, 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 times, and the best result is reported in Table 1.
The distill steps are , the distill epochs are , and the initial distill learning rate is . 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 and masking probability were chosen from the search spaces and . 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 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 , batch size , and learning rate . The batch size was the lowest in the considered range and the learning rate the largest in . Note that, the preserved accuracy, defined as the DOSFL accuracy over baseline performance, of all tasks is at least . 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 , the distill epochs , the distill batch size , and the starting distill learning rate is . Since there exists only 2 classes in IMDB dataset (positive or negative sentiment), non-IID performance is within 2% of IID. Approximately 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 . The amount of training data for the 2 client federated TREC-6 and 10 client federated IMDB are almost equal (). Similarly, 29 client federated TREC-6 is comparable with 100 client federated IMDB (). We have , , , and . 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 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 , distill epochs , , and initial learning rate .
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 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, more than vanilla DOSFL. This advantage disappears when the data becomes non-IID. The final accuracy of LP-DOSFL is only 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 be the fraction of the clients that participate each round. FedAvg sends server model parameters to each client, who responds with locally trained parameters. Let be the number of communication rounds.
For DOSFL, we only need to consider the single expense of sending distilled data to the server. Let be the number of elements in each data point and 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 . 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 is reported in Table 2.
Finally, we conclude our discussion of communication efficiency by comparing DOSFL with FedAvg under an iso-accuracy setting. Let 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 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 through . The results are given in Figure 7. Importantly, plain DOSFL and LP-DOSFL maintain their IID performance even as the shard count drops to . 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 decreases while vanilla DOSFL is flat. Beyond this point, test accuracy degrades quickly until both vanilla and LP-DOSFL have similar test accuracies once .