Think Locally, Act Globally: Federated Learning with Local and Global Representations
Paul Pu Liang, Terrance Liu, Liu Ziyin, Nicholas B. Allen, Randy P. Auerbach, David Brent, Ruslan Salakhutdinov, Louis-Philippe Morency
Introduction
Federated learning is an emerging research paradigm to train machine learning models on private data distributed in a potentially non-i.i.d. setting over multiple devices . A key challenge involves keeping private all the data on each device by training a global model only via communication of parameter updates to each device. This relies on the global model being sufficiently compact so that the parameters and updates can be sent efficiently over existing communication channels such as wireless networks . However, the recent demands in larger models pose a challenge for deploying federated learning on real-world tasks. In this paper, we propose a new federated learning algorithm, Local Global Federated Averaging (LG-FedAvg), which jointly learns compact local representations on each device and a global model across all devices. We perform a generalization analysis of federated learning which shows that a combination of local and global models reduces both variance in the data as well as variance across device distributions, which is more optimal than either extreme. To support our theoretical analysis, we perform a wide range of experiments that suggest local representation learning is beneficial for the following reasons:
1) Efficiency: Having local models extract useful, lower-dimensional representations means that the global model now requires fewer number of parameters, thereby reducing the number of parameters and updates that need to be communicated to and from the global model as well as the bottleneck in terms of communication cost. Our proposed method also maintains performance on publicly available datasets spanning image recognition (MNIST, CIFAR) and multimodal learning (VQA).
2) Heterogeneity: Real-world data is often heterogeneous (coming from different sources). A new device could contain sources of data that have never been observed before during training, such as images of a different domain or different texting styles on personalized mobile devices. Local representations allow us to process new device data using specialized encoders depending on their source modalities instead of using a single global model that might not generalize to new modalities and distributions . We show that our model learns personalized mood predictors from real-world private mobile data and better deals with heterogeneous data never seen during training.
3) Fairness: Real-world data often contains sensitive attributes and recent work has shown that it is possible to recover these attributes from data representations without access to the data itself . We show that local models can be modified to learn fair representations that obfuscate protected attributes such as race, age, and gender, a feature crucial to preserving the privacy of on-device data.
Related Work
Federated Learning aims to train models in massively distributed networks at a large scale , over multiple sources of heterogeneous data , and over multiple learning objectives . Recent methods aim to improve the efficiency of federated learning , perform learning in a one-shot setting , propose realistic benchmarks , and reduce the data mismatch between local and global data distributions . While several specific algorithms have been proposed for heterogeneous data, LG-FedAvg is a more general framework that can handle heterogeneous data from new devices, reduce communication complexity, and ensure fair representation learning. We compare with these existing baselines and show that LG-FedAvg outperforms them in heterogeneous settings.
Distributed Learning is a related field with similarities and key differences: while both study the theory and practice involving the partition of data and aggregation of updates , federated learning is additionally concerned with data that is private and distributed in a non-i.i.d. fashion. Recent work has improved communication efficiency by sparsifying the data and model , developing efficient gradient methods , and compressing the updates . These compression techniques are complementary to our approach and can be applied to our local and global models.
Representation Learning involves learning features from data for generative and discriminative tasks. A recent focus has been on learning fair representations , including using adversarial training to learn representations that are not informative of private attributes such as demographics and gender . A related line of research is differential privacy which constraints statistical databases to limit the privacy impact on individuals whose information is in the database . Our approach extends recent advances in adapting federated learning for heterogeneous data and fairness. LG-FedAvg is a general framework that can handle heterogeneous data from new devices, reduce communication complexity, and ensure fair representation learning.
Local Global Federated Averaging
At a high level, LG-FedAvg combines local representation learning with global model learning in an end-to-end manner. Each local device learns to extract higher-level representations from raw data before a global model operates on the representations (rather than raw data) from all devices. An overview of LG-FedAvg is shown in Figure 1. The local and global learning procedures are designed to be complementary: local representation learning aims to extract high level, compact features important for prediction, thereby allowing the global model to save parameters by operating only on lower dimensional representations. At the same time, the global model objective ensures that the global model must be able to classify data from all devices, thereby ensuring that the local representations are general enough instead of overfitting to the subset of data on each device. We begin by describing how local (§3.1) and global (§3.2) learning is performed. We then detail one example of adversarial local learning to learn fair local representations (Appendix B.1).
For each source of data , we learn a representation which should: 1) be low-dimensional as compared to raw data , 2) capture important features in that are useful towards the global model, and 3) not overfit to device data which may not align to the global data distribution. To be more concrete, we define features that should be captured using a good representation . In Figure 1(a) through 1(c) we summarize these local learning methods according to the choice of : (a) the labels (supervised learning), (b) the data itself (unsupervised autoencoder learning), or (c) some auxiliary labels (self-supervised learning). For simplicity, we focus the description on supervised learning but describe extensions to local adversarial learning of fair representations (Figure 1(d)) and unsupervised learning in Appendix B.1.
2 Global Aggregation
3 Inference at Test Time
Theoretical Analysis
In this section, we provide a theoretical analysis of using local and global models for federated learning. We show that 1) purely local models do not suffer from device variance but suffer from data variance, 2) the opposite holds true for purely global models, and 3) having both local and global models achieves a balance between both desiderata. All detailed proofs can be found in Appendix A. The link between the analysis and LG-FedAvg for deep networks is discussed at the end of the section and verified via comprehensive experiments in § 5.1.
Given datapoints , learning local and global parameters involve optimizing the following training objectives:
We denote the overall model as using both local and global models. The local empirical generalization error on device is defined as
The total empirical generalization error is defined as the mean of all local errors and the true generalization error is defined as the expectation taken over the randomness present in the data, devices, and noise:
which can further be manipulated to obtain a bias-variance decomposition for federated learning.
The generalization loss for federated learning can be decomposed as
Using only local models results in an unbiased estimator of . The bias term arises when learning global parameters since federated learning couples the estimation of both local and global parameters. The variance term comes from both the variance of both local and global parameter estimates.
As a simplified version, LG-FedAvg can be seen as a ensemble of local and global models, i.e. . In this case, one can show that:
For the linear setting we are considering, the above result can be further expanded as:
This shows that local models only control data variance at a rate of since they are only updated using local device data which may be limited in number and vary highly (both in quality and quantity) across devices. However, local models do not suffer from device variance.
The second method updates a joint global model (i.e. vanilla federated learning; ), which is equivalent to setting , i.e. . Its generalization error is:
The generalization error of the global model is .
Global models can control for data variance () at a rate of , decreasing with the total number of datapoints across all devices (since global parameters are updated using data across all devices), which is better than the rate for local models. However, it suffers from an extra term representing device variance so one global model is unable to account for very different devices.
The generalization error is .
This shows that using an ensemble of local and global models reduces both data variance and device variance. When is large (high device variance), one should prioritize local models that better model the local data distributions (larger ). Conversely, when is large (high data variance), one should prioritize a global model (smaller ).
While our theory holds true for linear models, we believe that it provides accurate insight into the practical generalization abilities of LG-FedAvg, where we use deep networks and treat as the split of the layers of between the local and global models. Empirically, we compare Figure 2(a)-(b) (test error for linear models using -interpolation on synthetic data) with Figure 2(d) (test accuracy for deep networks using -split on real-world mobile data). The close similarity implies that our theoretical analysis captures the correct relationship between local and global models.
Experiments
We evaluate how our method 1) verifies our theory under different data and device variances, 2) efficiently reduces parameters while retaining performance, 3) learns personalized models and handles data from heterogeneous sources, two settings with particularly high device variance, and 4) learns fair local representations that obfuscate private attributes. Code is included in the supplementary. Implementation details and sensitivity reports across hyperparameters are provided in Appendix C.
Synthetic Data: We first experiment on synthetic data to verify our theoretical analysis on ensembling local and global models. Data on device is generated by and teacher weights are sampled as , , where represents device variance. Labels are observed with noise, , , where represents data variance. We plot the average test error when local models perform better due to higher device variance (Figure 2(a), ) and when global models perform better due to lower device variance (Figure 2(b), ). In Appendix C.1, we discuss the performance of LG-FedAvg under several other variance settings. For all settings, using an -interpolation of both local and global models performs either close to the optimal extremes or better than either extreme.
CIFAR-10: Next, we verify our theory on deep networks over complex image classification problems using CIFAR-10. We focus on a highly non-i.i.d. setting and follow the experimental design in by assigning examples from at most classes to each device. The value of simulates device variance: represents highest device variance while represents an i.i.d. split of labels to devices. From Figure 2(c), we observe that LG-FedAvg consistently outperforms local only and FedAvg. The performance gap is higher as device variance increases, which supports our theory that local models deal with high device variance.
2 Model Performance & Communication Efficiency
CIFAR-10: We compare our approach with existing federated learning methods with respect to model performance and communication. We randomly assign each device to examples of two classes (highly unbalanced). We consider two settings during testing: 1) Local Test, where we know which device the data belongs to (i.e. new predictions on an existing device) and choose that particular trained local model. For this setting, we split each device’s data into train, validation, and test data, similar to . 2) New Test, in which we do not know which device the data belongs to (i.e. new predictions on new devices) , so we use an ensemble approach by averaging all trained local model logits before choosing the most likely class . For ensembling, all local model weights are sent to the server only once and averaged. We include this parameter exchange step and still show substantial communication improvement. We find that averaging model weights performs similarly ( for CIFAR) to averaging model outputs since deep ReLU nets are mostly linear). We choose LeNet-5 as our base model and we compare 4 methods: 1) FedAvg which is the traditional federated learning approach, 2) Local only as an extreme setting with only local models, 3) MTL which trains local models with parameter sharing in a multi-task fashion, and 4) LG-FedAvg which is our proposed method with local and global models.
The results in Table 1 show that LG-FedAvg gives strong performance with low communication cost. For CIFAR local test, LG-FedAvg significantly outperforms FedAvg since local models allow us to better model the local device data distribution. For new test, LG-FedAvg achieves similar performance to FedAvg while using around of the total parameters, and better performance with the same number of parameters. LG-FedAvg also outperforms using local models only and local models trained with multitask learning (MTL). This shows that our end-to-end training strategy for local and global models is particularly suitable for large neural networks. We show MNIST results and sensitivity analysis wrt data and model splits in Appendix C.2 and find similar observations.
Visual Question Answering (VQA): VQA is a large-scale multimodal benchmark with M images, M questions, and M answers . We split the dataset in a non-i.i.d. manner and evaluate the accuracy under the local test setting. We use LSTM and ResNet-18 unimodal encoders as our local models and a global model which performs early fusion of text and image features for answer prediction. In Table 2, we observe that LG-FedAvg reaches a goal accuracy of while requiring lower communication costs (more VQA results and details in Appendix C.2.3).
3 Learning Personalized Mood Predictors from Mobile Data
Data: We designed and collected a new dataset, Mobile Assessment for the Prediction of Suicide (MAPS), to determine real-time indicators of suicide risk in adolescents aged years. This study monitors adolescents including individuals who have recently attempted suicide, individuals who experience suicidal ideation, and psychiatric controls. Across a duration of months, data was collected from each participant’s smartphone using a keyboard logger which tracks all typed words. Participants were asked to rate their mood for the previous day on a scale ranging from , with higher scores indicating a better mood. All users have given consent for their mobile device data to be collected and shared with us for research purposes. MAPS is a realistic federated learning benchmark since it contains real-world data with privacy concerns and high device variance due to highly personalized use of mobile phones. We used a preliminary preprocessed version containing samples across participants. We discretize the scores into bins for -way classification. We use a random split for training/validation/testing, conduct all experiments times, and report the average accuracy and standard deviation (details in Appendix C.3).
Models: To assess how mobile text data can be used to make personalized mood predictions, we train a MLP classifier on top of a Bi-LSTM encoder on word embeddings. In addition to local only and FedAvg, we test LG-FedAvg across different splits of local and global model layers (i.e. ) while keeping the total parameter count constant.
Results: From Figure 2(d), consistent with our theoretical findings, an -split across local and global models leverages both personalized representations per devices as well as statistical strength sharing through data across all devices, outperforming either local or global extremes.
4 Heterogeneous Data in an Online Setting
Data: We test whether LG-FedAvg can handle heterogeneous data from a new source introduced during testing. We split MNIST across devices in both an i.i.d. and non-i.i.d. setting, then introduce a new device with training and test MNIST examples but rotated degrees. This simulates a drastic change in the data distribution.
Models: We consider 3 methods: 1) FedAvg, 2) FedProx : a method designed specifically for heterogeneous data by regularizing the local updates to reduce overfitting to local devices, and 3) LG-FedAvg: train on the original devices, and when a new device comes, learn local representations before fine-tuning the global model. We hypothesize that good local models can “unrotate” images from the new device to better match the data distribution seen by the global model. When learning on the new device, we also retrain on a fraction of the original devices: implies no fine-tuning and implies some fine-tuning ( implies retraining on all data which is impractical).
Results: We report results in Table 3 and observe that: 1) FedAvg suffers from catastrophic forgetting without fine-tuning (), in which the global model can perform well on the new device’s rotated MNIST but completely forgets how to classify regular MNIST . Only after fine-tuning () does the performance on both regular and rotated MNIST improve, but this requires more communication over the devices. 2) LG-FedAvg with local models relieves catastrophic forgetting. Augmenting local models indeed helps to improve online performance on rotated MNIST while allowing the global model to retain performance on regular MNIST , outperforming both FedAvg and FedProx. We believe LG-FedAvg achieves these results by learning a strong local representation that requires fewer updates from the trained global model.
5 Learning Fair Representations
Data: We examine whether local models can be trained to protect private attributes from the global model. We use the UCI adult dataset to predict whether an individual makes more than K per year based on their personal attributes. However, we want our models to be invariant to the sensitive attributes of race and gender instead of picking up on correlations that could exacerbate biases.
Models: We adapt adversarial learning to remove protected attributes from local models (see Appendix B.1. Specifically, we aim to learn fair local representations from which a fully trained adversarial network should not be able to predict the protected attributes. We report three methods: 1) FedAvg with only a global model and global adversary both updated using FedAvg. The global model is not trained with the adversarial loss since it is simply not possible: once local device data passes through the global model, privacy is potentially violated. 2) LG-FedAvg without penalizing the adversarial network, and 3) LG-FedAvgAdv which jointly trains local, global, and adversary models to learn fair local representations before global prediction.
Results: We report results according to: 1) classifier binary accuracy, 2) classifier ROC AUC score, and 3) adversary ROC AUC score. The classifier metrics should be as close to as possible while the adversary should be as close to as possible. From Table 4, LG-FedAvgAdv learns fair local representations that are unable to predict protected attributes ( adversary AUC) with only a small drop in global accuracy . In order to ensure that poor adversary AUC was indeed due to fair representations instead of a poorly trained adversary, we train a post-fit classifier from local representations to protected attributes and achieve similar random results.
Conclusion
We proposed LG-FedAvg combining local representation learning with federated training of global models. Our theoretical analysis shows that an ensemble of local and global models reduces both data variance and device variance. On a suite of real-world datasets, LG-FedAvg achieves strong performance while reducing communication costs, learns personalized models, better deals with heterogeneous data, and effectively learns fair representations that obfuscate protected attributes.
Acknowledgements
PPL and LM were partially supported by the National Science Foundation (Awards #1750439, #1722822) and National Institutes of Health. RS was supported in part by NSF IIS1763562, Office of Naval Research N000141812861, and Google focused award. Any opinions, findings, and conclusions or recommendations expressed in this material are those of the author(s) and do not necessarily reflect the views of National Science Foundation or National Institutes of Health, and no official endorsement should be inferred. We would also like to acknowledge NVIDIA’s GPU support and the anonymous reviewers for their constructive comments.
Broader Impacts
Federated learning provides tools for large-scale distributed training at unprecedented scales but at the same time requires more research on its implications to society and policy.
Broader applications: By 2025, it is estimated that there will be more than to 75 billion IoT (Internet of Things) devices all connected to the internet and sharing data with each other . Organizing, processing, and learning from device data will use federated learning techniques. It has already been shown to be a promising approach for applications such as learning the social activities of mobile phone users, early forecasting of health events like heart attack risks from wearable devices, and localization of pedestrians for autonomous vehicles . The societal impacts revolving around more invasive federated learning technologies have to be taken into account as we design future systems that increasingly leverage distributed mobile data.
Applications in mental health: Suicide is the second leading cause of death among adolescents. In addition to deaths, 16% of high school students report seriously considering suicide each year, and 8% make one or more suicide attempts (CDC, 2015). Despite these alarming statistics, there is little consensus concerning imminent risk for suicide . Given the impact of suicide on society, there is an urgent need to better understand the behavior markers related to suicidal ideation.
“Just-in-time” adaptive interventions delivered via mobile health applications provide a platform of exciting developments in low-intensity, high-impact interventions . The ability to intervene precisely during an acute risk for suicide could dramatically reduce the loss of life. To realize this goal, we need accurate and timely methods that predict when interventions are most needed. Federated learning is particularly useful in monitoring (with participants’ permission) mobile data to assess mental health and provide early interventions. Our data collection, experimental study, and computational approaches provide a step towards data intensive longitudinal monitoring of human behavior. However, one must take care to summarize behaviors from mobile data without identifying the user through personal (e.g., personally identifiable information) or protected attributes (e.g., race, gender). This form of anonymity is critical when implementing these suicide detection technologies in real-world scenarios. Our goal is to be highly predictive of STBs while remaining as privacy-preserving as possible. We outline some of the potential privacy and security concerns below and show some possibilities brought about from the flexibility of our local models.
Privacy: There are privacy risks associated with making predictions from mobile data. Although federated learning only keeps data private on each device without sending it to other locations, the presence of one’s data during distributed model training will likely affect model predictions. Therefore it is crucial to obtain user consent before collecting device data. In our experiments with real-world mobile data, all participants have given consent for their mobile device data to be collected and shared with us for research purposes. All data was anonymized and stripped of all personal (e.g., personally identifiable information) and protected attributes (e.g., race, gender).
Security: Communicating model updates throughout the training process could possibly reveal sensitive information, either to a third-party, or to the central server. Federated learning is also particularly sensitive to external security attacks from adversaries . Recent methods to increase the security of federated learning systems come at the cost of reduced performance or efficiency. We believe that our proposed local-global models make federated learning more interpretable and flexible since local models can be appropriately adjusted to be more secure. However, there is a lot more work to be done in these directions, starting by accurately quantifying the trade-offs between security, privacy and performance in federated learning .
Social biases: We acknowledge that there is a risk of exposure bias due to imbalanced datasets, especially when personal mobile data is involved. Models trained on biased data have been shown to amplify the underlying social biases especially when they correlate with the prediction targets . Our experiment showcased one example of maintaining fairness via adversarial training, but leaves room for future work in exploring other methods tailored for specific scenarios (e.g. debiasing words , sentences , and images ). These methods can be easily applied to local representations before input into the global model. Future research should also focus on quantifying the trade-offs between bias and performance .
Overall, we believe that our proposed approach can help quantify the tradeoffs between local and global models regarding performance, communication, privacy, security, and fairness. Its flexibility also offers several exciting directions of future work in ensuring privacy and fairness of local representations. We hope that this brings about future opportunities for large-scale real-time analytics in healthcare and transportation using federated learning.
References
Appendix
Appendix A Theoretical Analysis
In this section, we restate our main results and provide details proofs for them. We start with the short comment that, for the federated linear regression problem we are solving, gradient descent converges to the same solution as the commonly used ridgeless linear regression solution, as is shown in . This means that converges to and converges to in expectation. We also assume in our analysis that is small and can be summarized in the big- notation; however, this does not mean that the term is small because the noise present in the dataset might be a function of . For example, in the well-studied label noise literature, the noise rate is often proportional to , cancelling out the term in the denominator . In practice, this happens when the dataset has a fixed probability of wrong labels.
The generalization loss for federated learning can be decomposed as
The generalization error can be derived by taking expectation over the random variables:
which can further be manipulated to obtain a bias-variance decomposition for federated learning:
where we have omitted in the input to for notational conciseness. Using only local models results in an unbiased estimator of . The bias term arises when learning global model parameters since federated learning couples the estimation of both local and global parameters. The variance term comes from both the variance of both local and global parameter estimates.
As a simplified version, LG-FedAvg can be seen as a ensemble of local and global models, i.e. . In this case, one can show that:
When using linear models, we can further expand this result as follows:
Let , be learned through gradient descent algorithm, then , and can be expanded as follows:
Furthermore, by using the fact that , we have that
Using these results, we can compute the generalization error of several federated learning baselines as well as LG-FedAvg.
We begin by analyzing two baselines for federated learning.
This shows that local models only control data variance at a rate of since they are only updated using local device data which may be limited in number and vary highly (both in quality and quantity) across devices. However, local models do not suffer from device variance.
The second method updates a joint global model , which is equivalent to setting in our previous analysis, i.e. . We can compute its generalization error:
The generalization error of the global model is .
Set in equation 18 of corollary ‣ A, we obtain . ∎
Therefore, global models are able to control for data variance () at a rate of , decreasing with the total number of datapoints across all devices (since global parameters are updated using data across all devices), which is better than the rate for local models. However, it suffers from an extra term representing device variance so one global model is unable to account for very different devices.
A.2 Analysis of LG-FedAvg
Given that the above baselines achieve different generalization errors, one should be able to interpolate between the two methods to find the optimal tradeoff point. Our method therefore defines an interpolation between the local and global models:
where . The following theorem gives the generalization of this model.
The generalization error of is .
The proof follows by computing each term in equation 18 of corollary ‣ A
We need to find :
This can be solved to find the optimal that minimizes .
The optimal that minimizes is
Taking derivative w.r.t to , and setting to gives an exact expression for :
This shows that using an ensemble of local and global models reduces both data variance and device variance. When is large (high device variance), one should lean towards using local models that better model the local data distributions (larger ). Conversely, when is large (high data variance), one should lean towards using a global model whose parameters are updated using more data across all devices (smaller ).
Appendix B Fair Representation Learning
In this section we detail one extension of local representation learning to remove information that might be indicative of protected attributes. The data on each device is now a triple drawn non-i.i.d. from a joint distribution where are some protected attributes which the model should not predict. For example, although there exist correlations between race and income which could help in income prediction , it would be undesirable for our models to rely on these correlations since these would exacerbate racial biases.
In Figure 1, we also illustrate several other choices of local representation learning using auxiliary local models for (b) unsupervised autoencoding training, where reconstructs given , and (c) self-supervised learning (e.g. jigsaw solving ), where predicts auxiliary features given . However, when using local supervised learning, and are similar classification branches (Figure 1 (a)) both supervised by the target labels. LG-FedAvg can therefore be trained without to directly learn local representations for the global model to make predictions (equation 1). Thus, LG-FedAvg for supervised learning does not incur additional computational complexity while reducing communicated global parameters.
B.2 Theoretical Analysis of Local Fair Representation Learning
By our choice of the objective function we know that
which implies that we have the lower bound
where is a hyperparameter that controls the tradeoff between the prediction model and the adversary model.
Appendix C Experimental Details and Extra Results
Here we provide all the details regarding experimental setup, dataset preprocessing, model architectures, model training, and performance evaluation. Our anonymized code is attached in the supplementary material. All experiments are conducted on a single machine with 4 GeForce GTX TITAN X GPUs.
A note on hyperparameters used: For MNIST and CIFAR experiments, we would like to emphasize that our initial set of hyperparameters were taken directly from the default set of hyperparamters in https://github.com/shaoxiongji/federated-learning for fair comparison across all baselines and our approach. They were NOT manually tuned for LG-FedAvg to perform better.
For synthetic data, there are no hyperparamters involved.
For experiments on mobile data, since it is a new dataset, we start by training the best vanilla federated learning model using FedAvg and use the exact same set of hyperparameters for LG-FedAvg.
For experiments on fairness, we again use all the default hyperparameters as obtained from the tutorial https://blog.godatadriven.com/fairness-in-ml and associated code https://github.com/equialgo/fairness-in-ml.
We set , number of train samples per device as and the number of test samples per device as . Data on device is generated by and teacher weights are sampled as , , where represents device variance. Labels are observed with noise, , , where represents data variance. We plot the average test error when local models perform better due to higher device variance (Figure 2 left, ) and when global models perform better due to lower device variance (Figure 2 right, ). For both settings, using an interpolation of local and global models performs better than either extremes, which supports our analysis.
We also provide several other results demonstrating the effects of data and device variance on the performance local and/or global models. In particular, we fix data variance and gradually decrease device variance . This results in 4 cases: 1) when local models perform close to optimal (Figure 4 far left, ), 2) when local models perform better (Figure 4 middle left, ), 3) when global models perform better (Figure 4 middle right, ), and 4) when global models perform close to optimal (Figure 4 far right, ). For all settings, using an -interpolation of both local and global models performs either close to the optimal extremes (cases 1 and 4) or better than either extremes (cases 2 and 3).
C.2 Model Performance and Communication Efficiency
Details: In all our experiments, we train with number of local epochs and local minibatch size . We set . Images were normalized prior to training and testing. In our experiments, we take the last two layers to form our global model, reducing the number of parameters to (). Table 5 shows the of hyperparameters used. The dataset can be found here: http://yann.lecun.com/exdb/mnist/. We train LG-FedAvg with global updates until we reach a goal accuracy ( for MNIST) before training for additional rounds to jointly update local and global models. Our results are averaged over 10 runs. # FedAvg and LG Rounds are rounded to the nearest multiple of 5, which we use to calculate the number of parameters communicated. Standard deviations are also reported.
Extra results: In this section we provide results on MNIST in comparison with the baselines, see Table 6. Although MNIST is a slightly smaller dataset, we find that both local and global models help in maintaining performance while using fewer communication parameters.
C.2.2 CIFAR10
Details: We train with number of local epochs and local minibatch size . We set . Images are randomly cropped to size , randomly flipped horizontally with probability , resized to , and normalized. For our model architecture, we chose Lenet-5. We use the two convolutional layers for the global model in our LG-FedAvg method to minimize the number of parameters. We therefore reduce the number of parameters to (). Table 5 shows a table of additional hyperparameters used. The dataset can be found here: https://www.cs.toronto.edu/~kriz/cifar.html. We train LG-FedAvg with global updates until we reach a goal accuracy ( for CIFAR-10) before training for additional rounds to jointly update local and global models. Our results are averaged over 10 runs and we report standard deviations. # FedAvg and LG Rounds are rounded to the nearest multiple of 5, which we use to calculate the number of parameters communicated.
Extra results: In this section we provide more results and also show a sensitivity analysis to various hyperparameters especially regarding the data splits across devices and the local-global model split in our method. See Table 8. Our results are especially strong here: across different data splits (different number of users per device), LG-FedAvg consistently performs better on local test and new test while using fewer parameters.
C.2.3 VQA
Details: We adapt the baseline model from without norm I image channel embeddings. We also substitute the VGGNet used in the original baseline model with a pre-trained ResNet-18 . Finally we use the deep LSTM embedding, which is an LSTM that consists of two hidden layers. For the LG-FedAvg method, the global model uses the two final fully connected layers of the image and question channels, as well as the the additional two fully connected layers following the fusion via element-wise multiplication. The global model reduces the number of parameters to (). We use 50 devices and set number of local epochs , local minibatch size , fraction of devices sampled per round . To train and evaluate our models, we use the data from the following: https://visualqa.org/download.html. Table 10 shows a table of hyperparameters used, which strictly follows the baseline model from . We first trained the best FedAvg model and used the exact same hyperparamters to train LG-FedAvg as well.
Extra results: In Figure 5, we plot the convergence of test accuracy across communication rounds. LG-FedAvg outperforms FedAvg after rounds while requiring only of parameters in FedAvg and continues to improve.
C.2.4 Other Baselines
The MTL baseline is implemented by training local models each with parameters and adding a joint regularization term where and are hyperparmeters, and is a matrix whose -th column is the weight vector for the -th device. We choose which implements mean-regularized multitask learning , which assumes that all the tasks form one cluster, and that the weight vectors are close to their mean. For CIFAR-10, computing the full requires storing a matrix of size . Optimizing and storing this matrix becomes infeasible as the number of users or model size increases. Even for our experimental setting, where and (number of parameters in each local model), running MTL require the following relaxations. First, we reduce to , reducing the size of and therefore the number of parameters communicated per round.
C.3 Learning Personalized Mood Predictors from Mobile Data
Dataset details: We designed and collected a new dataset called Mobile Assessment for the Prediction of Suicide (MAPS). MAPS was designed to elucidate real-time indicators of suicide risk in adolescents ages years. Current adolescent suicide ideators and recent suicide attempters along with aged-matched psychiatric controls with no lifetime suicidal thoughts and behaviors completed baseline clinical assessments (i.e., lifetime mental disorders, current psychiatric symptoms). Following the baseline clinical characterization, a smartphone app, the Effortless Assessment of Risk States (EARS), was installed onto adolescents’ phones, and passive sensor data were acquired for -months. Notably, during EARS installation, a keyboard logger is configured on adolescents’ phones, which then tracks all words typed into the phone as well as they app used during this period. Each day during the -month follow-up, participants also were asked to rate their mood on the previous day on a scale ranging from , with higher scores indicating a better mood.
All users have given consent for their mobile device data to be collected and shared with us for research purposes.
MAPS is a realistic federated learning benchmark since it contains real-world data with privacy concerns and high device variance due to highly personalized use of mobile phones. We used a preliminary preprocessed version containing samples across participants. We discretize the scores into bins for -way classification. We use a random split for training/validation/testing, conduct all experiments times, and report the average accuracy and standard deviation.
Model details: To assess how mobile text data can be used to make personalized mood predictions, we train a MLP classifier on top of a Bi-LSTM encoder. The Bi-LSTM has 128 hidden units and the MLP has two hidden layers, each with size 512. We conduct our experiments over 10 iterations. Within each iteration, we use a random 80/10/10 split for training/validation/testing. We train and validate our model 5 times on this split and select the model that performs best on the validation set. We use the test accuracy of this best-performing model as the test accuracy for the iteration. We report the average accuracy and standard deviation over all 10 iterations in Figure 2(d). In addition to local only and FedAvg results, we plot the performance of LG-FedAvg across different splits of local and global models (i.e. ). Consistent with our theoretical findings, an -split across local and global models leverages both personalized representations per devices as well as statistical strength sharing through data across all devices, outperforming either local or global extremes.
C.4 Heterogeneous Data in an Online Setting
Details: Our experiments for the rotated MNIST follow the same settings and hyperparameter selection as our normal MNIST experiments (section C.2). However, we include an additional device, which randomly samples 3000 and 500 images from the train and test sets respectively and rotates them by a fixed 90 degrees. We show some samples of the rotated MNIST images we used in Figure 6, where the top row shows the normal MNIST images used during training and the bottom row shows the rotated MNIST images on the new test device.
C.5 Learning Fair Representations
Details: For method 1, FedAvg, we train the global model and global adversary for outer epochs, within which the number of local epochs . For methods 2 and 3 involving local models, we begin by pre-training the local models and local adversaries for epochs before joint local and global training for 10 epochs. Table 11 shows the table of all hyperparameters used. Experiments were run 10 times with the same hyperparameters but different random seeds. We aimed to keep the local, global, and adversary models as similar as possible between the three baselines for fair comparison. Apart from the number of local and global epochs all hyperparameters were kept the same from the tutorial https://blog.godatadriven.com/fairness-in-ml and associated code https://github.com/equialgo/fairness-in-ml. The data can be found at https://archive.ics.uci.edu/ml/datasets/Adult.
Appendix D Discussion and Future Work
We believe that LG-FedAvg is a general approach that offers several extensions for future work.
Firstly, combining LG-FedAvg with existing work on compressing the number of parameters and gradient updates could further improve the efficiency of federated learning. For example, existing work in sparsifying the data and model , developing efficient gradient-based methods , and compressing the updates can all be applied to our local and global models. In particular, sparsifying the model through techniques such as distillation and hashing could help to store the local models on devices with small memory and computational power.
Secondly, depending on the test time scenario (i.e. local test vs new test), there is a trade off between the ideal size of local models and the global model. If we know which device the test data belongs to, then having a more accurate local model would allow us to perform better prediction at test time. However, if we do not know which device the test data belongs to, it is important to use a more accurate global model to learn the true data distribution across all devices. Therefore, another step for future work would be dynamically learn the number of layers spread across the local and the global models, in a manner similar to learning dynamic computation steps in neural networks . Different devices which contain different data distributions could use different local models which are dynamically learnt rather than hand-designed by the user. Techniques in neural architecture search could also be helpful for this purpose.
Finally, learning fair representations is of utmost importance as our machine learning systems are deployed in real-life settings such as healthcare, law, and policy-making. In addition to the adversarial training method we described in this paper, there are a variety of methods for learning fair representations that can also be incorporated into our flexible local models. For example, recent work has shown that pre-trained word and sentence representations encode and exacerbate gender, race, and religious biases . Incorporating these debiasing methods for text data would be an important step towards learning fair and unbiased local representations in federated learning.