Federated learning is a distributed training approach that enables training deep neural networks on continuously generated user data across millions of devices while preserving data privacy, addressing key challenges including device heterogeneity, non-IID data distributions, unreliable network connections, and communication efficiency through algorithms like Federated Averaging and Federated Proximal Averaging (FedProx).
Federated Learning: Train Models Across Devices Without Sharing Data
Added:today we'll talk about a very interesting topic um and i'm not an expert on this topic and that's partially why i found it interesting to prepare the lectures because i learned a lot about it as well but it leads on nicely from distributed training it's taking the idea of distributed training and scaling it to many more devices and many more users and so um yeah we're in this you know last course module uh ml systems we talked about you know pre and post processing how that affects you know the speed of my system and what i should watch out for uh the key thing there is still amdahl's law that people often forget about until they deploy their system and they kind of profile it for performance then we talked about distributed training we talked about all reduce and how that doesn't really scale because the communication cost increases linearly with the number of gpus you're using for training we talked about parameter server and then the communication cost was again limited by the parameter server communication bandwidth itself and finally we talked about all reduce we also talked about model parallelism and other things so today we'll talk about federated learning which is a form of distributed training really but quite different and has different requirements and thanks to a good friend stefanos lascaridis so he's he's an expert on federated learning and i consulted him a lot before kind of putting these slides together so hopefully they look good yeah so that's kind of what we what we learned last time we talked about distributed training and we focused on the communication costs and how we can scale that in a way that's independent of the number of machines so that we're able to you know have a lot of gpus training at the same time without burdening the communication network too much um so so so i mean federated learning is not that different it's still a distributed training problem uh but it has its own requirements so in in conventional um deep neural networks um what we do first is that you know we gather a data set somehow um so you know the the standard example is imagenet you know we have a million images we take them all we label them by the image class and then we have this data set that we can use for training um and and so that's nice because you know we can we can easily set up this supervised training problem um we have the data set that we can try training on once and then try again and so on um but the fact of the matter is and reality is that you know data's being generated every second that's the whole idea of big data so we have all of these devices with fancy sensors on them capturing data as we speak so how do we take advantage of this data to train a model and make it smarter right so you know you have your phone you just took a photo i want to add that to my own imagenet which google is doing in some way because they have this huge um i think that's 300 million images now called jft like a big benchmark and you know if you have a drone and you want to make it kind of avoid obstacles better you want to gather that data from all the drones out there or if you have you know these fancy augmented reality goggles then again you want to be able to train whatever algorithm on it using the data that people are seeing right now and the same is true for self-driving cars especially like the main you know value proposition of the company called tesla is that you know they have all this data and they have all these cards gathering data all the time they're all connected to a central server where that data is being used to kind of improve the model for self-driving every second so that's that's what we want to achieve we want to go from this conventional view of there's a data set that we want to do really well on to a much more realistic setting where you don't have just a single central dataset stored somewhere that you bought once and that's it now you have continuously generated data that you want to continuously use to improve your model um so so that's kind of our question in this lecture so how do we leverage uh user generated data to train deep neural network so so that's what we want to do and so um in in the first setting that we'll talk about and what we'll talk about throughout this lecture actually is this federated setting where all of these devices are with the user and then you want to train a deep neural network that's stored centrally somewhere there are other cases as well where the neural network is not stored centrally but this is kind of the scenario that we will be talking about so i had to actually go look up what the definition of federated means um and it basically means distributed but reports to one central entity so in case your english is as bad as mine this is what it means so it's like you know federal governments for example so you have you know all of the state or provincial i guess state here it's called state uh government but they all still report to one central federal government and so the same is true here you have all of these devices and they are independent in some way they don't really belong to that central entity but they still report to it in some way right and so so that's what we want to do we want to gather data from all of these devices and they still have models on them that we can train on device but we want to improve a central model at the server somewhere and that's that's what we'll do through federated learning okay any questions about this short intro of what federated learning is or what the difference is to distributed training so far so so you don't actually know what the difference is to distribute training so that's one question but we will say we will kind of go over that right now so at the at the highest level um you know federated learning has kind of four steps and they look a lot like normal training but let's go through these steps one by one so first you know we have the central server that we talked about and let's just say this is google right and so they have a model that they trained here with a lot of data that they bought that they gathered before they curated and put somewhere and so that blue model the first step is you know i want to give access to that model to my users so we're downloading that model from the cloud to the user's devices and so now every uh every user has the same model it's identical it's a blue model and it's downloaded on my phone and i can use for inference like for detecting faces or classifying audio or doing whatever then the next step is you know every device is now independent it's not connected to the cloud it can operate offline and everything so you know each device is also gathering data as uh all the time so so device one is gathering different data from device two so maybe the person owning this phone really likes to take photos of you know landscapes but the person owning this phone really likes to take photos of their dog for example and so this model will be really good now at detecting dogs while the other one will be much better at detecting landscapes so we can train on device uh to improve each model um given the data that's available on that device so we don't have access to the global data pool anymore we just have access to our local model and the data that's on that device and independently this step step two is sometimes called personalization so you often take kind of a baseline model and you perform either transfer learning or top-up training to just personalize that model and make it really good at that specific user and people do that often for for example speech recognition and definitely for photo classification image classification and things like that as well but federated learning doesn't stop here so now that we've kind of specialized each model to each device and we've done on-device training which by the way is the first time we talk about training a model on the user's device so we'll get back to that but once we have these different models now on each device we want to find a way to get that knowledge or that kind of information and that data that's distilled into these models back to the central model to improve it as well right so we need to find a way to propagate you know the information that we learned here um back to the central server and so how can we do that so what's what's the simplest way of doing that so i mean i can i can just send the data back right and then we can train on the server that's one option are there other ways of doing it upload the gradient so similar to distributed training was that also what you were going to say or yeah okay so uh so do something more fancy so training on the device itself and then use that information somehow and propagate it to the server you said figure out which one is highest accuracy um that's i mean that's one way but remember we each device doesn't have like all of the training data it doesn't necessarily improve the model uh and performance of detecting for example all of the classes in a classification problem we're detecting all of the words in a speech recognition problem so every device is kind of specialized now so there's information to be gained from each one in some sense so these are all great answers so james do you have another one yeah that's that's a good question so so i did mention that step two alone can be called specialized specialization or personalization and people definitely do that for you know your keyboard app because you have a specific way of writing kind of text messages definitely for image classification and so on but that's only one part of these four steps that we call federated learning so for the sake of personalization alone we may actually not need to propagate anything back to the central server but our goal here actually is that we're google and we want to improve that model right we we're not the users in the other end uh eventually kind of the better model gets downloaded to the users again and so on um but but um [Music] the goal is not just personalization personalization is part of it i would say okay so um yeah so these are all great answers and we'll see in the next slide exactly what we do to make uh to kind of propagate that information back to the cloud and then finally we update the global model so we we get that information that we learned whether it's you know gradients whether it's updated parameters whether it's data and we use it to train that or to continue training that new model and and these four steps are actually not just four steps they continue in a loop um and you know this loop in normal machine learning or conventional machine learning you would say these are different epochs of training but here they call them rounds of training so every time kind of you do this whole thing it's one round of training and you you do it multiple times so so you continuously train as new data is being generated and so on so before we dive into the specifics um of how it's done let's kind of discuss quickly why or what the use cases are um so one one very important use case um is you know using health data to train models and so machine learning is very successful in many problems in medicine um but but it does face a lot of problems one one key problem is that you don't have access to all of the data there is such a thing as user privacy and you know not all patients want to disclose their you know private information and so on and not all laws actually allow you to disclose that private information and so uh so so training a model based on health data and health diagnostics is one key um one key usage of federated learning um another one which is i think the obvious one is kind of just personal data from your personal device so basically look at your phone any data that's being gathered about you from that phone especially offline data so things like you know your your imu data so how the phone moves and when when when for example you want to detect a fault you want to be better at detecting falls which apple does i think now with their smart watches and stuff so that's an example of kind of imu data which is really to health as well fitness data so maybe you know you get the user to put in his weight after doing a specific exercise and then they can better correlate you know the weight loss for a specific person after you know working out you know running 500 meters or whatever um or 5k or something against speech so just figuring out you know how to detect speech and transcribe it into text or something like that so all of this personal data is very amenable to this idea of federated learning just because it's being continuously generated and you don't necessarily have anything to gain from just uploading your data somewhere to the cloud right or giving access to people then finally um you know we talked about self-driving cars and sensor data from smart home or a smart city again you know you have all of these sensors dotted around the city now so things like at traffic lights detecting pedestrians in in crosswalks or something like that so you have a lot of data being continuously generated and you want to use that data to improve your models can people think of other use cases actually that are not listed here any anything else that would make sense yeah go ahead weather data so yeah i mean if you're collecting it from user owned devices or something then yeah that would be a good example as well oh i see like they put them on the roof or something uh and they detect wind yeah i think so i don't think the problem is as big um in terms of just problem size the number of users that have but yeah that could be an example as well [Music] sorry i can't really hear you industrial iot systems so yeah so iot in general like the word iot goes really well with federated learning because iot implies lots of sensors scattered around a large space and so um that would be a good use case as well so um okay all right so yeah so at this point it's good to kind of just you know um to say exactly you know what is the difference to distributed training to have kind of a mental model of the problem scale and so the first thing to keep in mind is that the devices in something like distributed training are identical typically identical or even if there are variations they're very small variations like you know 1080 versus 2080 gpu or something have more cuda cores or less memory but with federated learning you can have like i don't know a nokia old nokia phone versus you know the latest iphone or even different generations of iphone or samsung phones or something and so the device type here really varies it could be a completely different device like we said you know augmented reality glasses versus a phone or a smart watch or something and so what does how does that affect me so it affects me because we're doing on-device training and so we need to actually for this system to work we actually need to find a way to train the model on each of these devices and these different devices also have different communication links some devices are only connected to wi-fi for example some devices would be connected over a 3g or 4g or 5g network so it depends on how often you're also connected to the network and then a really big one is the number of devices so in distributed training you'd be lucky if you're training using tens or gpus or hundreds of gpus if you're really rich but and again if you're like a huge company then maybe thousands of gpus for a specific training run but for federated learning we're talking you know millions to billions of devices so so again the the problem scale is much larger and this really dictates many of the training algorithms that we'll see in the next few slides the network connection again in data centers we have these currently 40 gig 40 gigabits per second links between server racks so it's a really fast very reliable network you don't drop packets really on that network it's a bunch of cables unless something really goes wrong whereas with federated learning you know you have your phone connected to a mobile base station and sometimes your phone is off or the battery is finished or you put it on airplane mode or you're underground or something so you don't have a reliable connection so you can't kind of really rely on that network as you would in distributed training the data so that's another another big difference actually all of these are big differences between the two things but this is also a very important one in distributed training um you you pay some attention to having kind of a representative mini batch on each device so if it's a classification problem again you make sure you have a mix of classes being trained in each mini batch and so the data coming from each device or being trained on each gpu and distributed training is almost the same it's still representative it has the same data distribution that's what we mean by kind of similar data so the training problem is almost the same on each device whereas here we talked about personalization and really you know the data is not um is not identical in any way across these devices so people call that non-iid independently and identically distributed so this distribution of data is different so the data itself is different that's you know for sure because we're training on different devices that's fine but you don't have the same mix of classes if you have 100 classes you may have you know 99 samples of one class on my device so in the case of images it's probably like photos of my son have 99 of those and then one photo of i don't know a building or something in another person's phone it may be you know 99 photos of a dog and one photo of himself or herself so so the data is very skewed um and and this presents a major challenge actually and we'll talk about how we can fix that finally the design goals so we we didn't we didn't kind of say this explicitly yet but you know a common theme across all of these federated learning use cases is that the data needs to stay private um so we don't want to actually transmit the data itself so we want to make sure the data stays secure and private and we want to be able to scale that system to millions or billions of devices whereas in distributed training our design goal is kind of just speed how fast can we do this obviously accuracy is a design goal in both but you know where it differs is you know speed versus scale and privacy i would say okay any comments about this part or or anything let's uh that's kind of we should discuss more so yeah so i mean to sum up you know we want to be able to train a model using data that's being continuously generated without having access to that data and being able to scale it to millions or billions of devices so that's that's kind of where federated learning is and so now it's starting to look you know much different compared to distributed training yeah so that's a great question so how often do we need to communicate in the case of federated learning versus distributed training i mean the first hint is that you know we have a much worse network uh to be able to communicate over and that question will be answered in a couple of slides so maybe we can move forward and before we talk about that we i should point out at this stage that you know we what i implied is that we want to do on-device training um we haven't actually seen that yet so in fact like in your assignment 2 for example when you had a small device you could barely perform inference on that device with a tiny model called tinycon right like it was the smallest model we could ever come up with it did a very simple task and we just did inference and you know you profiled it it took a long time to do that inference in terms of milliseconds at least so when we compare kind of you know the devices that we usually use for training versus the devices that we're talking about now like the most powerful kind of device in the federated learning setting will look like this to be kind of you know an apple chip with an npu about three or five watts of power which is two orders of magnitude less than kind of an nvidia gpu which is typically used for training so keep in mind that we can't actually do much training on that device right so so far you will see if you read about federated learning you'll see that it's used for things like you know making the google keyboard better or keyword detection on alexa or or the google home or something like that so it's actually still being performed for very simple tasks just because on-device training is not there yet we don't have uh hardware that's powerful enough to perform on-device training as of now so we can we can do inference fairly easily on mobile phones but not training basically and so yeah as i said one one one thing is to keep in mind is that you know we use that for simple tasks for uh low dimensional data things like health data is quite low dimensional it's kind of you know it's it's over time but it's it's a specific thing that you're measuring so whether it's you know steps or speed or or something like that and a few ways of kind of making it slightly better is that you know some some of these algorithms actually kick in when you're charging your phone so when you're connected to power basically um and again to kind of circumvent some of the communication uh problems uh you know you don't want to also be paying the bill for you know 5g especially if you have kind of a limited plan or something so sometimes you know these data transmissions only happen when wi-fi is available and so one of the readings is about the google keyboard thing that i mentioned and that's what they do basically they wait until you're connected to wi-fi you're plugged in they do the top-up training and they can they send something back to the server um so so just keep that in mind uh if you're thinking about federated learning for a problem um think of your device and whether if we want to do on-device training how much training can that device actually support right i will also say that you know many of these socs on your phone for example um do have dedicated npus but i haven't seen training libraries for these npus so currently unless you're writing your own low-level code you can't just say you know train this model or model.backwards you don't have that command yet on these devices so people would often write their own kernel or they would run it on the cpu of that device in the case of google they did deploy a small version of tensorflow to their phones and it's actually different than tf lite i went and checked so it's it's another thing that enables on-device training for very simple models um so they basically have to write their own code and it's not a public library yet okay and yeah and we touched upon this a bit but you know if you have different devices you'll have different data so this idea of non-iid the the data is not identically distributed this is really important and you know here are some examples of images they're all from my phone actually but i would say 99 of my images look like this or it's my son doing something and you know the rest is just random random stuff there is probably a disproportionate number of cows in my phone but just just based on where i've been over the past few years uh you can ask me about that later but um yeah but you'll find kind of um mostly just kid photos of my phone someone else may have more kind of building photos or food photos i know um i know someone who takes a lot of photos of food so i'm sure his data will be very different than mine so text data is also quite different even though even if we're speaking the same language or using the same alphabet we may say things in different ways right so these are all you know ways of saying the same thing in different languages or just dialects and so you want to learn from all of that different data sound so for things like noise cancellation you're looking at background noise and people have different background noise obviously if you're in a city it's different than from when you're in a countryside people speak different languages as i mentioned and different voices and you want to learn from all of that data so so yeah so i'm i'm repeating myself at this point but it's very important to realize that this is a main challenge of federated learning you know the data volume the number of photos and the data type is different from each device and not representative of the global data that's being used to train that model so yep yeah so that's definitely a concern so devices will transmit at different speeds and at different times and so we will see something called stragglers basically devices that can't make it in time for the next round of training so we will talk about that again okay so at this point it's maybe worthwhile to do kind of a quick discussion of you know why can't we just use distributed training so distributed training is where we propagate back the gradients and we update the global model actually can we use that for federated learning and if not then why not so so the network so that's definitely you know one thing so in when you're propagating oh that's still my text okay um so the first thing is network uh with distributed training you're propagating gradients every mini batch of data so that's a lot of communication we want to kind of make that less in federated learning because your network is not as good okay what else um i see so your gradients become stale because oh i see so it's you don't uh you're not operating on the latest model anymore basically that's what you mean have gradients trained on an older version of the model for example yeah so that's true if we're doing the updates asynchronously which is probably what we would need to do in this case so that's a good point again it's related to you know this idea of a network and you know i talked about gradients and and three mentioned gradients kind of being uploaded to that central server but then we have to download the average gradients as well to update our local model and so distributed training would require a lot of network so funny so yeah the input data is highly imbalanced so even if we actually do distributed training using that input data it might not converge properly so we maybe we need to change stochastic gradient descent in some way to make it scale to this much larger models so that's number two update uh imbalance this is very bad it's too bad it's a bad marker i don't have an eraser it's not good okay whatever okay what else what else do we have that's that's an issue why can't we just do distributed training so number of devices okay so that's that's a great point because i mean a number of devices would also kind of go over here um you know do we have enough bandwidth to actually receive all of those gradients at once in a central server so we talked about how in the parameter server thing it scales linearly with the number of devices so you know we probably can't handle it if we have a million devices sending you know gradient updates to one server and then it needs to broadcast the model back that's great so so i mean privacy is also an issue and you know someone is thinking if i'm propagating gradients which is what i do in distributed training why is privacy an issue it's no longer the data right but turns out people can somewhat easily now extract the data back from gradients and so that's um yeah there are many papers about it people can do it and so just transmitting the gradients is not secure anymore so you can't do that what else are we missing anything so network data imbalance and privacy no yeah i think that's it actually and yeah the main one that i wrote down here at least is communication cost but forget the slide because we have a better summary here okay so now let's kind of uh now that we understand what federated learning is and why it's different from distributed training let's look at the first algorithm that was proposed in this space called federated averaging just one second so this is one of your readings and this is actually i believe the first paper to kind of look at federated learning um and to kind of attempt to do this on a large scale can someone guess from which company it came from yeah from google so i mean one thing you will notice um is that at least half the paper is about federated learning or from google um and one main reason is that there aren't many researchers that have access to millions of devices right or millions of users and so doing realistic research on this work is actually very hard in my opinion unless you're taking a sub problem that you can test in the lab so so how does this algorithm work so remember this picture where we had four steps you know we have a global model we download it we train it on local data we update something to the global model and improve it and so on and we keep doing that again and again so here is what it looks like in this algorithm called federated averaging so we have the baseline model and in this algorithm here we're initializing it so it could be randomly initialized or it could be pre-trained using some large data set and then we repeat the following training loop for t rounds and like i said you can think of rounds as epochs basically in the equivalent of epochs in normal training however these rounds usually take days just because of the communication overhead and waiting for all of the devices to come back and so on as opposed to epochs which could take hours or minutes depending on the data size and so first what you do using those two statements here is that you sample some devices so you have millions of devices out there maybe let's say one million as an example you sample some of them and realistically you usually sample thousands so between five and ten thousand i think is what i've seen in real deployments uh in each round so you sample one thousand of them you wait until a thousand of them respond back with something before you perform the next step and so for each of those clients you're training the model on each device and updating its weights so we're doing normal training on the device using all of the data available in the device and we're getting a new set of weights w t plus one and this is you know the updated model parameters on each device so at this stage we have each device with a you know updated model different you know represented by a different color here and it's trained on the device itself and then what we do is that you know each of those clients will send their updated parameters so after you know doing forward pass backward pass wait update doing it multiple times we have an updated model we want to send that updated model to the central model or the central server right and then what we do on the central server is that we basically average all the weights together so it's much simpler than than doing anything else like we just take all of the weights coming from all of the clients and we average them we also perform a weighted average depending on the number of training samples that was available on each device right so let's say you know the first device has a hundred images the second device had only two so we put that factor here and we linearly scale each one of them as a fraction of the whole updated data set right and then we use that to update the model and we perform another round and another round and so on so it's the concept is quite simple right so we um we're training locally on each device uh to minimize communication costs we're not propagating the gradients or sending anything after each mini batch on the device now we're waiting until we finish training with the whole data before we send it back so um below here it shows the client update and you know we're training for multiple batches um just doing normal training for multiple epochs and for all of the data basically so we're waiting until each model converges on each device before we send any update back and then at the at the server we're just averaging all the weights together and we're seeing if we get anything meaningful out of it and so just by doing that by having each of the devices perform training for multiple epochs locally we're reducing you know communication by between 10 and 100 x so uh instead of you know communicating every mini batch we're now not communicating even every epoch we're communicating every training ground uh so to speak so it's a simple algorithm but it works to some extent and we'll see what the problems with it are and it ad it really addresses this communication issue at the first order it reduces the communication by orders of magnitude it also addresses the data imbalance a bit by adding this scaling factor which enables the training so instead of weighing all of the weight updates the same way depending on the number of samples we will apply that weight right um [Music] yeah so so yeah so this is kind of the baseline now this is how we can do this federated learning and scale to millions of devices while you know relying on only kind of very weak communication links so any questions about fed average okay so there are some issues though right so the first issue is what harris was um was alluding to like stragglers so what if you know i'm doing this training i need all of the data to return at some point to do that averaging on the server update my model and send the model back to the clients and so on um but what if the device doesn't return the data on time right um between you know those devices that i sampled so what fed average does in that case is that it simply drops the stragglers so if there are any stragglers from these 5000 or 10 000 devices that i sampled it just and doesn't come back in time i just drop them and go to the next round of training right and what if there's a huge data imbalance right so what does that say so we said you know we scale each model uh update linearly based on the data in balance but basically is this really representative anymore because you know data from one device may be much fewer samples but it's still as important and so there are questions of fairness and kind of equality that arise here in terms of at least model training itself um but we um you know we we're currently favoring devices that generate more data it's not necessarily better data it's just more security risks like what oh oh i see so like maliciously you're trying to break the central model oh yeah people have done that and people have written papers about that as well so they usually call it adversarial attacks or something so basically finding a way to you know inject some data into the system to break the global model and people have done this in various ways actually one interesting way is that you know if you take a sign of a speed limit i don't know if you read any articles about this but if it says 50 for example and you know a self-driving car would go next to the sign and it would read it as 50.
but if you just change three or four pixels here the self-driving car would read it at 80 or 90 or something like that and you know it didn't really change the actual numbers but it changed the image enough to exploit a weakness in the model that would make it kind of um um i'd incorrectly classify or incorrectly detect this text so this is an example of an adversarial attack they call it um adversarial attack and it's actually much so that's fairly complicated or complex and people do it based on looking at the gradients and how training works and exploiting various you know local minima in the model so to speak but with federated learning it's actually much easier like you said you can just you know overwhelm the model with uh photos of a specific thing that you want to skew the model towards and yeah so people have done that and that's one weakness of the model basically yeah yeah you can so uh so people call that adversarial training and so you basically you reproduce what that versatile attack could have done and you added your data augmentations and train it to the right label so yeah maybe that's actually a good topic for one of those extra classes like adversarial attacks security so so yeah so these two things are major kind of drawbacks of federated averaging accuracy is also one thing like we don't actually uh yeah like if we compare this federated averaging to normal training if we took all of that data and just trained with it directly we are losing a lot of accuracy by doing this federated averaging i don't have the exact numbers but maybe when you read the paper you'll find them so so what can we do about that so so one thing that people have done is you know the very second algorithm that people invented there from the same company is called fedprocs a proximal proximal federated learning or something like that and this is an example of a really incremental change a very small change to the algorithm but had a profound effect on on kind of all the data being learned and so the first thing is that this algorithm doesn't drop stragglers so instead of dropping stragglers it would just take you know a snapshot of the model if the if the device didn't finish training if it's too slow or something it would just say okay give me what you have right now so that improves fairness a bit and increases the data pool a little bit compared to kind of just dropping the stragglers and not using their data at all so using partial results and then a second thing it does to to prevent models from dominating one of those rounds of weight update is that we limit the delta uh of of the parameters we limit the change in parameters on each device and so we want to discourage large weight updates so as i said you know we update the model wt plus one something you know after training of wt so we want to have those previous and new parameters close to each other we don't want them to differ too much because if they're close to each other that means they're also close to the global model and that also means that all the numbers won't kind of skew away too much and so the way they do that is that you know instead of just minimizing your standard loss so f here just represents you know your standard loss function they add this extra term in your loss function which is you know just the difference um the square differences of all of the parameters so basically during training on each device i'm trying to have my weights not changed by too much at least in terms of magnitude and so when i propagate the updated weights to the central server even if i trained on a bunch of parameters or a bunch of samples on this device it didn't actually change the weight values too much so when i average them they're all still close to that original weight value and so it kind of puts all the weights in the same scale so to speak so these are the two simple modifications to fed averaging to address particularly those two drawbacks um that were presented here and so uh so fed prox actually works a lot better than fed averaging and we'll see i do have plots for those i will say at this point that there are other algorithms in this space there's one called q fed average there's one called perfed average and there are many other algorithms that try to improve on this but i would say that this is kind of one of the biggest increments even though the changes are very small okay so so is fed procs now clear the two changes it does to address you know stragglers and to address just large weight updates from a specific device makes sense so you will see kind of plots like these in the paper where they're trying it on different data sets across many devices and in the first row it's showing the training loss so lower is better with zero percent stragglers and pink here is fed procs and blue is fed prox with this regularization term to limit the weight and then orange is fed averaging so with zero percent stragglers this pink line so fed procs with mu equals zero remember mu is just this weight factor next to the regularization term right it's the exact same thing so that's why you won't see the orange line the first line in the first row because it's all overlapping with this pink line but when we added that regularization term it actually made training a lot smoother and a lot better compared to that term not being there because it limits weight updates it limits kind of drastic changes in any direction at any given time and so it made training converge much better and we'll see that throughout kind of all of these plots as well but then when we look at an increased number of stragglers federated averaging doesn't really work well because at least compared to fed procs which takes some partial data and the reason for that is that you know you just have less data to train with you have a poorer representation of the model because again each device may be focused on a specific class for example or a specific type of data so your data becomes really skewed when you're just simply dropping stragglers and here we see with like a really high number of stragglers which is that's i don't know if it's realistic or not maybe we should um kind of consult one of the systems papers related to federated learning but you know you definitely don't necessarily have the data from all devices and so this is kind of an extreme case where you don't have it from 90 of the devices and you're dropping all of them in your training um so yeah so in this case you know blue is always better than orange and so that's what fed prox did and as i said it's very small change but had a big effect on training okay so um so this is kind of the technical part of the lecture you know federated averaging and proximal federated averaging again and so these are two of the standard ways that people do this now and like i said i think fed procs is being used much more often just because it's much simpler but works much better and there are other algorithms there and so a couple of quick notes about you know privacy and security you know with federated learning the data still stays on your device which is a good thing if you're worried about that data being stolen or being hacked into or something it could still be hacked into your device but at least it's not you still own it and you can encrypt it however you want or something and you have the keys [Music] and as i said you know if you're just propagating gradients there are ways of recovering the data from these gradients and you know the latest papers in here actually also recover the data from um you know federated averaging weight updates so just from model updates if people can intercept those model updates being transmitted in an insecure way they can actually still recover some information about the user and so if you're interested in that topic i would recommend you go read that paper it's quite interesting so basically it's saying you know the current way we're doing federated learning is not private don't don't think it is uh there are ways to attack it um and so you also see a lot of kind of research from the security side of things where um you know instead of um you know instead of just transmitting these weight updates in an unencrypted way saying all their weight updates no one can understand them instead people actually encrypt them so you have you don't have data you have weight updates based on data and then you encrypt those as well and then a step further is that you only decrypt that data on the server after averaging that data in an encrypting in an encrypted way with a thousand other updates so uh so what does that mean remember you know our federated averaging algorithm would take all of the weight updates coming in and it would average it all together um so instead of you know uh like decrypting the data and then averaging there are ways of actually performing arithmetic operations on encrypted data the specific kind of encryption called homomorphic encryption and then we would only decrypt the data after we average a thousand of them together so that makes it even harder for people to steal the data and understand kind of the data coming from each user so there's an interesting paper again from google about this but practical secure aggregation for privacy preserving machine learning so that's quite interesting so um so what i'm saying is that you know as is as it is always the case with these systems you know don't just assume it becomes secure because you're obfuscating or withholding the data in some way there are ways to attack this even in this really twisted cases of you know average model updates that are being sent so and um yeah and this published stuff anyone can use it so it's quite easy to do right now there is a lot of research that actually also looks into minimizing the communication costs [Music] so even though we have those much more sparse weight updates as opposed to you know gradient updates every mini batch with distributed training we still want to optimize those and minimize the load on the network and battery usage on users phone and so on but i don't have much to say there because the research that's published in this area is all kind of the standard stuff you know compress your model updates before you send them quantize them to a lower precision prune away all the zeros and all of that stuff so it's it's all standard model compression methodologies that people use to minimize the communication bandwidth and then um one thing that you know at this stage when i was preparing the slides you know what what happens um with all of your label data how do we get the labels so i talked about you know how we get the data from users devices but how can we actually get the labels i mean most machine learning right now is done in a supervised way and so if i'm training a supervised model a model using supervised training then how do i get the labels so because usually with training you have you know data and labels as well so how do i get those labels to be able to train on device any ideas i mean some of the ideas are already written here and so you know in some cases you don't need labels right so unsupervised learning is increasing and it's it's generally tipped to be kind of the future way of doing deep learning just because data labeling is so error-prone to begin with right but also another thing is that you know these companies are really trying to incentivize root users to sit down and label their data somehow so um yeah so there are some problems which don't even need an explicit data label like next word prediction or something like that or when you're you know when you're speaking to your phone and you wanted to transcribe your text in a text message or something if you see something wrong you will correct it so they'll capture that and they'll use that to train their algorithms so that's easy but in some cases like google photos for example i use google photos and it's always asking me like who was with you in that picture or were you in switzerland in that other photo or something like that or did we mislabel this if i search by a category or something so all of these things are being kind of fed back into their model to be able to do that training so incentivizing users or using unsupervised approaches is what people do there okay so we set out to answer this question you know how do we leverage you know continuously generated user data to train a central uh single deep neural network model uh so by now you know we looked at how to do that using federated learning we talked about the implications of of the setting given you know privacy issues scalability issues communication bandwidth and so on and we talked about you know two algorithms which successfully um were able to actually perform that training uh federated averaging and federated proximal federated something fed procs that's called call it fitprox so at this stage it's important to also kind of put this in context of two other approaches and these are the two last slides where you can still do this kind of thing you can still train a central model um in in kind of from from data generated on user devices so the first one uh maybe just worth mentioning again is just offloading right so you know as soon as you have some data you send it back to the server training happens on the server and the reason this is good is that it kind of avoids all of the issues with federated learning like for example many of these devices as i said they can't do training on large problems and so in this case you still have to kind of transmit your data back and in many cases now you find that your data actually sits on the clouds somehow like another example is google photos where they really make it easy for you to back up your data to the cloud for example or apple photos or whatever it may be so so in many cases the data isn't actually on your device even though you can access it from your device and so it's not actually that complicated to just you know do training with it on the cloud of course the main drawbacks here um yeah what are the main drawbacks main drawbacks is that you don't own your data anymore you're giving access for someone else to use it and see it and you know debug with it and everything so that's one of the main problems scalability is also another issue but we can we can typically deal with a lot of these scalability issues with stuff we learned in the last lecture basically buying more gpus and scaling the problem of training to more gpus so so this is still kind of i would say the main way large complicated models are being trained even though federated learning on paper is very scalable the devices are still very weak so they can't actually do much training on device [Music] another advantage is potentially you know you don't put too much load on the device so you know the battery life um and you know the hardware cost of just supporting training you don't need to do that anymore and then another emerging approach which i didn't find much about actually but it's called collaborative machine learning and i think this is this term will be used um quite generally for federated approaches as well and so if you're doing a literature search this may not be the best keyword um but there are a couple of papers just uh finding ways of training a model in a distributed way without a central server so we're no longer updating a central you know model stored somewhere in the cloud we are propagating data between devices and incrementally improving the model on each device of course this is an even more challenging you know setting because there is no one device that has a global view of all of the model updates anymore or all of the data and each model needs to stay intact it needs to converge after you know doing this top-up training with new updates and so on but i think it would be cool if people got this to work basically you know it has all of the challenges of federated learning uh plus and then i forgot to continue this bullet point but basically you you don't have this global view you no longer have a bunch of model updates that you can average together um and so i think this is kind of a good goal to have um but it's um it's currently not mature enough from what i've seen yeah so so i mean what this reminds me of is just distributed training right so if we want to have an all-to-all communication it will never be scalable and so people have looking into this again ring topology where each device kind of sends and receives from one other device and we can propagate updates that way so um so i saw blockchain in one other person's kind of topic proposal as well um so blockchain uh i'm not the person to ask about it so many other people who can answer this but basically it's a way of proving ownership right and so if you want to monetize your data in some way or um or or trace it back to where it was created then the blockchain can be useful in a federated setting and there is a paper if you guys go through the kind of the um yeah the the paper category on ed so gia ga right so gia posted one paper that kind of mixes you know blockchain and federated learning i haven't tried it yet but it looks super interesting so maybe that's a good way to kind of figure out where blockchain would fit in but i don't think it would necessarily help the training problem itself it would just allow you to trace back kind of where the data was generated or where or where it has been or something like that so yeah so so yeah i uh i think this is one alternative to federated learning that could be quite interesting to watch out for uh but currently nowhere near um development ready or production ready rather okay so what did we talk about today we talked about federated learning uh it's a method to train a model using millions of devices we still have the central server where we receive updates and average them um we talked about fed average and fed procs [Music] and you know keep in mind this thing where you know not many people have actually access to these kind of to this kind of scale so there isn't much realistic research unfortunately um especially when it comes to the variations or the real world variations of multiple devices and then you know this collaborative ml idea is an emerging paradigm which could even get rid of that idea of a central server similar to how all reduced ring all reduce kind of got rid of the parameter server in my mind at least here's some further reading i recommend especially reading the first one because it's very easy to go through it's just a blog post but these are great representative papers about federated learning i would say and that's it i'll stick around if there are any questions
Up Next

Decentralized AI: Data, Governance, and Edge Personalization
@vanishinggradients
184 views•2025-07-29

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
























![[Webinar] Level Up with Privacy Enhancing Technologies - 3 February 2026](https://i.ytimg.com/vi/AMLl7ZzuxXo/maxresdefault.jpg)
![[deep learning] Federated Learning - training on decentralized data](https://i.ytimg.com/vi_webp/KxZXhzfDgik/maxresdefault.webp)













