Client Selection Method for Federated Learning Based on Grouping Reinforcement Learning
Guoming Li, Waixi Liu, Zhen-zheng Guo, Dao-xiao Chen · 2024
Federated learning (FL) is currently the most widely adopted machine learning model collaborative training framework under privacy constraints. However, it still faces issues such as client data heterogeneity and communication bottlenecks with the server. To address these problems, this article proposes a client selection method for federated learning based on the grouping reinforcement learning (CSFL). Initially, it groups client populations with jointly trainable data distributions through clustering model parameters from clients. Within each group, a client selection model is trained to intelligently select clients for participation in each round of federated learning, mitigating bias introduced by non-IID data and accelerating convergence. The client selection model within each group uses a double Deep Q-Network to select an optimal subset of clients in each communication round. Certainly, we introduce an adjustment strategy for reinforcement learning, aiming to reduce the impact caused by the initial random client selection in reinforcement learning. Furthermore, through theoretical analysis, we have determined the optimal number of clients to select for each cluster in each round. Through extensive experiments conducted in PyTorch, we demonstrate that CSFL can reduce the required communication rounds by up to 37% on the MNIST dataset and 21% on CIFAR-10 compared to the Federated Averaging algorithm.