Federated Learning Client Selection Method Based on Constraint Factors
By initializing the global model and constraint factor in federated learning and selecting the client based on the constraint factor and cosine similarity, the problem of accuracy and fairness in federated learning in non-independent homogeneous distribution scenarios is solved, and more efficient model training and convergence are achieved.
Patent Information
- Application Number
- CN202311132778.3
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2023-09-04
- Publication Date
- 2025-07-01
- Estimated Expiration
- 2043-09-04
AI Technical Summary
The prior art in federated learning under non-independent and homogeneous distribution scenarios leads to lower accuracy and poor fairness of the global optimal model.
By initializing the global model and constraint factors in the central server and selecting clients based on the constraint factor and cosine similarity, ensure that each client's participation and contribution to the global model are fully considered.
It improves the accuracy and fairness of the federated learning model, avoids the defect of the model being easily trapped in local optimality, and enhances the convergence speed of the global model.
Smart Images

Figure CN117217328B_ABST
Abstract
Description
Technical Field
[0001] The present invention belongs to the technical field of federated learning, and relates to a method for selecting federated learning clients based on constraint factors, which can be applied to problems such as federated learning classification and prediction in scenarios where client data is non-independent and identically distributed in fields such as healthcare and finance. Background Art
[0002] Federated learning is a distributed machine learning method. In each round of training, the central server broadcasts global model parameters to clients. The clients use the datasets they hold to locally train the global model to obtain their own local models, and then upload the local model parameters to the central server. Subsequently, the central server aggregates the local models of the clients to obtain a new round of global model. After multiple rounds of training, the global model converges to the global optimal model. In federated learning, clients do not need to upload their local data, which protects the data privacy of clients. Therefore, it is widely applied in fields such as healthcare and finance that require protecting client privacy.
[0003] Due to the heterogeneity of resources such as the computing power and storage capacity of different clients and the ways of collecting data, the statistical characteristics of the data held by clients in federated learning are usually non-independent and identically distributed. However, using such non-independent and identically distributed data for local training will cause the problem of divergence of local model parameters, and aggregating such local models will reduce the accuracy and convergence speed of the global model.
[0004] How to improve the accuracy and convergence speed of the global optimal model and ensure the fairness of the clients participating in the training is the key point and difficulty of federated learning. During the process of federated learning, it is necessary to consider the contribution of each client to the training of the global model. Only by ensuring its fairness can each client be motivated to use more and better data to participate in the training of the global model. Existing technologies can improve the accuracy of the federated learning model and accelerate the model convergence speed by selecting some clients for model parameter aggregation in the non-independent and identically distributed scenario. For example, the patent application with the publication number CN115695429A and the title "Federated Learning Client Selection Method for Non-IID Scenarios" discloses a federated learning client selection method for Non-IID scenarios. The invention first initializes the global model by the central server and randomly selects a subset of clients from all available clients, and broadcasts the global model to the clients in the subset. Each client receives the global model from the server, trains the local model using the local original data, obtains the local update, and calculates the average loss. Then the client sends the local update and the loss value to the server, and the server selects the clients based on the loss values obtained by the clients. The invention filters out low-quality clients through the loss values before aggregating the model, and selects the local models of high-quality clients for aggregation to improve the model accuracy and convergence speed in the non-independent and identically distributed scenario. However, since the invention does not consider the participation degree of the clients and their contribution degree to the global model, the model is prone to fall into the local optimum rather than the global optimum. Summary of the Invention
[0005] The purpose of the present invention is to overcome the defects existing in the above-mentioned prior art, and propose a federated learning client selection method based on a constraint factor to solve the technical problems of low accuracy of the global optimal model obtained by federated learning and poor fairness of federated learning in the prior art.
[0006] To achieve the above purpose, the technical solution adopted by the present invention includes the following steps:
[0007] (1) Initialize the federated learning system:
[0008] Initialize the federated learning system including a central server and N clients. The number of local training iterations of the clients is t, the maximum number of iterations is T, and the global constraint factor threshold is and let t = 1, where N ≥ 2, T ≥ 1;
[0009] (2) Each client obtains a training sample set:
[0010] Each client C n holds Z pieces of data X that are non-independent and identically distributed with any other client nand its corresponding label Y n constitute the training sample set D n ={X n , Y n}, where Z≥50, X n ={x n1 , x n2 ,..., x nz ,..., x nZ}, x nz represents the z-th data held by C n , and the corresponding label of x nz is y nz ;
[0011] (3) The central server initializes the global model and broadcasts the model parameters:
[0012] The central server initializes the global model for its r-th round of communication with each client C n with parameters θ r , and broadcasts θ r to each client C n . Meanwhile, it initializes the constraint factor for the r-th round of communication for each client C n
[0013] (4) Each client conducts local training on the global model:(4) Each client conducts local training on the global model:
[0014] Each client C n uses as the parameters of the global model when t = 1, and conducts T times of local iterative training on the global model through the training sample set D n , and uploads the parameters of the locally trained model to the central server;
[0015] (5) The central server obtains the client selection result based on the constraint factor and cosine similarity:
[0016] The central server selects K clients whose constraint factors n of client C meet the global constraint factor threshold as alternative clients, and calculates the cosine similarity between the local model parameters k of each alternative client C and the global model parameters θ r . After updating the constraint factors of the M clients with the highest cosine similarity, these M clients are taken as the final selection result. Compared with the prior art, the present invention has the following advantages:
[0017] Compared with the prior art, the present invention has the following advantages:
[0018] By comparing the constraint factor of each client in each round with the constraint factor threshold, the central server of the present invention obtains alternative clients, and selects the clients with the highest cosine similarity among all alternative clients as the final selection result, fully considering the participation degree of each client and its contribution degree to the global model, avoiding the defect that the model is prone to fall into local optimum in the prior art, and effectively improving the accuracy of the model and the fairness of federated learning. BRIEF DESCRIPTION OF THE DRAWINGS
[0019] Figure 1 It is a flowchart for implementing the present invention. DETAILED DESCRIPTION OF THE INVENTION
[0020] The present invention will be further described in detail below with reference to the accompanying drawings and specific embodiments.
[0021] Refer to Figure 1 , the present invention includes the following steps:
[0022] (1) Initialize the federated learning system:
[0023] Initialize the federated learning system including a central server and N clients. The number of iterations of local training for the clients is t, the maximum number of iterations is T, and the global constraint factor threshold is and let t = 1, where N ≥ 2, T ≥ 1;
[0024] In this embodiment, the number of federated learning clients N = 50, and the constraint factor threshold The maximum number of iterations for local training T = 10;
[0025] (2) Each client obtains a training sample set:
[0026] Each client C n holds Z pieces of data X that are non-identically distributed with any other client n and their corresponding labels Y n to form a training sample set D n = {X n , Y n}, where Z ≥ 50, X n = {x n1 , x n2 ,..., x nz ,..., x nZ}, x nz represents the z-th piece of data held by C n , and the label corresponding to x nz is y nz ;
[0027] Each client C in the present invention n holds Z pieces of data, which can be any one of image data, audio data, and text data. In this embodiment, Z = 100 is taken, and the MNIST dataset is used as the image dataset to train a classification model; the MNIST dataset is a grayscale handwritten digit image dataset widely used in the field of federated learning, which contains 60,000 training image data, each of which is a grayscale handwritten digit image of size 28×28, and its label is the true digit corresponding to the handwritten digit, taking values from 0 to 9; in this embodiment, each client C n respectively holds 100 different grayscale handwritten digit images X n with true label Y n . In order to ensure that the statistical characteristics of the data held by different clients satisfy non-independent and identically distributed, in this embodiment, the data held by each client C n is set to 4 out of 10 labels;
[0028] (3) The central server initializes the global model and broadcasts the model parameters:
[0029] The central server initializes the global model with parameters θ n for its r-th round of communication with each client C r , and broadcasts θ r to each client C n . At the same time, it initializes the constraint factor for the r-th round of communication for each client C n .
[0030] In this embodiment, the global model initialized by the central server is a classification model, and a convolutional neural network model including two composite layers and a fully connected layer stacked in sequence is adopted. Each composite layer is composed of a convolutional layer, an activation function layer, and a max pooling layer stacked in sequence. The sizes of the convolutional kernels of the two convolutional layers are both 3×3. The number of channels of the first convolutional layer is 10, and the number of channels of the second convolutional layer is 20. The activation function layer adopts the ReLU function, and the size of the pooling kernel of the max pooling layer is 2×2.
[0031] When r = 1, the parameters of the global model are initialized by the central server, and the constraint factor corresponding to each client is initialized to zero. When r > 1, the parameters θ r of the global model are aggregated by the central server from the local model parameters of the clients selected in the r-th round of communication, and the constraint factor of each client is obtained by updating.
[0032] (4) Each client performs local training on the global model:
[0033] Each client C n will As the parameters of the global model at t = 1, and through the training sample set D n Perform T local iterative trainings on the global model, and upload the parameters of the locally trained model to the central server. The specific implementation steps are as follows:
[0034] (4a) Each client C n uses the training sample set D n as the input of the global model for forward propagation to obtain the predicted result y' of each data x nz for the current iteration, and uses the cross-entropy loss function to calculate the loss value L of the global model through y' nz,t and the label y of x nz,t : nz where Σ represents the summation operation; nz n,t
[0035]
[0036]
[0037] (4b) Each client C n updates n,t through the partial derivative of L with respect to to obtain the global model for the current iteration. Among them, the update formula of with respect to is: where η represents the learning rate;
[0038]
[0039]
[0040] (4c) Each client C n judges whether t = T holds. If so, obtain the local model with the parameters of the T-th iteration as and upload as the parameters of the local model for the r-th round of communication to the server. Otherwise, set t = t + 1 and execute step (4a).
[0041] (5) The central server obtains the client selection result based on the constraint factor and cosine similarity:
[0042] The central server selects K clients whose constraint factors n of client C meet the global constraint factor threshold as the alternative clients, and calculates the local model parameters of each alternative client C k : with the global model parameters θ r cosine similarity After updating the constraint factors of the M clients with the highest cosine similarities, these M clients are used as the final selection results.
[0043] The present invention compares the constraint factor of each client with the global constraint factor threshold to limit the upper limit of the number of times each client participates in the selection, fully considering the participation degree of each client and the contribution degree to the global model, improving the fairness of federated learning. The central server calculates the local model parameters of each alternative client C k of with the global model parameters θ r cosine similarity and updates the constraint factors of the M clients with the highest cosine similarities. Finally, these M clients are selected as the selection results to participate in the subsequent aggregation steps of this round. In this embodiment, M = 20 is taken. The calculation formula and the update formula of the constraint factor are respectively:
[0044]
[0045]
[0046] where ||·|| is the L2 norm operation, and <·> represents the inner product operation. represents the constraint factor of the m-th client before update, 1 ≤ m ≤ M.
Claims
1. A method for selecting a federated learning client based on a constraint factor, characterized in that It includes the following steps: (1) Initialize the federated learning system: Initialize a federated learning system including a central server and N clients. The number of local training iterations of the clients is t, the maximum number of iterations is T, and the global constraint factor threshold is and set t = 1, where N ≥ 2, T ≥ 1; (2) Each client obtains a training sample set: Each client C n will hold Z pieces of data X that are non-identically distributed with any other client n and their corresponding labels Y n to form a training sample set D n ={X n ,Y n}, where the Z pieces of data are any one of image data, audio data, and text data, Z≥50, X n ={x n1 ,x n2 ,...,x nz ,...,x nZ}, x nz represents the z-th piece of data held by C n , and the label corresponding to x nz is y nz ; (3) The central server initializes the global model and broadcasts the model parameters: The central server initializes its connection with each client C n The parameter for the r-th round of communication is θ r of the global model, and broadcasts θ r to each client C n , and at the same time initializes the constraint factor for the r-th round of communication for each client C n (4) Each client conducts local training on the global model: Each client C n uses as the parameters of the global model at t = 1, and performs T local iterative trainings on the global model through the training sample set D n and uploads the parameters of the locally trained model to the central server; (5) The central server obtains the client selection result based on the constraint factor and cosine similarity: The central server takes the client C n 's constraint factor and the global constraint factor threshold to satisfy the K clients as alternative clients, and calculates the local model parameters k of each alternative client C and the cosine similarity r with the global model parameter θ After updating the constraint factors of the M clients with the highest cosine similarity, these M clients are used as the final selection result, where the calculation formula, and the update formula for the constraint factor of the m-th client are as follows: where ||·|| is the L2 norm operation, <·> represents the inner product operation, and 1 ≤ m ≤ M.
2. The method according to claim 1, characterized in that, The global model described in step (3) adopts a convolutional neural network including two composite layers and a fully connected layer stacked in sequence. Each composite layer consists of a convolutional layer, an activation function layer, and a max pooling layer stacked in sequence. The sizes of the convolutional kernels of the two convolutional layers are both 3×3. The number of channels of the first convolutional layer is 10, and the number of channels of the second convolutional layer is 20. The activation function layer adopts the ReLU function, and the size of the pooling kernel of the max pooling layer is 2×2.
3. The method according to claim 1, characterized in that, Each client C described in step (4) n Using the training sample set D n Perform T iterations of training on the global model, and the implementation steps are as follows: (4a) Each client C n forwards the training sample set D n as the input of the global model for forward propagation, obtaining the prediction result y' of this iteration corresponding to each data x nz and uses the cross-entropy loss function to calculate the loss value L of the global model through y' nz,t and the label y of x nz,t : nz nz n,t where ∑ represents the summation operation; (4b) Each client C n through L n,t partial derivative of with respect to for is updated to obtain the global model for this iteration, where the update formula for where η represents the learning rate; (4c) Each client C n Determine whether t = T holds. If so, obtain the local model with the T-th iteration parameter as and use as the parameter of the local model for the r-th round of communication Otherwise, set t = t + 1 and execute step (4a).
Citation Information
Patent Citations
Robustness federated learning model aggregation method based on truth value discovery
CN114186237A
Federal learning client selection method for Non-IID scene
CN115695429A