Federated Learning (FL) addresses the need to create models based on proprietary data in such a way that multiple clients retain exclusive control over their data, while all benefit from improved model accuracy due to pooled resources. Recently proposed Neural Graphical Models (NGMs) are Probabilistic Graphical Models that utilize the expressive power of neural networks to learn complex non-linear dependencies between the input features. They learn to capture the underlying data distribution and have efficient algorithms for inference and sampling. We develop a FL framework which maintains a global NGM model that learns the averaged information from the local NGM models while the training data is kept within the client’s environment. Our design, FedNGM, avoids the pitfalls and shortcomings of neuron matching frameworks like Federated Matched Averaging that suffers from model parameter explosion. Our global model size doesn’t grow with the number or diversity of clients. In the cases where clients have local variables that are not part of the combined global distribution, we propose a Stitching algorithm, which personalizes the global NGM model by merging additional variables using the client’s data. FedNGM is robust to data heterogeneity, large number of participants, and limited communication bandwidth.

错误:搜索内容不能为空,请输入英文关键词
错误:关键词超出字数限制,请精简
高级检索

Federated Learning with Neural Graphical Models

  • Urszula Chajewska,
  • Harsh Shrivastava

摘要

Federated Learning (FL) addresses the need to create models based on proprietary data in such a way that multiple clients retain exclusive control over their data, while all benefit from improved model accuracy due to pooled resources. Recently proposed Neural Graphical Models (NGMs) are Probabilistic Graphical Models that utilize the expressive power of neural networks to learn complex non-linear dependencies between the input features. They learn to capture the underlying data distribution and have efficient algorithms for inference and sampling. We develop a FL framework which maintains a global NGM model that learns the averaged information from the local NGM models while the training data is kept within the client’s environment. Our design, FedNGM, avoids the pitfalls and shortcomings of neuron matching frameworks like Federated Matched Averaging that suffers from model parameter explosion. Our global model size doesn’t grow with the number or diversity of clients. In the cases where clients have local variables that are not part of the combined global distribution, we propose a Stitching algorithm, which personalizes the global NGM model by merging additional variables using the client’s data. FedNGM is robust to data heterogeneity, large number of participants, and limited communication bandwidth.