Heterogeneous federal learning training method with efficient communication
By using the initial model weight and probability mask for local training and uploading binary masks in heterogeneous federated learning, the problem of server-to-client communication overhead is solved, efficient federated learning training is achieved, and communication costs are reduced and model performance is improved.
Patent Information
- Application Number
- CN202510009813.5
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-01-03
- Publication Date
- 2025-05-06
- Estimated Expiration
- 2045-01-03
AI Technical Summary
In the heterogeneous federated learning scenario, the prior art has problems such as large server-to-client communication overhead and low model performance, which is difficult to effectively reduce the communication cost of federated learning and improve the availability of the model.
A heterogeneous federated learning training method with efficient communication is proposed. The initial model weight and probability mask are generated by the server, the client conducts local training and uploads binary masks, and the server performs clustering and grouping aggregation, reducing communication overhead and improving model performance.
This method can significantly reduce the communication cost of federated learning, improve the usability and reasoning performance of the model, especially on low-end devices, and increase the enthusiasm of low-end devices to participate in federated learning training.
Smart Images

Figure CN119940479A_ABST
Abstract
Description
Technical Field
[0001] The present invention belongs to the field of machine learning, and in particular relates to a heterogeneous federated learning training method with efficient communication. Background Art
[0002] Federated learning is an emerging distributed machine learning paradigm that helps many clients collaborate to optimize a global model without sharing local sensitive data. Due to its privacy-preserving properties, federated learning has become popular in both academia and industry, especially in the fields of medical image analysis, face recognition, and personalized recommendation systems. In federated learning, a large amount of communication is required between the server and the client for iterative training. This communication cost becomes particularly significant when dealing with large models. Taking large language models and Q-learning as examples, the current trend of model architecture involves unprecedented scale, containing billions or even trillions of learnable parameters. Although complex models have excellent performance, the huge demand for communication bandwidth and throughput of such large models makes traditional federated learning frameworks unsuitable for training at this scale. In addition, data heterogeneity is common in actual federated learning, which severely degrades the performance of the model. For example, a mobile application may collect user behavior data, such as application usage frequency, location information, social interactions, etc. Since user preferences and behaviors may vary due to geographical locations and cultural differences, these data may show significant heterogeneity among different user groups. Therefore, how to improve performance and communication efficiency in heterogeneous federated learning scenarios has become a difficult problem that needs to be solved urgently.
[0003] In their paper "Heterofl: Computation and communication efficient federated learning for heterogeneous clients (ICLR 2021)", Diao et al. proposed that the server allocates global model sub-models for local training to the client based on the client's current computing and communication resources. The server divides the client's computing and communication resources into P levels, each level corresponding to a category of global sub-models. The global sub-model reduces the input and output channels of each layer at a fixed ratio according to different levels. The server sends the divided sub-models to the corresponding clients for local training, and then uploads them to the server. The server constructs a global model through specific aggregation rules, and then divides the aggregated global model in turn, repeating the iterative training until convergence. The shortcomings of this method are: on the one hand, since low-end devices are allocated with smaller model structures, the inference performance of low-end devices is seriously reduced, affecting the enthusiasm of low-end devices to participate in federated learning training; on the other hand, high-end devices are allocated with almost complete models, and the communication overhead of high-end devices has not been reduced, and the communication overhead is still heavy.
[0004] Shandong University disclosed a method for optimizing the performance and communication efficiency of heterogeneous federated learning in its patent application "A communication method and system for heterogeneous federated learning in the industrial Internet of Things" (patent application number: CN202410487452.0, authorization announcement number: CN118101501B). This method triggers the multi-granularity quantization module to quantize the updated local model parameters, and the central server receives the model updates uploaded by each device node, and passes the quantized model updates to the dequantization module for dequantization. Then, the central server sends the dequantized results to the stage aggregation module to realize model parameter aggregation, and repeats the iteration until this federated learning reaches the predetermined number of training rounds. The shortcomings of this method are: on the one hand, the global model after dequantization and aggregation by the central server has poor inference performance on devices with small data volume; on the other hand, this method only reduces the communication overhead from the client to the server, and the communication overhead from the server to the client is not reduced, which means that low-end devices are still unable to participate in the training of complex models.
[0005] In summary, existing methods such as model gradient quantization and compression still have problems such as high communication overhead from server to client and low model performance in heterogeneous federated learning scenarios. Therefore, there is an urgent need for a communication-efficient heterogeneous federated learning training method to ensure the availability and correctness of the model when performing detection tasks in specific fields, while reducing the communication cost of federated learning. Summary of the invention
[0006] In order to solve the above problems existing in the prior art, the present invention provides a heterogeneous federated learning training method with efficient communication. The technical problem to be solved by the present invention is achieved by the following technical solutions:
[0007] A communication efficient heterogeneous federated learning training method is applied to a system consisting of a server and multiple clients, the method comprising:
[0008] S100, the server generates the initial model weight w0 and initial weight score s of the neural network model g,0 , by communicating with the client, each client obtains the initial model weight w0 to initialize the local model; the server uses the initial weight score s g,0 Calculate the initial probability mask θ g,0 , sent to all clients;
[0009] S200, in the tth round of global training iteration, the server selects K t Clients participate in the federated learning model training;
[0010] S300, each selected client k performs E rounds of local model training based on the local data set, and updates and optimizes the local weight score s k,t , and upload the binary mask m k,t to the server; wherein the local data set includes a plurality of images;
[0011] S400, the server clusters the binary masks uploaded by all clients, and allocates cluster IDs to the corresponding binary masks to implement cluster ID allocation for the clients;
[0012] S500, the server groups and aggregates the binary masks of each cluster after clustering to obtain the probability mask within the cluster and the global probability mask θ of the tth round of global training iteration g,t ;
[0013] S600, the server notifies the client to perform a consistency check on the allocated cluster ID;
[0014] S700, the server sends the intra-cluster probability mask to the client in the corresponding cluster, and sends the global probability mask to the client to which the cluster ID is not assigned, so that the client recovers the probability mask and performs local training until the model converges;
[0015] S800, each client obtains a binary mask according to the probability mask within the cluster, and uses the obtained binary mask and the initial model weight w0 to obtain the final sub-network model, thereby obtaining a trained local model; the local model is used to complete the preset image detection task.
[0016] In one embodiment of the present invention, S100 includes:
[0017] S110, the server uses the random seed S to generate the initial model weight w0 and initial weight score s of the neural network model g,0 ; Among them, the weight score represents the importance of the weight on the neural network model. The larger the value, the more beneficial the weight is to model reasoning;
[0018] S120, the server sends the random seed S to all clients, so that each client constructs the same initial model weight w0 locally to initialize the local model;
[0019] S130, the server end calculates the initial weight score s g,0 and logistic function to calculate the initial probability mask θ g,0 , and the initial probability mask θ g,0 Sent to each client; where the probability mask represents the probability of the weight being used in the forward propagation.
[0020] In one embodiment of the present invention, S300 includes:
[0021] S310, in the tth round of global training iterations, each client k uses the logic function Logit and the global probability mask θ received from the server g,t-1 , calculate the weight score s of the local model g,t ; Among them, for the first round of global training, the global probability mask θ g,t-1 is the initial probability mask θ g,0 ;
[0022] S320, start the current round of local training, each client k will set the weight score s of the local model k,t Use the logistic function to map to the probability mask θ k,t ; Among them, in the first round of local training of the tth round of global training, the weight score s of the local model k,t is the weight score s of the local model g,t ;
[0023] S330, each client k pairs of probability masks θ k,t Perform Bernoulli sampling to generate binary mask m k,t ;
[0024] S340, each client k receives a binary mask m k,t And the initial model weight w0, by multiplying the corresponding elements of the matrix, calculate the sub-network model w for forward propagation k,t ; When the value of the binary mask is 1, it means that the corresponding weight participates in the forward propagation, and when the value of the binary mask is 0, it means that the corresponding weight does not participate in the forward propagation;
[0025] S350, each client k uses the local data set and sub-network model w k,t , calculate its gradient value after forward propagation, and then update the local weight score s k,t ;
[0026] S360, each client k repeats steps S320 to S350 E times to complete E rounds of local training and obtain an updated and optimized weight score s k,t ;
[0027] S370, each client k adds the weight score s obtained in step S360 k,t Convert to optimized binary mask m k,t , and upload it to the server.
[0028] In one embodiment of the present invention, S400 includes:
[0029] S410, the server receives Kt Binary mask uploaded by client
[0030] S420, the server receives K t Randomly select Q binary masks from the binary masks as the initial centroid of each cluster;
[0031] S430, the server assigns each binary mask to the cluster where the centroid closest to the binary mask is located, and obtains a preliminary clustering result; wherein the preliminary clustering result records the ID of the client in each cluster;
[0032] S440, for each cluster, a centroid is calculated so that the sum of the distances between the centroid and each binary mask of the cluster is minimized, thereby re-obtaining the centroid of each cluster;
[0033] S450, repeat steps S430 to S440T1 times, and obtain K t A binary mask M K,t The final clustering result is obtained, and at the same time, a corresponding cluster ID is assigned to each client k according to the final clustering result.
[0034] In one embodiment of the present invention, S430 includes:
[0035] S4301, calculate each binary mask m k To the center of mass m q The distance d q,k :
[0036]
[0037] Among them, L is the number of layers of the neural network model, n l represents the number of model parameters in the lth layer, Represents the binary mask m k The parameters of the lth layer, Represents the exclusive OR operation;
[0038] S4302, for each binary mask m k , determine the cluster ID of the nearest centroid, expressed as:
[0039] q = argmin q∈[Q] d q,k ;
[0040] S4303, the binary mask m k Assign it to its nearest cluster q, and its operation is expressed as:
[0041] C[q].insert(k);
[0042] Among them, C represents the clustering result, and C[q] records the IDs of all clients in cluster q.
[0043] In one embodiment of the present invention, S500 includes:
[0044] S510, the server aggregates the binary masks in each cluster q to obtain the cluster probability mask θ after cluster q aggregation. g,t [q], calculated as
[0045] S520, the server aggregates all intra-cluster probability masks to obtain a global probability mask for the tth round of global training iteration
[0046] In one embodiment of the present invention, S600 includes:
[0047] S610, the server uses the aggregated probability mask θ g,t [q] Perform Bernoulli sampling to obtain Q clusters of binary masks M t , then M t Sent to each client; where M t ={m q,t |m q,t = Bern(Θ g,t [q])} q∈[Q] ;Bern means Bernoulli sampling;
[0048] S620, each client generates a binary mask M in each cluster. t [q] Obtain the sub-network model of forward propagation, and then calculate its loss value based on the local data set, and obtain the cluster with the smallest loss value to determine the locally inferred cluster ID; wherein, when the cluster ID locally inferred by the client is consistent with the cluster ID assigned by the server, the corresponding client will no longer participate in the process of clustering and assigning cluster IDs. When the cluster ID locally inferred by the client is inconsistent with the cluster ID assigned by the server, the cluster ID locally inferred by the client shall prevail.
[0049] In one embodiment of the present invention, S700 includes:
[0050] S710, the server sends an integer and integers Sent to the client with cluster ID q; represents the number of clients in each cluster, represents the sum of the binary masks in each cluster;
[0051] S720, each client receives and Then recover the global probability mask θg,t , and perform local training T times to make the model converge.
[0052] In one embodiment of the present invention, S800 includes:
[0053] S810, each client performs a probability mask Θ on the cluster g,T [q] Perform Bernoulli sampling to obtain a binary mask m q = Bern(Θ g,T [q]);
[0054] S820, each client generates a binary mask m according to the binary mask m. q And the initial model weight w0, by multiplying the corresponding elements of the matrix, a sub-network model with reasoning ability is generated to constitute a trained local model.
[0055] In one embodiment of the present invention, the server is a central server, and the client is a server of a medical institution; the local data set of the client includes the medical imaging data of the patient and is marked with a tumor type label;
[0056] The preset image detection task includes: the client uses the trained local model to detect the medical image data to be tested to obtain the tumor type.
[0057] Beneficial effects of the present invention:
[0058] The embodiment of the present invention provides a communication-efficient heterogeneous federated learning training method. On the one hand, during local model training, the client optimizes locally and uploads a binary mask to the server, which can improve the upload communication efficiency (each parameter occupies 1 bit). The client downloads the integer clustering mask to improve the download communication efficiency, so that low-end devices can also participate in complex federated learning model training, improve the availability of the model, and significantly reduce the communication cost of federated learning. On the other hand, the binary mask uploaded by the client can well reflect the characteristics of data distribution, and the clients with different data distributions are estimated according to the similarity of the binary mask. Clients with heterogeneous data can be clustered without additional information, so that clients with different data distributions have different models, which improves the model reasoning performance in heterogeneous scenarios. By clustering and distributing the binary mask model on the server, the communication efficiency and model performance of heterogeneous federated learning can be improved. BRIEF DESCRIPTION OF THE DRAWINGS
[0059] Figure 1 A flow chart of a heterogeneous federated learning training method with efficient communication provided by an embodiment of the present invention;
[0060] Figure 2A flowchart of another communication-efficient heterogeneous federated learning training method provided by an embodiment of the present invention. DETAILED DESCRIPTION
[0061] The present invention is further described in detail below with reference to specific embodiments, but the embodiments of the present invention are not limited thereto.
[0062] Federated learning can unite multiple users to jointly learn and train a global model without leaking privacy. However, the heterogeneity of federated learning data distribution and the limited communication bandwidth of low-end devices seriously affect the convergence and performance of federated learning models. At present, the existing federated learning training methods have the following main shortcomings:
[0063] (1) Existing technologies rely on compression or quantization methods to only reduce the communication overhead from the client to the server. The communication overhead from the server to the client is still heavy, and the communication cost of model training is high.
[0064] (2) The models trained in heterogeneous federated learning in the existing technology have low model inference performance and poor usability when the data volume is small or the equipment is low.
[0065] In order to solve the above problems, an embodiment of the present invention provides a communication-efficient heterogeneous federated learning training method, which is applied to a system consisting of a server and multiple clients, and aims to use the trained local model to achieve detection tasks in a specific field, where the specific field may include medical and other fields.
[0066] See also Figure 1 , and combined with Figure 2 It is understood that the method may include the following steps S100 to S800:
[0067] S100, the server generates the initial model weight w0 and initial weight score s of the neural network model g,0 , by communicating with the client, each client obtains the initial model weight w0 to initialize the local model; the server uses the initial weight score s g,0 Calculate the initial probability mask θ g,0 , sent to all clients;
[0068] The purpose of S100 is to enable each client to generate an initial neural network model through server initialization, that is, to obtain an initial local model.
[0069] In an optional implementation manner, S100 may include:
[0070] S110, the server uses the random seed S to generate the initial model weight w0 and initial weight score s of the neural network model g,0 ;
[0071] In the embodiment of the present invention, the weight score represents the importance of the weight on the neural network model, and the larger the value, the more beneficial the weight is to the model reasoning;
[0072] Among them, the initial model weight w0 remains unchanged throughout the process, so the initialization method of the model weight has a great impact on the performance of the model. The present invention uses the "Kaiming constant" method to initialize the model weight. Specifically, the absolute value of each weight value of each layer of the neural network model is a constant σ, and its sign is randomly selected as positive or negative. The constant σ is the standard deviation of the Kaiming normal distribution, and its calculation formula is n l is the number of model parameters of the lth layer of the neural network model. In other words, the constant σ of each layer is determined according to the number of model parameters of the layer.
[0073] The random seed S is a common method, which will not be described in detail here. The method of generating the initial model weight w0 can be random.
[0074] S120, the server sends the random seed S to all clients, so that each client locally constructs the same initial model weight w0 to initialize the local model;
[0075] After the server sends the random seed S to all clients, the clients can use the random seed S to generate the initial model weight w0 of the neural network model, thereby constructing the same neural network model and realizing local model initialization.
[0076] As mentioned earlier, during the client local training process, the initial model weight w0 remains unchanged.
[0077] S130, the server end calculates the initial weight score s g,0 and logistic function to calculate the initial probability mask θ g,0 , and the initial probability mask θ g,0 Sent to each client;
[0078] Among them, the probability mask represents the probability that the weight is used in the forward propagation.
[0079] Specifically, calculate the initial probability mask θ g,0 The formula is: g,0 =Logistic(s g,0 ), Logistic is the logistic function. The logistic function is to convert the initial weight score s g,0 Mapping from (-∞, +∞) to the initial probability mask θ g,0 The (0,1) interval.
[0080] S200, in the tth round of global training iteration, the server selects K t Clients participate in the federated learning model training;
[0081] Among them, select K t A client can be randomly selected.
[0082] S300, each selected client k performs E rounds of local model training based on the local data set, and updates and optimizes the local weight score s k,t , and upload the binary mask m k,t To the server; wherein the local data set includes a plurality of images; according to different image detection task requirements, the local data set can be set accordingly.
[0083] E is a natural number greater than 0 and can be set as needed.
[0084] In an optional implementation manner, S300 may include:
[0085] S310, in the tth round of global training iterations, each client k uses the logic function Logit and the global probability mask θ received from the server g,t-1 , calculate the weight score s of the local model g,t ;
[0086] Among them, for the first round of global training, the global probability mask θ g,t-1 is the initial probability mask θ g,0 ;
[0087] Calculate the weight score s of the local model g,t The formula is:
[0088] s g,t =Logit(θ g,t-1 );
[0089] Among them, the logical function Logit is a function that maps real numbers to the interval (0,1).
[0090] S320, start the current round of local training, each client k will set the weight score s of the local model k,t Use the logistic function to map to the probability mask θ k,t ;
[0091] The specific calculation formula is: k,t =Logistic(s k,t );
[0092] Among them, in the first round of local training of the tth round of global training, the weight score s of the local model k,t is the weight score s of the local modelg,t ;
[0093] S330, each client k pairs of probability masks θ k,t Perform Bernoulli sampling to generate binary mask m k,t ;
[0094] The specific calculation formula is: k,t = Bern(θ k,t );
[0095] Among them, Bern represents Bernoulli sampling.
[0096] S340, each client k receives a binary mask m k,t And the initial model weight w0, by multiplying the corresponding elements of the matrix, calculate the sub-network model w for forward propagation k,t ;
[0097] The specific calculation formula is: k,t =m k,t ⊙w0;
[0098] Among them, ⊙ represents the multiplication of the corresponding elements of the matrix; when the value of the binary mask is 1, it means that the corresponding weight participates in the forward propagation, and when the value of the binary mask is 0, it means that the corresponding weight does not participate in the forward propagation;
[0099] S350, each client k uses the local data set and sub-network model w k,t , calculate its gradient value after forward propagation, and then update the local weight score s k,t ;
[0100] The specific calculation formula is:
[0101] Among them, ← means to update the left side with the right side; the s on the right side k,t is the original weight score, and the left side is the updated weight score; η represents the learning rate, B represents the batch size (i.e., the number of samples processed simultaneously during each model training iteration), and L represents the loss function. Indicates the calculation of the gradient value, D k represents the local dataset of client k, D k,b Represents the bth batch of samples of size Batchsize in the local dataset of client k.
[0102] S360, each client k repeats steps S320 to S350 E times to complete E rounds of local training and obtain an updated and optimized weight score s k,t ;
[0103] S370, each client k adds the weight score s obtained in step S360k,t Convert to optimized binary mask m k,t , and upload it to the server.
[0104] Specifically, the client k first converts the weight score s obtained in step S360 into k,t , mapped to a probability mask θ k,t , the calculation formula is θ k,t =Logistic(s k,t ), then the probability mask θ k,t Perform Bernoulli sampling to obtain the optimized binary mask m k,t , the calculation formula is m k,t = Bern(θ k,t ), and finally upload it to the server.
[0105] It should be noted that the client of the present invention uploads a binary mask m k,t , each of which can be represented by 1 bit. In the traditional federated learning mechanism, the client uploads the model gradient, and each parameter is a floating point number, which needs to be represented by 32 bits. Therefore, compared with the traditional federated learning mechanism, the present invention can reduce the communication overhead uploaded by the client by 32 times.
[0106] S400, the server clusters the binary masks uploaded by all clients, and allocates cluster IDs to the corresponding binary masks to implement cluster ID allocation for the clients;
[0107] In an optional implementation manner, S400 may include:
[0108] S410, the server receives K t Binary mask uploaded by client
[0109] S420, the server receives K t Randomly select Q binary masks from the binary masks as the initial centroid of each cluster;
[0110] Step S420 is equivalent to determining Q clusters.
[0111] S430, the server assigns each binary mask to the cluster where the centroid closest to the binary mask is located, and obtains a preliminary clustering result; wherein the preliminary clustering result records the ID of the client in each cluster;
[0112] Specifically, S430 may include:
[0113] S4301, calculate each binary mask m k To the center of mass m q The distance d q,k:
[0114]
[0115] Among them, L is the number of layers of the neural network model, n l Indicates the number of model parameters in the lth layer, m l k Represents the binary mask m k The parameters of the lth layer, Represents the exclusive OR operation;
[0116] S4302, for each binary mask m k , determine the cluster ID of the nearest centroid, expressed as:
[0117] q = argmin q∈[Q] d q,k ;
[0118] Among them, argmin means finding the minimum value.
[0119] S4303, the binary mask m k Assign it to its nearest cluster q, and its operation is expressed as:
[0120] C[q].insert(k);
[0121] Among them, C represents the clustering result, and C[q] records the IDs of all clients in cluster q.
[0122] S440, for each cluster, a centroid is calculated so that the sum of the distances between the centroid and each binary mask of the cluster is minimized, thereby re-obtaining the centroid of each cluster;
[0123] S450, repeat steps S430 to S440T1 times, and obtain K t A binary mask M K,t The final clustering result is obtained, and at the same time, a corresponding cluster ID is assigned to each client k according to the final clustering result.
[0124] Specifically, the final clustering result can still be represented by C, and a corresponding cluster ID q is allocated to each client k according to the final clustering result.
[0125] S500, the server groups and aggregates the binary masks of each cluster after clustering to obtain the probability mask within the cluster and the global probability mask θ of the tth round of global training iteration g,t ;
[0126] In an optional implementation manner, S500 may include:
[0127] S510, the server aggregates the binary masks in each cluster q to obtain the cluster probability mask θ after cluster q aggregation. g,t [q], calculated as
[0128] S520, the server aggregates all intra-cluster probability masks to obtain a global probability mask for the tth round of global training iteration
[0129] S600, the server notifies the client to perform a consistency check on the allocated cluster ID;
[0130] Since clustering errors may occur on the server side, the present invention performs consistency check on the cluster ID on the client side.
[0131] In an optional implementation manner, S600 may include:
[0132] S610, the server uses the aggregated probability mask θ g,t [q] Perform Bernoulli sampling to obtain Q clusters of binary masks M t , then M t Sent to each client so that the client can perform a local consistency check on the cluster ID;
[0133] This step is to reduce communication overhead;
[0134] Among them, M t ={m q,t |m q,t = Bern(Θ g,t [q])} q∈[Q] ;Bern means Bernoulli sampling;
[0135] S620, each client generates a binary mask M in each cluster. t [q] Obtain the forward propagation sub-network model, then calculate its loss value based on the local data set, and obtain the cluster with the smallest loss value to determine the locally inferred cluster ID;
[0136] The specific calculation formula is: q'=argmin q∈Q L(X,M t [q]⊙w0);
[0137] Among them, q' is the locally inferred cluster ID; X is the local dataset; L is the loss function;
[0138] Among them, when the cluster ID inferred locally by the client is consistent with the cluster ID assigned by the server, the corresponding client will no longer participate in the process of clustering and allocating cluster IDs. When the cluster ID inferred locally by the client is inconsistent with the cluster ID assigned by the server, the cluster ID inferred locally by the client shall prevail, that is, q' will replace the originally assigned q.
[0139] Specifically, once the inferred cluster ID of the client is consistent with the inferred cluster ID of the server, the client cluster ID is determined to be q, and no further inference estimation is performed on the client cluster ID. Otherwise, the client updates and uploads the binary mask using its probability mask corresponding to the estimated cluster ID, and the server re-evaluates the client cluster ID until the two estimates of the cluster ID of the client and server are consistent.
[0140] S700, the server sends the intra-cluster probability mask to the client in the corresponding cluster, and sends the global probability mask to the client to which the cluster ID is not assigned, so that the client recovers the probability mask and performs local training until the model converges;
[0141] In an optional implementation manner, S700 may include:
[0142] S710, the server sends an integer and integers Sent to the client with cluster ID q;
[0143] To further improve communication efficiency, the server can send only integers. and integers For a client with cluster ID q, represents the number of clients in each cluster, represents the sum of the binary masks in each cluster. The maximum value is This can be encoded as Therefore, the server can only transmit Where N w Indicates the number of parameters in the binary mask.
[0144] S720, each client receives and Then recover the global probability mask θ g,t , and perform local training T times to make the model converge.
[0145] Recover the global probability mask θ g,t The calculation formula is:
[0146] T is a natural number greater than 0 and can be set as needed.
[0147] S800, each client obtains a binary mask according to the probability mask within the cluster, and uses the obtained binary mask and the initial model weight w0 to obtain the final sub-network model, thereby obtaining a trained local model; the local model is used to complete the preset image detection task.
[0148] Step S800 is that the client generates a final sub-network model for each cluster. Specifically,
[0149] S800, which can include:
[0150] S810, each client performs a probability mask Θ on the cluster g,T [q] Perform Bernoulli sampling to obtain a binary mask m q = Bern(Θ g,T [q]);
[0151] After the Tth training iteration, the client performs cluster probability mask Θ g,T [q] Perform Bernoulli sampling to obtain a binary mask m q The binary mask m q Able to generate a sub-network model with reasoning capabilities from the initial network model.
[0152] S820, each client generates a binary mask m according to the binary mask m. q And the initial model weight w0, by multiplying the corresponding elements of the matrix, a sub-network model with reasoning ability is generated to constitute a trained local model.
[0153] The specific calculation formula for generating the sub-network model is: q =m q ⊙w0;
[0154] Among them, w q The generated sub-network model has reasoning capabilities.
[0155] It can be seen that the final model can be deployed efficiently. It only needs to obtain the initial random seed S and binary mask m to obtain the sub-model w with reasoning ability. The random seed is used to construct the initial random network model parameter w0, and the binary mask is used to extract the sub-model w with reasoning ability from the random network model. The calculation formula is w = m⊙w0.
[0156] From the above content, it can be seen that the present invention implements a set of efficient federated learning training mechanisms. The client in the present invention updates the weight score locally, obtains the binary mask through Bernoulli sampling, and then uploads the binary mask to the server, which greatly reduces the communication efficiency from the client to the server. At the same time, the client can download the integer mask from the server to improve the communication efficiency from the server to the client. Moreover, the server of the present invention clusters clients with different data distributions by using binary masks, and then groups them together to improve the model reasoning performance of heterogeneous federated learning. The server obtains K t A binary mask M K,t The final clustering result is C. At the same time, the corresponding cluster ID q is assigned to client k, and the probability mask Θ within the cluster after cluster q is aggregated is obtained. g,t [q]. In addition, the present invention provides a binary mask clustering consistency check method, which improves the accuracy of clustering of heterogeneous federated learning models and improves the robustness of the model. The server performs Bernoulli sampling on the aggregated intra-cluster probability mask to obtain Q intra-cluster binary masks M t , and then send it to the client to perform a local consistency check on the cluster ID. If the local and server inference results are consistent, clustering will not be performed in subsequent iterations to reduce computing overhead; otherwise, clustering will still be performed in the next iteration.
[0157] The embodiment of the present invention provides a communication-efficient heterogeneous federated learning training method. On the one hand, during local model training, the client optimizes locally and uploads a binary mask to the server, which can improve the upload communication efficiency (each parameter occupies 1 bit). The client downloads the integer clustering mask to improve the download communication efficiency, so that low-end devices can also participate in complex federated learning model training, improve the availability of the model, and significantly reduce the communication cost of federated learning. On the other hand, the binary mask uploaded by the client can well reflect the characteristics of data distribution, and the clients with different data distributions are estimated according to the similarity of the binary mask. Clients with heterogeneous data can be clustered without additional information, so that clients with different data distributions have different models, which improves the model reasoning performance in heterogeneous scenarios. By clustering and distributing the binary mask model on the server, the communication efficiency and model performance of heterogeneous federated learning can be improved.
[0158] The communication-efficient heterogeneous federated learning training method provided in the embodiment of the present invention can be applied to scenarios such as smart mobile devices, finance, healthcare and advertising recommendation systems, and drone multi-task collaboration to improve the communication efficiency and model performance of heterogeneous federated learning.
[0159] Taking the medical field as an example, due to the need for privacy protection, patients' medical data cannot be shared. Different hospitals or medical institutions can use medical imaging data (such as X-rays, CT scans), as well as patient historical medical records, biomarkers and other data to train local models locally as clients, and then upload the binary mask of the local model to the central server. The central server groups and aggregates the model binary masks uploaded by each medical institution, and then sends the aggregated model integer masks to different medical institutions, repeating iterative training until the model converges. Since the medical institution transmits binary masks to the server, and the present invention groups and aggregates different models, the present invention can reduce the communication cost between the medical institution and the central server, while improving the accuracy of the model in predicting different diseases.
[0160] Based on the specific detection requirements in the medical scenario, local data sets can be selected and labeled to achieve different detection purposes.
[0161] A specific example of the embodiment of the present invention in the medical field is given below:
[0162] The server is a central server, and the client is a server of a medical institution; the local data set of the client includes the medical imaging data of the patient and is marked with a tumor type label; wherein the tumor type label can be the size of the tumor, the grade of classification, etc., which can be marked in advance by professionals as needed.
[0163] The preset image detection task includes: the client uses the trained local model to detect the medical image data to be tested to obtain the tumor type.
[0164] For this example, the client trains a local model based on the local data set and the server using the communication-efficient heterogeneous federated learning training method of an embodiment of the present invention. Then, when the client obtains the medical imaging data to be tested of a patient, it can input its trained local model, and the local model will output the tumor type of the medical imaging data to be tested, thereby completing the detection task.
[0165] Please refer to the previous step description for the specific training process, which will not be repeated here.
[0166] Similarly, the communication-efficient heterogeneous federated learning training method of the embodiment of the present invention can also be applied in multi-UAV collaborative task execution scenarios to improve the communication efficiency and model reasoning capabilities of the UAVs.
[0167] When a drone cluster performs tasks, such as environmental monitoring, search and rescue, drone delivery, etc., multiple drones need to work together to complete the task. Through the present invention, each drone can process its sensor data and model training locally, and then upload the binary mask of the model to the central server for grouping aggregation to improve communication efficiency and enhance the overall task execution capability of the drone cluster. Specifically, a small drone cluster can jointly train the navigation and obstacle avoidance algorithm through the present invention, and share the model binary mask to improve the overall obstacle avoidance capability and path planning optimization of the cluster. Another small drone cluster can use sensor data (such as cameras, radars, etc.) to train the target recognition algorithm locally, and summarize the updated model binary mask to the central server for improvement. Through the present invention. On the one hand, in the present invention, the drone only needs to regularly upload the binary mask of the model instead of a large amount of floating-point data, and the communication efficiency from the drone to the central server can be improved by 32 times, which significantly saves bandwidth and reduces latency, solving the problem of low communication bandwidth during drone model training; on the other hand, drones can share learning results in multi-task execution scenarios, and use the binary mask grouping aggregation idea to improve the model reasoning ability in drone multi-task collaborative scenarios, ensuring the availability and accuracy of drone models.
[0168] Based on the specific detection requirements in the medical scenario, local data sets can be selected and labeled to achieve different detection purposes.
[0169] The following is a specific example of an embodiment of the present invention performing environmental monitoring in the field of drones:
[0170] The server is a central server, and the client is an onboard computer of a drone; the local data set of the client includes environmental image data collected by the drone and marked with a target type label; wherein the target type label can be information such as the category, size, and location of the target, which can be marked in advance by professionals as needed.
[0171] The preset image detection task includes: the client uses the trained local model to detect the image data of the environment to be tested to obtain the target type.
[0172] For this example, the client trains a local model based on the local data set and the server using the communication-efficient heterogeneous federated learning training method of an embodiment of the present invention. Then, when the client obtains new environmental image data, it can input its trained local model, and the local model will output the target type of the environmental image data to be tested, thereby completing the detection task.
[0173] Please refer to the previous step description for the specific training process, which will not be repeated here.
[0174] In summary, compared with the prior art, the method of the present invention has the following beneficial effects:
[0175] First, the present invention provides a communication-efficient federated learning training method, in which the client only needs to upload binary masks to the server, without uploading floating-point parameters, which can greatly reduce the communication overhead from the client to the server. At the same time, the client can download integer clustering masks from the server, further improving communication efficiency and reducing communication costs.
[0176] Second, the present invention provides an efficient clustering learning method for heterogeneous federated learning. The present invention estimates clients with different data distributions based on the similarity of binary masks, clusters clients with heterogeneous data without additional information, and allows clients with different data distributions to have different models, thereby improving the model reasoning performance in heterogeneous scenarios.
[0177] It should be noted that, in the description of the present invention, the description with reference to the terms "one embodiment", "some embodiments", "example", "specific example", or "some examples" etc. means that the specific features, structures, materials or characteristics described in conjunction with the embodiment or example are included in at least one embodiment or example of the present invention. In this specification, the schematic representations of the above terms do not necessarily refer to the same embodiment or example. Moreover, the specific features, structures, materials or characteristics described may be combined in any one or more embodiments or examples in a suitable manner. In addition, those skilled in the art may combine and combine the different embodiments or examples described in this specification.
[0178] The above description is only a preferred embodiment of the present invention and is not intended to limit the protection scope of the present invention. Any modification, equivalent replacement, improvement, etc. made within the spirit and principle of the present invention are included in the protection scope of the present invention.
Claims
1. A communication efficient heterogeneous federated learning training method, characterized in that: Applied to a system consisting of a server and multiple clients, the method includes: S100, the server generates the initial model weight w0 and initial weight score s of the neural network model g,0 , by communicating with the client, each client obtains the initial model weight w0 to initialize the local model; the server uses the initial weight score s g,0 Calculate the initial probability mask θ g,0 , sent to all clients; S200, in the tth round of global training iteration, the server selects K t Clients participate in the federated learning model training; S300, each selected client k performs E rounds of local model training based on the local data set, and updates and optimizes the local weight score s k,t , and upload the binary mask m k,t to the server; wherein the local data set includes a plurality of images; S400, the server clusters the binary masks uploaded by all clients, and allocates cluster IDs to the corresponding binary masks to implement cluster ID allocation for the clients; S500, the server groups and aggregates the binary masks of each cluster after clustering to obtain the probability mask within the cluster and the global probability mask θ of the tth round of global training iteration g,t ; S600, the server notifies the client to perform a consistency check on the allocated cluster ID; S700, the server sends the intra-cluster probability mask to the client in the corresponding cluster, and sends the global probability mask to the client to which the cluster ID is not assigned, so that the client recovers the probability mask and performs local training until the model converges; S800, each client obtains a binary mask according to the probability mask within the cluster, and uses the obtained binary mask and the initial model weight w0 to obtain the final sub-network model, thereby obtaining a trained local model; the local model is used to complete the preset image detection task.
2. The communication efficient heterogeneous federated learning training method according to claim 1, characterized in that: S100, including: S110, the server uses the random seed S to generate the initial model weight w0 and initial weight score s of the neural network model g,0 ; Among them, the weight score represents the importance of the weight on the neural network model. The larger the value, the more beneficial the weight is to model reasoning; S120, the server sends the random seed S to all clients, so that each client constructs the same initial model weight w0 locally to initialize the local model; S130, the server end calculates the initial weight score s g,0 and logistic function to calculate the initial probability mask θ g,0 , and the initial probability mask θ g,0 Sent to each client; where the probability mask represents the probability of the weight being used in the forward propagation.
3. The communication efficient heterogeneous federated learning training method according to claim 2, characterized in that: S300, including: S310, in the tth round of global training iterations, each client k uses the logic function Logit and the global probability mask θ received from the server g,t-1 , calculate the weight score s of the local model g,t ; Among them, for the first round of global training, the global probability mask θ g,t-1 is the initial probability mask θ g,0 ; S320, start the current round of local training, each client k will set the weight score s of the local model k,t Use the logistic function to map to the probability mask θ k,t ; Among them, in the first round of local training of the tth round of global training, the weight score s of the local model k ,t is the weight score s of the local model g,t ; S330, each client k pairs of probability masks θ k,t Perform Bernoulli sampling to generate binary mask m k,t ; S340, each client k receives a binary mask m k,t And the initial model weight w0, by multiplying the corresponding elements of the matrix, calculate the sub-network model w for forward propagation k,t ; When the value of the binary mask is 1, it means that the corresponding weight participates in the forward propagation, and when the value of the binary mask is 0, it means that the corresponding weight does not participate in the forward propagation; S350, each client k uses the local data set and sub-network model w k,t , calculate its gradient value after forward propagation, and then update the local weight score s k,t ; S360, each client k repeats steps S320 to S350 E times to complete E rounds of local training and obtain an updated and optimized weight score s k,t ; S370, each client k adds the weight score s obtained in step S360 k,t Convert to optimized binary mask m k,t , and upload it to the server.
4. The communication efficient heterogeneous federated learning training method according to claim 3, characterized in that: S400, including: S410, the server receives K t Binary mask uploaded by client S420, the server receives K t Randomly select Q binary masks from the binary masks as the initial centroid of each cluster; S430, the server assigns each binary mask to the cluster where the centroid closest to the binary mask is located, and obtains a preliminary clustering result; wherein the preliminary clustering result records the ID of the client in each cluster; S440, for each cluster, a centroid is calculated so that the sum of the distances between the centroid and each binary mask of the cluster is minimized, thereby re-obtaining the centroid of each cluster; S450, repeat steps S430 to S440T1 times, and obtain K t A binary mask M K,t The final clustering result is obtained, and at the same time, a corresponding cluster ID is assigned to each client k according to the final clustering result.
5. The communication efficient heterogeneous federated learning training method according to claim 4, characterized in that: S430, including: S4301, calculate each binary mask m k To the center of mass m q The distance d q,k : Among them, L is the number of layers of the neural network model, n l represents the number of model parameters in the lth layer, Represents the binary mask m k For the parameters of the lth layer, ⊕ represents the XOR operation; S4302, for each binary mask m k , determine the cluster ID of the nearest centroid, expressed as: q=argmin q∈[Q] d q,k ; S4303, the binary mask m k Assign it to its nearest cluster q, and its operation is expressed as: C[q].insert(k); Among them, C represents the clustering result, and C[q] records the IDs of all clients in cluster q.
6. The communication efficient heterogeneous federated learning training method according to claim 5, characterized in that: S500, including: S510, the server aggregates the binary masks in each cluster q to obtain the cluster probability mask θ after cluster q is aggregated. g,t [q], calculated as S520, the server aggregates all intra-cluster probability masks to obtain a global probability mask for the tth round of global training iteration 7. The communication efficient heterogeneous federated learning training method according to claim 6, characterized in that: S600, including: S610, the server uses the aggregated probability mask θ g,t [q] Perform Bernoulli sampling to obtain Q clusters of binary masks M t , then M t Sent to each client; where M t ={m q,t |m q,t = Bern(Θ g,t [q])} q∈[Q] ;Bern means Bernoulli sampling; S620, each client generates a binary mask M in each cluster. t [q] Obtain the sub-network model of forward propagation, and then calculate its loss value based on the local data set, and obtain the cluster with the smallest loss value to determine the locally inferred cluster ID; wherein, when the cluster ID locally inferred by the client is consistent with the cluster ID assigned by the server, the corresponding client will no longer participate in the process of clustering and assigning cluster IDs. When the cluster ID locally inferred by the client is inconsistent with the cluster ID assigned by the server, the cluster ID locally inferred by the client shall prevail.
8. The communication efficient heterogeneous federated learning training method according to claim 7, characterized in that: S700, including: S710, the server sends an integer and integers Sent to the client with cluster ID q; represents the number of clients in each cluster, represents the sum of the binary masks in each cluster; S720, each client receives and Then recover the global probability mask θ g,t , and perform local training T times to make the model converge.
9. The communication efficient heterogeneous federated learning training method according to claim 8, characterized in that: S800, including: S810, each client performs a probability mask Θ on the cluster g,T [q] Perform Bernoulli sampling to obtain a binary mask m q = Bern(Θ g,T [q]); S820, each client generates a binary mask m according to the binary mask m. q And the initial model weight w0, by multiplying the corresponding elements of the matrix, a sub-network model with reasoning ability is generated to constitute a trained local model.
10. The communication efficient heterogeneous federated learning training method according to claim 1 or 9, characterized in that: The server is a central server, and the client is a server of a medical institution; the local data set of the client includes the medical imaging data of the patient and is marked with a tumor type label; The preset image detection task includes: the client uses the trained local model to detect the medical image data to be tested to obtain the tumor type.
Citation Information
Patent Citations
A communication method and system for heterogeneous federated learning in industrial Internet of Things
CN118101501B
Federal learning system and method based on model pruning and transmission compression optimization
CN115564062A
Federal learning sparse training method and system based on comparative learning
CN115829027A
Federal learning method and device for adaptive communication in dynamic bandwidth scene
CN117938690A
Federal learning sampling method and system under heterogeneous wireless network
CN118446331A
Cited By
Heterogeneous federal model adjusting method based on importance sampling
CN120725101A
A heterogeneous federated model adjustment method based on importance sampling
CN120725101B