Federated learning enables collaborative machine learning across decentralized mobile devices by having each device train a local model on its own data and only sharing aggregated updates with a central server, while differential privacy adds carefully calibrated noise to these updates to provide formal privacy guarantees, allowing organizations to improve models using distributed data without compromising individual user privacy.
Federated Learning & Differential Privacy for Mobile ML
Added:right yeah I'm Brenda McMahon and I'm gonna talk a little bit about machine learning mobile devices and privacy this I should say this is work from a lot of different people at at Google not just me and Google does a whole bunch of different things to kind of protect user security and privacy and this is just kind of one set of techniques and it's still I would say more on the research side of things so I'm going to talk a little bit more about the research that we've been doing and not really say a lot about actual applications at Google at this point so kind of the goal of the team I work on is imbuing mobile devices with state of the art machine learning systems and we want to do it in a way see if we can do it without centralizing the kind of that core training data and with privacy by by default why is this important all I think as everybody here knows kind of our phones now have access to a wealth of extremely personal information it's extremely useful information from a machine learning point of view or from a product point of view because access to all of the kind of sensors and digital action interactions and everything that goes on your phone I mean you know your location it's almost always with you there's a lot of value in that data a lot of ways it can be used to make the experience of using the phone a lot smarter and a lot better but the flip side of that is that data can be very privacy sensitive so what we'd like to do is kind of have the best of both worlds I mean this is similar to in the previous talk you know you want to get the value from the data without having to to pay a cost in terms of user privacy that works so I'm going to talk about machine learning deep learning in particular so we want to train you know not just only simple logistic regressions but kind of state-of-the-art machine learning models that with non-convex loss functions millions of parameters and kind of complex internal structures and since we want to do this in a distributed in a decentralized way we have to do this on very large kind of decentralized datasets and the the approach we're gonna take here as kind of this federated decentralized and trusted aggregate or who is kind of integrating these individual updates to produce a model for for everybody to share so I'm going to give just a very quick to slide introduction to deep learning essentially what we what we want to do is learn a function from some input to an output so a function from pixels to you know is it the digit five or what digit is it and to do that we're gonna learn we're gonna pick a function out of a parameterised function class so you know all of the these matrices in a multi-layer neural network architecture and we define a loss function that basically says what what do good predictions look like we want and so that defines kind of an optimization problem and so we adjust the parameters to minimize that loss function right that's that's the kind of problem we want to solve here and kind of the standard recipe for doing this and we'll come back to this later is stochastic gradient descent so we pick a random subset of the training data we kind of compute a descent direction that makes the the loss better and then we take a step in the parameter space in that direction so the interesting question is how do we do this in a decentralized and privacy-preserving way so before I get into that I want to kind of step back and say well what's kind of the the kind of cloud centric machine learning workflow for mobile devices look like as it's typically done so we start with current model parameters in the in the cloud and we have some training data and so we train that model based on the data that's a available in the cloud and then we've got a bunch of mobile devices these circles here that come along and they want to use that model to make predictions so I mean you can think about something like a Google search or you know you go to YouTube you get predictions for what movies you want to watch that kind of thing so you send a request to the the cloud we use that model in the cloud to say come up with recommended YouTube movies and that prediction gets sent back to the device and and this works really well and there's actually a really nice positive feedback cycle here because then that that interaction with the service produces training data that we can use to improve it so we will see what YouTube movies actually get watched and that can be used to improve the prediction model going forward and for a lot of cases you know like YouTube movie recommendations this cloud based workflow makes a lot of sense we can't keep the full inventory of all the available YouTube models on a device and you're going to have to screen the device the movie from the cloud anyway so it makes sense to use this kind of cloud centric view for that for that type of problems so we use that that in additional interaction of data to make the to make the model better in many cases though we can do something a little bit different a little bit more mobile centric and that is we can we can move those models to device so instead of keeping the model and in the cloud and having this request prediction cycle we actually distribute the model so we're basically bringing the model to you rather than bringing your request to the cloud and that has a bunch of advantages right now now the service works offline for one because the the model is on device it can save battery it can save bandwidth because we download the model you know when you're on your home Wi-Fi and then you can use it wherever you are it's actually nice for app developers as well because the data is already kind of localized the models right there and it opened up it opens up new product opportunities and this also you know there's a privacy advantage here immediately because now instead of those requests having to go to the cloud none of your data has to leave the device because all of the predictions happen happen locally and in some applications that's quite important what I want to focus on in this talk though is not just making predictions actually kind of bringing model training on to mobile phones and that kind of extends a bunch of these same advantages we can we can take these these advantages even farther if we move model training onto the device so that brings us to what we're calling federated learning so the idea here is we want to this is really a problem not a particular algorithm I should say first and it's this problem of training a shared global model under the coordination of a central server but then we from from a loose Federation of participating moldable devices which each maintain control of their own data so the data is decentralized but we want to kind of enable these devices to collaborate to train a shared model so let's kind of go through the same cartoon and think about what this looks like now we're going to assume that the the training data a kind of each each mobile device has its own little corpus of training data and we start out with a model in the cloud but you can think of this as initially being just randomly initialized now at any given time many devices will be offline and what we're gonna do is we're gonna select just a sample of maybe a hundred of those devices to participate in one round of the training and we want to be a little bit careful about which devices we select we don't want to put any additional burden or impact the the experience of using the phone so we're gonna look for phones that are say connected to a free Wi-Fi network and are plugged in with charge battery so think of your phone you know dreaming wallet while it sits on your nightstand at night so we'll select kind of a small number of devices to participate and then each one of those devices is going to download the current model parameters and now we're gonna run the core training on each device kind of independently on that local cache of training data and then those devices will each contribute kind of a suggested update back to the to the global model and then the the role of the server in the in the cloud here is just to aggregate those updates so for what problems can we apply this there's a few criteria for that makes an application particularly suited to this approach one is that that kind of training data that you have on device is exactly what you want it's it's more relevant than day you already have in the server second you know that that data is either large in in volume it there's a lot of it or it's privacy sensitive so there's a good reason to prefer leaving it on the mobile devices and not having to centralized it and then finally you need to be able to because you're not going to kind of be showing that data to a human and having them label it in some way directly as part of the training process the labels need to be naturally inferred from the user interaction on the device so a couple of examples here that are good for motive a motivation one is language modeling so when you type into your the keyboard on your mobile device it will kind of often predict the next word you're going to type which saves you from typing even if you're not using that next word prediction feature there's a language model that decodes the taps and gestures you make because you're not always perfect about it and so it needs to kind of figure out what word you actually meant so learning those kind of language models are a really good example of where federated learning can make a lot of sense for one thing the on device data is more relevant because people use their mobile devices differently than say somebody writing Wikipedia articles or even somebody writing Gmail right we use a lot more abbreviations we tend to say different things so having a language model that's specialized for mobile keyboard usage is is really valuable on the same time I think it's easy to see that the the data you type on your mobile keyboard can kind of be arbitrarily privacy sensitive writing truths passwords URLs text messages social media pretty much everything you enter you to you're on your mobile device right comes in through the keyboard so we really care about the privacy of that data and here the data is also kind of naturally self labeled right so as you're typing kind of the next word you type is exactly what you're what you're typing and if we don't predict it so we learn what the correct prediction is in each case federated learning presents a bunch of challenges from an optimization perspective I'm not actually going to go through these because I want to make sure I get to the privacy part of it which is the main of the talk which is the main focus today but this was really some of the motivation motivation for for introducing the phrase because this it is in some sense a distributed optimization problem to train these models but the standard assumptions that get made in the majority of the distributed optimization literature just don't apply here we're kind of on on different scales different assumptions about how the data is distributed across the compute nodes and so on so before we can talk about privacy we need to talk about how we can kind of solve this problem without privacy so I want to introduce the algorithm we've been we've been using for this it's actually kind of funny that this we initially implemented this algorithm expecting it to be kind of a baseline because it was an obvious first thing to try and it ended up working well enough that we haven't had to spend a lot of time developing more sophisticated algorithms so this steps through kind of in the the procedure I just described the server is going to select a random subset say a hundred devices to participate in the training in parallel each of those devices downloads the current model parameters and then locally essentially each client is going to run stochastic gradient descent over their local data to compute the update so if you're just gonna kind of run stochastic gradient descent all each client would do would be kind of compute one gradient and send that back to the server but instead what we're gonna do is take multiple local steps of stochastic gradient descent to hopefully compute kind of a better update and the reason that's advantageous is that in federated learning our main kind of constraint is communication bandwidth right you have to have a mobile device download the current model which can be fairly large think you know five megabytes or something for a for a moderately sized mobile model and then upload the resolving update which is about the same size so compared to doing that same kind of computation in the data center you know we're I don't know three or four orders of magnitude slower in terms of the network so it's advantageous to do a little bit more local computation which is relatively cheap once you've downloaded the model in order to get a better update and then hopefully have to make fewer iterations of the overall protocol yeah using second-order methods at all I mean obviously computing the heat like the the Hessian is unrealistic here right using some kind of second order information to reduce the number of rounds yeah we mean you could do that potentially locally just in kind of the inner optimization or globally over all of them so we've certainly thought about it but we haven't tried it yet like I said this has been kind of working well enough that it so far hasn't hasn't been the bottleneck for us and then so once the the client is kind of computed these updates just the change in the model parameter is recommended based on the local data it sends those to the server and the server's job is really easy it just kind of computes a data weighted average of those updates and applies it to the current model so just to quickly give some experimental results here this is a large scale next word prediction problem and compared to a baseline that's already pretty good this kind of large batch federated SGD using this this federated averaging approach we get a 23 X decrease in communication cost so we're able to train a fairly reasonable model in like 35 it rounds of this of this procedure which is which is awesome similar results for C far this is a standard kind of image classification task that neural networks are very good at here compared to kind of a baseline of like data center style SGD again we get like a 50x decrease in the number of rounds of communication we need to get a reasonably accurate model and more importantly kind of the number of rounds we need you know hundreds to thousands to get a reasonable model that's the kind of thing that would let you train a good model in the real world and you know a few days or a weeks even if or less than a week even if each round was taking say five minutes or something like that it varies but but you know one to five megabytes or something you know depending on you know think think a million parameters five million parameters and there's tricks you can play with you know quantizing those models and and compressing the updates in various ways that I'm not good that would decrease the actual size but but think a few million parameters as being as being typical STD is actually you just kind of pretend you had the data in the data center and pick a ignore kind of any user partitioning and just kind of take one step on a mini batch of like 50 examples federated SGD is essentially doing kind of a large batch gradient descent where you take all all of the data from say a hundred users so that would be like essentially essentially like taking an SGD step on a few thousand examples instead of 50 so that actually just using those larger batches if you if you tune the learning rate a little bit actually helps a fair bit and then federated averaging adds in instead of kind of just computing a single gradient on each of those selected users you actually kind of iteratively take a bunch of gradient descent steps before averaging so that's optimizing these kind of models is an inherit of inherently kind of iterative process to get it right and so by kind of we're able to move kind of some of those sequential steps onto a single device I think that's the the best intuition for why this speeds speeds things up alright so now we've kind of got the basic story of how we want to train these models let's let's think about this from a privacy perspective so the the place we need to come back to is where the server is start aggregating these updates this is the first time that any information leaves the client mobile device and goes to a server and so the obvious question is you know do these updates that gets into the server potentially contain privacy sensitive data and the answer is obviously yes they could it's it's less information than if we sent the whole raw data set right substantially less information is from an information theoretic point of view but that's not to say that you couldn't infer quite a bit from a user if you inspected one of these individual updates the details obviously depend heavily on the particular model you're optimizing kind of how hard or easy that that attack would be I think it can range from totally trivial like you know you see a particular coefficient is zero and that means the user used this particular word or this particular phrase for other kinds of models I think it would be quite a bit harder to attack the updates but I'm sure it's it's doable so so there definitely can be sensitive information and these updates but already before it kind of introducing anything new we already have two pretty important things going for us one is that these updates are ephemeral so unlike you know when you log kind of training data if you're gonna train models in the cloud typically you would hold on to that data for some amount of time so you could train different models and use it here we can kind of apply these updates aggravate these updates and apply them to the model and then discard them immediately so we never have to store them and kind of going back to the previous cost the previous talk storing private information actually comes in it with with substantial risks or I should say maybe a substantial responsibility to protect that data so if we don't need to do that it's nice to be able to to discard those updates as soon as we've we've used them second and it's already hinted at this this is very focused collection we're not singing all the raw data we're sitting trying to send really the minimal thing that is sufficient to to improve the current model so that's that's again better from a privacy point of view the next point is that that that aggregation step that the server had to do was very simple it was just a simple data weighted average of the individual updates and that would that was by design so in fact all the server really needs to know is the average of those updates or the sum of the updates it doesn't care about any individual update so this so you know the ideal thing here wouldn't it be great if Google kind of could provably not see the individual updates and only got that aggregate so this gets into secure multi-party computation secure aggregation and we've actually done some work on efficient protocols in this in this space one of the key challenges for doing secure aggregation with mobile devices is the limited availability of the devices the devices can kind of disappear at any time so you need a protocol that not only handles kind of securely aggregating not a handful values but these kind of million dimensional model parameter vectors but you also need a protocol that's robust two devices participating kind of dropping out you still want to be able to recover so you want to be able to select you know a hundred and thirty devices but still complete the protocol if only a hundred of them stick around and then finally we can actually ask a different question say we're already doing secure aggregation we're not worried or we're not worried about Google seeing the individual updates because they're ephemeral and they're they're not stored but we're still kind of training a final model here and that's gonna get sent down to mobile devices so we might worry that the that that model might kind of memorize a particular users data this gets it kind of some of the recent work on membership attacks in model and version attacks and so on and so this is where differential privacy comes into the story the basic idea is instead of just doing the the kind of exact some of the models we're going to add we want to be able to add some noise right that's that's proportional to anyone users update in order to actually get a differential privacy guarantee though we're actually going to need to make a few changes to this federated averaging algorithm some of these are really critical and some of them are more technical to to let us prove the appropriate guarantees the first one is instead of selecting say a batch of exactly a hundred users we want to select users independently with a probability Q so we'll get an expectation say a hundred users but now we actually have a variable size set of users on each round the fact that that we're doing the sampling that way is mostly for technical reasons but because we're adding noise we're actually going to need to select a much larger number of users per round as well because the noises that we need to add is is noise proportional to one users value and so we can kind of decrease the relative effect of that noise on the average by adding more users and that turns out to be a critical degree of freedom we have here the next thing is we need to bound the influence of any one users data on the final average and if one user's data could you know be a billion there's no way to do that so we need to apply a clipping step so we bound the maximum l2 norm of any one users data and there's again a few different ways you can do that but the key is that you if you view it just as a single parameter vector you assure that each user is only contributing a model update that has bounded l2 norm then because we're doing an average here we have to be again for kind of some technical reasons a little careful about how exactly we compute that average we want to do something more like a sum if if there's kind of user data or we're trying to use a private estimate of what's in the denominator we have to worry about exactly how we compute the sensitivity of that function so this is another kind of more technical point and then of course the final change and maybe the most critical one is we have to add the appropriate amount of Gaussian noise to the final model update so there's really four big changes to the algorithm we've we've made here and the question is how are these going to influence those convergence properties of the algorithm the the main technical tool we're using to do this is the moments accountant that lets us kind of track our privacy cost in a very efficient manner over multiple rounds or multiple iterations of the protocol so I'm not going to say too much about the technical details in the talk the reference is there if you want to to look at the details but this this is a key point a key piece of theory that we're building on and it's what necessitates say the random sampling with with probability Q for each user now I want to go back to this side where we said hey we had this this huge decrease in communication costs that means we only had to run this protocol 23 times fewer this also has privacy RAM mutations you know from a privacy point of view this means we only need to query the database the private database this decentralized data set say 35 times instead of 800 times so there's a really nice synergy here that using an algorithm that's more efficient from a communication point of view actually ends up also being very beneficial from a privacy point of view one of the challenges then of making this work as we now have some additional hyper parameters we have to we have to worry about so this is looking at the effect of that clipping parameter and so we can see it's if we're too aggressive with the clipping you can ignore the two different kinds of lines this is just accuracy after different numbers of rounds of training that that clipping can really slow the convergence of the algorithm but we can look at this plot and say hey if we pick a value of like 20 clipping doesn't really hurt so so far we haven't done any privacy or anything we're just in looking at the impact of that one parameter now he can fix say clipping at that level and look at the effective noise we're still kind of using only a small number of users per round and we're looking at just the total the absolute amount of noise were adding and again we see if we add too much noise training kind of completely blows up accuracy goes down to basically what an untrained model would get but if we pick a reasonable amount of noise like something like a standard deviation of point zero zero six in the final average then it doesn't seem to hurt training too much so we can put these ideas together and the key is that those initial experiments get us into a reasonable a reasonable parameter range and so by picking something like a standard deviation of noise about point zero zero three and an ax clipping level of say fifteen we can get a model that's that's very comparable with a non private baseline in terms of accuracy and we can get a pretty good level of differential privacy so four point six an epsilon to four point six for the for the actual data set size one of the key observations though is if we just run this on a bigger data set in actual kind of internet scale data set the privacy is a lot better thanks to amplification via sampling what's what's different than the baseline is for the baseline we were using only a hundred users per round in order to get the noise level small enough to get down to that point zero zero three we would need to use like five thousand users per round so here privacy is not really costing us anything in accuracy we've got a very comparable model but it's actually costing this quite a bit of in terms of computation because we need to run a few more rounds of training because of the noise and we had to use 50 times as many users per round but this is actually kind of nice because computation is getting cheaper and so being able to kind of pay computation and or instead of paying accuracy is really much much better from from the point of view of of being able to deploy this in real systems and I think I'll pretty much in there so I should say this this was differential privacy in the trusted aggregator setting we are interested in moving towards more of a local differential privacy guarantee but that is an in combining kind of differential privacy with secure aggregation and being able to say something about the guarantees that provides but that's that's kind of ongoing work and kind of in conclusion that differential privacy is very complimentary to the to the benefits that federated learning was was already providing so if you want to learn a little bit more we had a blog post on federated learning back in April that talks a little bit about how we're using this technology in G board now not with differential privacy or secure a great but but we do have the system deployed in some real some real products and there's plenty more work to do so I'll leave it there thanks yes that's right that that single communication round I think our protocol is like five kind of rounds inside that yeah it's a fairly complicated protocol so there is a pretty substantial cost in terms of at least the complexity of the protocol to implement secure aggregation that's right Surya how much of the conclusion they had about about it costing only more computation as opposed to accuracy is it do you think that's coming just because of the vast amounts of its yeah right the caveat there is for a sufficiently large data set where large here means large number of users I should say I didn't mention this explicitly one of the advantages of this approach is we're actually getting user level differential privacy not example level because we're bounding the contribution of each user that's kind of the atom of privacy where we're working with but you need a lot of users to make that work the advantage is at Google we tend to have a lot of a lot of users oh yeah I can't remember yeah yeah yeah which is not which is not unreasonable for a lot of but even for a million users like we get we can start to get something reasonable that seems to be where at least for these models you know we could get the noise low enough that's right that's right the other thing that we kind of made things harder for ourselves by training these models from scratch training purely from random initialization for most of these applications I mean like language modeling yeah there's no reason to train from scratch and practice what you would do in practice is you a trained kind of a reasonable language model on public data and then you would just refine it to work on the particular distribution you know to work on mobile usage in particular and so that should let you use kind of quite a bit many fewer rounds of communication even than we needed here that's right that's it that's a really important question for open work I think like right now getting this to work right we had to run quite a few experiments to get the right parameters and of course you have to do a fair bit of parameter attorney to need to get neural network training working in general I think what we have in this paper up on archive now is a reasonable recipe for doing that parameter tuning what I'd really like to see is more adaptive techniques that would make that parameter tuning unnecessary and we started doing some work on that but there's there's still more to do yeah [Music] change that's right that's right whether we're tuning that parameter adaptively and most of the initial experiments we didn't we ran one experiment where we kind of restarted training and and and changed it once and that does seem to help so either using kind of an appropriately chosen schedule for adjusting those those clipping parameters over time seems likely to help of course that means are now and even a much higher dimensional parameter tuning space or maybe even better would be kind of using an adaptive process to kind of measure the actual nor norms of kind of the unclipped gradients and use that to guide that's right yeah exactly so so the the clipping and the noise and the users go together so if you if you use more users per round you can afford to just use a larger clipping parameter so in some sense more users and more computation if you have that if you have the users in the computation solves most of your problems here and I think that's one of the biggest takeaways okay let's take the remaining questions to the break so let's think Brendan again [Applause]
Up Next

Logical Agents and Entailment in Knowledge Bases | AI Tutorial
@fiacobelli
66.9K views•2015-07-14

Secure Multiparty Computation (MPC): Foundations & Challenges
@SimonsInstitute
7.3K views•2015-05-28

Bypassing Tor Censorship: Bridges and Pluggable Transport Guide
@Coding_ForEveryone
397 views•2024-06-11

Neural Networks Explained: Math, Layers, and Learning Fundamentals
@3blue1brown
21.9M views•2017-10-05
Related Study Plans & Knowledge Roadmaps
Structured learning paths in Artificial Intelligence























![FedAVG, SCAFFOLD, FedDyn [20210405, Nam Hyeon-Woo]](https://i.ytimg.com/vi/7hMIInjmKiU/sddefault.jpg?sqp=-oaymwEmCIAFEOAD8quKqQMa8AEB-AHuDoACmgiKAgwIABABGGUgZShlMA8=&rs=AOn4CLD6JyA6tkD8L5QhY8qNWBvUNL4hdA)















