Federated learning with improved aggregation via optimal transport algorithm
Dawei Chen, Yuan Lu · 2024
With the increasing awareness of privacy protection, federated learning is increasingly applied to distributed training scenarios. However, in the process of executing federated learning on distributed clients, due to the non-independent and identical distribution between clients, the accuracy and convergence speed of the global model will decrease in the model aggregation process of ordinary federated learning. To solve this problem, a federated learning framework is proposed to improve the model aggregation method. During training, the server extracts the feature parameters of the local model, and performs entropy regularization on each layer of the local model to obtain the optimal transmission feature parameters. Finally, the optimal transmission and other federated learning global model feature parameters are generated by fusion. In the distributed training of the two data sets, the data distribution in four different cases was simulated for federated training comparison. The results show that compared with ordinary federated learning, the proposed algorithm significantly improves the model performance and the convergence speed of the global model, which can be used for distributed user training and commercial privacy-sensitive scenarios.