Variational Federated Multi-Task Learning

Luca Corinzia, Ami Beuret, Joachim M. Buhmann

I Introduction

Large scale networks of remote devices, like mobile phones, wearables, smart homes, self-driving cars, and other IoT devices are becoming a significant source of data to train statistical models. As a consequence, there has been a growing interest to develop machine learning paradigms that can take into account distributed data-structure, despite the several challenges arising in this setting. (i) Security: Data generated by remote devices is often privacy-sensitive and its centralized collection and storage is governed by data protection regulations (e.g GDPR and the Consumer Privacy Bill of Rights ). Learning paradigms that do not access user data directly are hence desired. (2) System: Remote devices in these networks have typically important storage and computational capacity constraints, limiting the complexity and the size of the model that can be used. Moreover, the communication of information between devices or between the central server and devices, mostly happens on wireless networks. Hence communication cost can become a significant bottleneck of the learning process. (3) Statistical: The devices of the network typically generate samples with different user-dependent probability distributions, making the setting in general strongly non-IID. While it is a challenge to achieve high statistical accuracy for classical federated and distributed algorithms in this setting, a multi-task learning (MTL) approach can tackle heterogeneous data more naturally. Every device of the network requires a task-specific model, tailored for its own data distribution, to boost the performance of each task.

Federated learning (FL) has emerged to address the scenario of learning models on private distributed data sources. It assumes a federation of devices called clients that both collect the data and execute an optimization routine, and a server that coordinates the learning process by receiving and sending updates from and to the clients. This paradigm has been applied successfully in many real-world cases, e.g, to train smart keyboards in commercial mobile devices and to train privacy-preserving recommendation systems . Federated Averaging (FedAvg) is the state-of-the-art for federated learning with non-convex models and requires all clients to share the same model. Hence, it does not address the statistical challenge of strongly skewed data distributions, and while it has been shown to work well in practice for a range of (non-federated) real-world datasets, it performs poorly in heterogeneous scenarios . We address this problem by introducing VIRTUAL (VarIational fedeRaTed mUlti tAsk Learning), a new framework for federated MTL. In VIRTUAL, the central server and the clients form a hierarchical Bayesian network and the inference is performed using variational methods. Every client has a task-specific model that benefits from the server model in a transfer learning fashion with lateral connections. A part of the parameters are shared between all clients, and another part is private and tuned separately. The server maintains a posterior distribution that represents the plausibility of the shared parameters. In one step of the algorithm, the posterior is communicated to the clients before the training starts, while during training the clients update the posterior given the likelihood of their local data. Finally, the posterior update is sent back to the central server.

Our main contributions are twofold: (i) We address for the first time the problem of federated MTL for generic non-convex models, designing an additional MT metric and proposing VIRTUAL, an algorithm to perform federated training with strongly non-IID client data distributions. (ii) We perform extensive experimental evaluation of VIRTUAL on real-world federated datasets, showing that it outperforms the current state-of-the-art in FL, and simultaneously allowing lower communication costs.

II The VIRTUAL algorithm

In FL, KK clients are associated with KK datasets D1,…,DK\mathcal{D}_{1},\dots,\mathcal{D}_{K}, where Di≔{xi(n),yi(n)}n=1Ni\mathcal{D}_{i}\coloneqq\{\textbf{x}_{i}^{(n)},y_{i}^{(n)}\}_{n=1}^{N_{i}} is in general generated by a client dependent probability distribution function (pdf) and only accessible by the respective client. It is natural to fit KK different models, one for each dataset, enforcing a relationship between models using parameter sharing . This approach has been investigated extensively, and it has been shown to boost effective sample size and performance in MTL for neural networks .

Let assume a star-shaped Bayesian network with a server SS with model parameters θ\bm{\theta}, as well as KK clients with model parameters {ϕi}i=1K\{\bm{\phi}_{i}\}_{i=1}^{K}. Assume that every client is a discriminative model distribution over the input given by p(yi(n)∣xi(n),θ,ϕi)p(y_{i}^{(n)}|\bm{x}_{i}^{(n)},\bm{\theta},\bm{\phi}_{i}) (a straight-forward extension of the work could consider also generative models). Each dataset Di\mathcal{D}_{i} has a likelihood that factorizes as p(Di∣θ,ϕi)=∏n=1Nip(yi(n)∣xi(n),θ,ϕi)p(\mathcal{D}_{i}|\bm{\theta},\bm{\phi_{i}})=\prod_{n=1}^{N_{i}}p(y_{i}^{(n)}|\bm{x}_{i}^{(n)},\bm{\theta},\bm{\phi}_{i}). Following a Bayesian approach, we assume a prior distribution over all network parameters p(θ,ϕ1,…,ϕK)p(\bm{\theta},\bm{\phi_{1}},\dots,\bm{\phi_{K}}). The posterior distribution over all parameters, given all datasets D1:K≔{D1,…,Dk}\mathcal{D}_{1:K}\coloneqq\{\mathcal{D}_{1},\dots,\mathcal{D}_{k}\} reads then

where we enforce that client-data is conditionally independent given server and client parameters, p(D1:K∣θ,ϕ1,…,ϕK)=∏i=1Kp(Di∣θ,ϕi)p(\mathcal{D}_{1:K}|\bm{\theta},\bm{\phi_{1}},\dots,\bm{\phi_{K}})=\prod_{i=1}^{K}p(\mathcal{D}_{i}|\bm{\theta},\bm{\phi_{i}}) , and a factorization of the prior as p(θ,ϕ1,…,ϕK)=p(θ)∏i=1Kp(ϕi)p(\bm{\theta},\bm{\phi_{1}},\dots,\bm{\phi_{K}})=p(\bm{\theta})\prod_{i=1}^{K}p(\bm{\phi_{i}}) . The Bayesian network is illustrated in Figure 1a.

II-B The optimization procedure

The posterior given in Equation 1 is in general intractable and hence we have to rely on an approximation inference scheme (e.g. variational inference, sampling, expectation propagation ). Here we propose an expectation propagation (EP) like approximation algorithm that has been shown to be effective and to outperform other methods when applied in the continual learning (CL) setting . Let us denote the collection of all client parameters by ϕ=(ϕ1,…,ϕK)\bm{\phi}=(\bm{\phi}_{1},\dots,\bm{\phi}_{K}). Then we define a proxy posterior distribution that factorizes into a server and a client contribution for every client ii as

The fully factorization of both server and client parameters allows us to perform a client update that is independent from other clients, and to perform a server update in the form of an aggregated posterior that preserves privacy.

Given a factorization of this kind, the general EP algorithm refines one factor at each step. It first computes a refined posterior distribution where the refining factor of the proxy is replaced with the respective factor in the true posterior distribution. It then performs the update minimizing the Kullback-Leibler (KL) divergence between the full proxy posterior distribution and the refined posterior. The optimization to be performed for our particular Bayesian network and factorization is given by the following.

Assuming that at step tt the factor ii is refined, then the proxy pdf si(t)(θ)s_{i}^{(t)}(\bm{\theta}) and ci(t)(ϕi)c^{(t)}_{i}(\bm{\phi}_{i}) are found minimizing the variational free energy function Li≔L(si(θ),ci(ϕi))\mathcal{L}_{i}\coloneqq\mathcal{L}(s_{i}(\bm{\theta}),c_{i}(\bm{\phi}_{i})), with

where s(t)(θ)=si(θ)∏j≠iKsj(t−1)(θ)s^{(t)}(\bm{\theta})=s_{i}(\bm{\theta})\prod_{j\neq i}^{K}s_{j}^{(t-1)}(\bm{\theta}) is the updated posterior over the server parameters.

The free energy in Proposition 1 can be optimized using gradient descent and the reparametrization trick . For simplicity, we use a Gaussian mean-field approximation of the posterior, hence for server and client parameters, the factorization reads respectively si(θ)=∏d=1DsN(θd∣μids,σids)s_{i}(\bm{\theta})=\prod_{d=1}^{D^{s}}\mathcal{N}(\theta_{d}|\mu_{id}^{s},\sigma_{id}^{s}) and ci(ϕi)=∏d=1DicN(ϕid∣μidc,σidc)c_{i}(\bm{\phi}_{i})=\prod_{d=1}^{D_{i}^{c}}\mathcal{N}(\phi_{id}|\mu_{id}^{c},\sigma_{id}^{c}), where DsD^{s} and {Dic}i=0K\{D_{i}^{c}\}_{i=0}^{K} are respectively the total number of parameters of the server and client networks. A depiction of the full graphical model of the approximated variational posterior is given in Figure 1b. The pseudo-code of VIRTUAL is described in Algorithm 1. The structure of the algorithm is equivalent to the FedAvg and the FedProx algorithms. At each round, a subset of clients is selected and trained with the local free energy given in Proposition 1. The client update is then computed as the ratio of the client parameter distribution before and after the training. The ratio corresponds to the simple difference of the sufficient statistics in the case of exponential family distributions (see Appendix B), that hence is equivalent to the delta computation in the FedAvg setting. In the general setting, it can be computed as the un-normalized pdf Δi=si(t)(θ)si(t−1)(θ)\Delta_{i}=\frac{s_{i}^{(t)}(\bm{\theta})}{s_{i}^{(t-1)}(\bm{\theta})}. The client ii communicates the delta Δi\Delta_{i} to the main server, that aggregates all the received updates into a single update Δ=∏i∈CtΔi\Delta=\prod_{i\in\mathcal{C}_{t}}\Delta_{i}. A major difference with the typical non-MTL setting, where the new server model is given by averaging all selected active clients at any given round, in our case we aggregate updates and the information of non-active clients is effectively retained in the server posterior. In the simple case of mean-field Gaussian approximation of the posterior distribution, the delta computation and the aggregation easily generalize the subtraction and average used in FedAvg, taking into account the uncertainty in the parameters represented by the standard deviation of the learned Gaussian (see Appendix B). Notice that similarly to FedAvg, privacy is preserved since at any time the server can get access only to the overall posterior distribution s(θ)s(\bm{\theta}) and to the aggregated update Δ(θ)\Delta(\bm{\theta}), and never to the individual factor si(θ)s_{i}(\bm{\theta}) and ci(θ)c_{i}(\bm{\theta}), that are visible only to the respective client.

We can further notice an interesting similarity of the VIRTUAL algorithm to the Progress&Compress method for CL introduced in , where a similar free energy is obtained heuristically by composing CL regularization terms and distillation cost functions .

In the experimental section we will make use of a slight modification of the free energy given in Proposition 1 where the KL divergence terms are weighted by a regularization multiplier β\beta. The Kl multiplier has been widely used already in other scenarios, e.g., disentanglement in unsupervised learning , where it has been shown that at different values of β\beta achieves various degree of disentanglement in the embedding space. We will show also in this case that a tuning of the KL multiplier can enhance the performance of the model.

III Related Work

We here provide a brief survey of work in the area of distributed/federated learning and of transfer/continual learning, in light of the problem at hand described in Section I and of the tools used in deriving VIRTUAL.

Distributed learning is a learning paradigm for which the optimization of a generic model is distributed in a parallel computing environment with centralized data . Early work on this paradigm propose various learning strategies that require iterative averaging of locally trained models, typically using Stochastic Gradient Descent (SGD) steps in the local optimization routine . Distributed learning typically consider the learning to be set in a computational cluster, hence with few computing devices, fast and reliable communication between devices, and centralized unbalanced datasets. FL eliminates all these latter constraints and it is framed as a paradigm that encompasses the new challenges and desiderata listed in Section I. FedAvg has been proposed as a straightforward heuristic for the FL. At each step of the algorithm, a subset of online clients is selected randomly, and these are then updated locally using SGD. The models are then averaged to form the model at the next step, which is maintained in the server and transmitted back to all clients. Despite working well in practice, it has been shown that the performance of FedAvg can degrade significantly for skewed non-IID data .

Some heuristics have been proposed recently to solve the statistical challenges of FL. In particular recently it has been proposed to share part of the client-generated data or a server-trained generative model to the whole network of clients. These solutions are however questionable since they require significant communication effort and do not comply with the standard privacy requirements of FL. Another solution for this problem has been proposed in , where the authors extend FedAvg into FedProx, an algorithm that prescribes clients to optimize the local loss function, further regularized with an quadratic penalty anchored on the weights of the previous step. Despite showing improvements on the FedAvg algorithm for highly data-heterogeneous settings, the method is strongly inspired by early works on continual and transfer learning (see, e.g., Elastic Weight consolidation in and the literature review in the next paragraph) and hence can be further refined.

The first contribution to highlight the possibility of naturally embedding FL in the MTL framework has been reported by MOCHA , that extends some early work on distributed MTL-like CoCoA and variations . In this work a federated primal-dual optimization algorithm is derived for convex models with MTL regularization, and it is shown for the first time that the MTL framework can enhance the model performance, with the MTL model outperforming global models (trained with centralized data) and local models as well, on real world federated datasets. This method, however, can be only used on convex models, hence it does not constitute a usable benchmark for deep learning models that are used in the experimental section here.

More recently, multiple efforts have been made to align and match the weights of client models before aggregation in a layer-wise fashion. This is done to minimize the effect of averaging model weights that do not correspond to each other, due to the overparametrization and the symmetry of neural network parametrization. Note that, despite such efforts are applied to standard algorithms as FedAvg, they can be extended to our framework seamlessly.

A Bayesian approach for distributed datasets, similar to the method proposed in this paper, is developed in , where Expectation Propagation (EP) and its variations are performed on a generic partition of the dataset for distributed inference. However, the authors propose only a global model, hence with no structure of shared and non shared parameters between the server and clients. The training is further performed according to a single loss function, hence not in a MTL setting. Moreover, inference is performed using heavy MCMC methods to estimate the moment of the local distributions, limiting the scale of the model considered. In , a further variational framework for generic partitioned data is described. It can be noted that the frameworks proposed encompasse also our method if applied to the particular MTL BayesNet in Figure 1a. However the case study and the experimental section are focused on a classic EP algorithm based on moment matching and heavy MCMC simulation for moment estimation, and hence it can only address limited size models.

Transfer and Continual Learning

The transfer of knowledge in neural network, from one task to another, has been used extensively and with great success since the pioneering work in of transferring information from a generative to a discriminative model using fine tuning. The application of this straightforward procedure is however difficult to apply in scenarios where multiple tasks from which to transfer from are available. Indeed a good target performance can be obtained only with a priori knowledge of task similarity, that is usually not known, while learning of sequential tasks causes knowledge of previous tasks to be abruptly erased from the network in what has been called catastrophic forgetting .

Many methods have been introduced to overcome catastrophic forgetting, and to enable models to learn multiple task sequentially retaining a good overall performance, and transferring effectively to new tasks. Many early works proposed different regularization terms of the loss function anchored to the previous solution in order to achieve new solutions that generalize well on old tasks . These methods have been first introduced as heuristics, but have been found to be applications of well-known inference algorithms like Laplace Propagation and Streaming Variational Bayes , which led to further generalizations . New approaches focused on other components, like architecture innovations, introducing lateral connections that allow new models to reuse knowledge from previously trained models with layer-wise adaptors , and memory enhanced models with generative networks . A recently introduced online Bayesian inference approach served as inspiration for our work. It frames the continual learning paradigm in the Bayesian inference framework, establishing a posterior distribution over network parameters that is updated for any new task in light of the new likelihood function. It has been shown that this method outperformed all previously known methods for CL.

IV Experiments

In this section we present an empirical evaluation of the performance of VIRTUAL on three real-world federated datasets that well represent both the challenges of federated training and of multi-task learning. Due to the fact that in VIRTUAL clients retain a private model at every round, experiments performed on a simulated network on a single GPU have a memory cost that scales linearly with the number of clients (compared to FedAvg that has a constant memory cost on simulation as well). For this reason, the number of clients is bounded in all experiments to 100. Note however that this drawback does not extend to a real network of devices, as in the latter case the client model is retained by the device, and the memory cost at the server side is constant w.r.t the number of clients.

FEMNIST: This dataset consists of a federated version of the EMNIST dataset , maintained by the LEAF project . Different clients correspond to different writers. We subsample 100 random writers and use only the 10 digit labels. Train and test split is provided by the distribution.

Vehicle Sensors Network (VSN)http://www.ecs.umass.edu/~mduarte/Software.html: A network of 23 different sensors (including seismic, acoustic, and passive infra-red sensors) are place around a road segment to classify vehicles driving through. . The raw signal is featurized in the original paper into 50 acoustic and 50 seismic features. We consider every sensor as a client and perform the binary classification of assault amphibious and dragon wagon vehicles.

Human Activity Recognition (HAR)https://archive.ics.uci.edu/ml/datasets/Human+Activity+Recognition+Using+Smartphones: Recordings of 30 subjects performing daily activities are collected using a waist-mounted smartphone with inertial sensors. The raw signal is divided into windows and featurized into a 561-length vector . Every individual corresponds to a different client and we perform classification of 12 different activities (e.g., sitting, walking). For both the VSN and the HAR, a 75%-25% train-test split is performed.

MNIST: The classic MNIST dataset , randomly split into 100 different sections, one section per client. Every client has 600 training samples and 100 test samples. This dataset represents an atypical federated dataset with very homogeneous clients, both in terms of dataset sizes and in term of statistical properties of samples.

Permuted MNIST (PMNIST): The MNIST dataset is randomly split into 100 sections as above, and a random permutation of pixels is applied to every single client dataset. This dataset has been introduced in in the context of continual learning and represent a strongly non-IID federated dataset, with low level features being very dissimilar between clients.

Shakespeare (NLP): This dataset is built concatenating the whole literary production of William Shakespeare . The task here considered is English words spelling, hence the next character prediction task, over a vocabulary size of 86. The characters are arranged in sequences of 80 and further aggregated in batches of 10 sequences. Every role of a play is considered as a individual client, and we excluded roles that do not contain a single full batch.

A comprehensive description of the statistics of the datasets used is available in Table I.

IV-B Experimental setting

We consider a multilayer perceptrons (MLP) with two hidden dense layers (with local reparametrization for the Bayesian counterpart ) with 100 units and ReLU activation functions in the hidden layers, and softmax activation at the output layer. For the NLP task, we use a two-layer LSTM classifier with 100 hidden units per layer and an 8D embedding layer. The Bayesian model makes use of Bayesian LSTMs and Bayesian Gaussian Embeddings . We further use a convolutional neural network for the FEMNIST dataset with two convolutional layers with kernel size 5 and number of filter respectively 32 and 64. We use max-pooling after both convolutional layers, then we adopt an MLP as described above on the flattened activations.

Using the notation of the original paper , all methods are evaluated in all experiments with a number of updated clients per round C=10C=10 and a number of epochs per round E=20E=20, that is high enough to be meaningful on a FL setting, to guarantee convergence in all scenarios while being challenging on complex tasks, like in the case of the Shakespeare dataset. Hyperparameter optimization is performed over a grid of 5 log-spaced client learning rates. Vanilla SGD optimizer is used as the local optimizer for every client. Implementation of VIRTUAL is based on tensorflow and tensorflow distributions packages.

IV-C Metrics

At every round and for any given metric (cross-entropy loss or accuracy), we report two different variations: (i) The centralized metric measured by the central server (S) tested on all client test data; (ii) the Multi-Task (MT) metric measured as the average test metric of all clients models. In both cases, the average is weighted by the client dataset size. For the case of a non-MTL approach (like FedAvg and FedProx), we measure the MT metric using the model that every client deployed last, while in the case of an MTL setting as in Virtual, every client maintains a private model that is tested at every round.

IV-D Effect of the kl divergence weight

We first study the effect of the KL divergence multiplier β\beta on the performance of an MLP network in the FEMNIST dataset. We can observe from Figure 2 that values of β\beta in the range 10−6−10−310^{-6}-10^{-3} do not impair the performance of the model. For a higher value of β\beta the leading term in the free energy Proposition 1 is given by the regularization term and the reconstruction loss is negligible. From Figure 2 we can further observe that in the range of adequate values of β\beta, the server performance is not affected, while increased generalization is achieved in the MT loss, with the β=10−5\beta=10^{-5} being the best performing model. In all following experiments, we tune the β\beta parameter and the l2l_{2} regularizer in the case of the FedProx baseline (with β=0\beta=0 corresponding to FedAvg) in a log-spaced grid of 5 values. We denote as FedProx the best model with a l2l_{2} multiplier strictly larger than 0.

IV-E Full results

In this section, we evaluate the performance of our proposed method and the baselines FedAvg and FedProx in terms of both server and multi-tasks performance.

Metric comparison. From Figure 3 we can first observe that the MT metric is uniformly more stable than the respective server variant since the former assumes clients retaining a private model until further training, while the latter assumes a server model that is updated at every single round, hence experiencing more stochasticity during the training process. Moreover, the MT metric is typically delayed compared to the server counterpart since clients are updated on average only every K/CK/C rounds. This is responsible for the slower progress of the MT metric, with the values of the losses in the first stages of training being up to a factor of 10 larger than the server counterpart. We further notice from Figure 3 (last column) that the MT metrics are typically superior to the centralized metrics (with some exceptions that are examined thoroughly in the following), showing that, at convergence, clients can personalize the model to a specific private dataset.

Method comparison. In Table II we additionally report the maximum accuracy achieved at convergence by the baselines FedAvg and FedProx and our method in all datasets. We can observe that our method outperforms both baselines in almost all datasets (except in the PMNIST datasets) and with all neural network architectures considered (MLPs, convolutional and RNNs). Virtual is able to achieve up to +2%+2\% and +1%+1\% in maximum accuracy respectively in the MT and S variant in the FEMNIST, MNIST and Shakespeare datasets, and marginally outperforms the baselines in VSN. The Shakespeare dataset is particularly crucial. In this dataset, the S metrics are superior then the MT metrics, implying that, at convergence, the clients that train further on the private datasets impair their performance. This follows likely from the high heterogeneity of the size of the Shakespeare datasets (see the last row in table I), that can cause clients with a small dataset to over-fit the private data and under-perform the central server model. Our proposed method performs particularly well in this scenario, causing only a slight reduction of the MT compared to the S accuracy.

IV-F Inducing sparse updates

The use of a Bayesian framework allows us to estimate the importance of the client ii using the property of the client posterior distributions and client un-normalized updates distributions Δi\Delta_{i}. In fig. 4 we examine the cumulative distribution function (CDF) of the signal-to-noise (SNR) of the weights of the client posteriors. A CDF located on the right-hand side of the plot represents a compressible network, with only a small ratio of the weights with high SNR being determinant in the model performance (see for an application of this concept to simple neural network pruning). For comparison, we also implement a variation of the Virtual method that retains the client loss function used in Virtual, and in addition initializes the weights of the client with the server posterior at the beginning of every training round (Virtual + FedAvg init in fig. 4). In these plots, we observe that the server initialization forces all clients to learn a complex model (larger CDF), while with no such initialization the clients specialize to the task at hand with a small ratio of weights.

We, therefore, design a simple updated pruning procedure that sparsifies the client updates setting to zero all Δi\Delta_{i} elements that have a SNR smaller than a given percentile. The results are reported in section IV-F. Virtual greatly outperforms the FedAvg initialization variant and retains a superior performance compared to the FedProx method (compare with Table II, first and second row) up to an induced sparsity of 75%. This is equivalent to a 50% communication reduction cost compared to FedProx, factoring the use of a Bayesian Neural Network that uses twice as many parameters as a standard deterministic network.

Appendix B The Gaussian case

The Virtual method requires to compute the so-called (un-normalized) cavity distributions s(θ)si(θ)\frac{s(\theta)}{s_{i}(\theta)}, the client deltas Δi=s(θ)si(θ)\Delta_{i}=\frac{s(\theta)}{s_{i}(\theta)}, and deltas aggregation Δ=∏i∈CtΔi\Delta=\prod_{i\in\mathcal{C}_{t}}\Delta_{i} (see Algorithm 1 in the main text). Using a factorized Gaussian distribution over the weights as si(θ)=∏d=1DsN(θd∣μids,σids)s_{i}(\bm{\theta})=\prod_{d=1}^{D^{s}}\mathcal{N}(\theta_{d}|\mu_{id}^{s},\sigma_{id}^{s}), we can easily observe that the factorization extends to all three terms listed above. In turn, in order to implement the Virtual algorithm with factorized Gaussian distributions, we need to compute univariate Gaussian products and ratios that read respectively (see [42, Sec 8.1.8]):

where SpS_{p} and ZrZ_{r} are normalization constants and

with σ1<σ2\sigma_{1}<\sigma_{2} in the latter case (ratio). It is easy to observe that using the natural parameterization of the Gaussian distribution with sufficient statistics given by

products and ratios of Gaussians translates respectively into the sum and difference of the natural parameters χ\chi and ξ\xi. The implementation used in the code uses natural parameter Gaussian distributions, and sum and difference of its parameters to obtain the required products and ratios. This parameterization and the consequences on product and ratios easily generalize to any exponential family distribution, but it is here presented in the Gaussian case for convenience.

At step tt the global posterior for server parameters is s(t)(θ)=si(θ)∏j≠iKsj(t−1)(θ)s^{(t)}(\bm{\theta})=s_{i}(\bm{\theta})\prod_{j\neq i}^{K}s_{j}^{(t-1)}(\bm{\theta}) and analogously the client parameters distribution reads c(t)(ϕ)=ci(ϕi)∏j≠iKcj(t−1)(ϕj)c^{(t)}(\bm{\phi})=c_{i}(\bm{\phi}_{i})\prod_{j\neq i}^{K}c_{j}^{(t-1)}(\bm{\phi}_{j}). Then the EP-like update for the model described is given by minimizing the following KL divergence w.r.t si(θ)s_{i}(\bm{\theta}) and ci(ϕi)c_{i}(\bm{\phi}_{i})

where the second equality comes from the normalization of client and server pdfs and from Bayes rule p(θ,ϕi∣Di)∝p(Di∣θ,ϕi)p(ϕi)p(θ)p(\bm{\theta},\bm{\phi}_{i}|\mathcal{D}_{i})\propto p(\mathcal{D}_{i}|\bm{\theta},\bm{\phi}_{i})p(\bm{\phi}_{i})p(\bm{\theta}). Notice also that si(θ)s(t)(θ)=si(t−1)(θ)s(t−1)(θ)\frac{s_{i}(\bm{\theta})}{s^{(t)}(\bm{\theta})}=\frac{s_{i}^{(t-1)}(\bm{\theta})}{s^{(t-1)}(\bm{\theta})} because of the factorization in Equation 2, and hence Proposition 1 is proved. ∎

In all experiments reported here and in the main text we use a batch size B=20B=20, except for the Shakespeare dataset, where B=10B=10. For both baselines and Virtual we perform a careful log spaced grid-search over 5 values of the regularization multiplier β\beta and of the client learning rate ηc\eta_{c}, verifying that optimal values do not lie at the boundaries of the grid. Similarly we use a linearly spaced grid for the server learning rate ηs\eta_{s} of the type {0.2,0.4…,1.}\{0.2,0.4\dots,1.\}. In the case of Virtual, we use an additional damping factor γ∈\gamma\in, frequently used in the case of message passing algorithms (see ), to prevent oscillations. The damping factor acts on the client updates as si(θ)t+1←si(θ)t+1γsi(θ)t1−γs_{i}(\theta)_{t+1}\leftarrow s_{i}(\theta)_{t+1}^{\gamma}s_{i}(\theta)_{t}^{1-\gamma}. In order to retain the same number of hyper-parameters as the baselines FedProx and FedAvg we fix the value of the damping factor γ\gamma as 1−ηs1-\eta_{s} that resulted in good performance of the overall method.