A heterogeneous data federated learning method based on an improved aggregation algorithm

CN118821909BActive Publication Date: 2026-09-25NANJING UNIV OF POSTS & TELECOMM
View PDF 2 Cites 0 Cited by

Patent Information

Application Number
CN202410823922.6
Authority / Receiving Office
CN · China
Patent Type
Patents(China)
Current Assignee / Owner
Filing Date
2024-06-25
Publication Date
2026-09-25
Estimated Expiration
2044-06-25

AI Technical Summary

Technical Problem

SCAFFOLD通过控制变量减少训练偏差,但其维护成本高且通信开销大

Benefits of technology

[0048]本发明提出一种基于改进聚合算法的异构数据联邦学习方案,适用于在数据非独立同分布情况较为严重的场景下进行联邦学习。该方案建立在传统联邦平均算法(FedAVG)的基础上,首先评估参与本轮次学习的客户端相较于上一轮次参与客户端的重合度;然后根据各客户端上传模型的更新差异计算各参数的重要性得分,再通过归一化各客户端之间相同参数位置上的得分得到加权系数;最后根据加权系数和模型重合度对各客户端模型进行加权得到更新后的全局网络模型,在保证模型收敛的同时降低客户端本地的计算复杂度,能够有效缓解训练数据分布不平衡导致的全局模型训练不足和通信开销大的问题。

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN118821909B_ABST
    Figure CN118821909B_ABST
Patent Text Reader

Abstract

The application discloses a heterogeneous data federation learning method based on an improved aggregation algorithm, and relates to the technical field of machine learning. Firstly, a server initializes a global model and allocates a unique identifier to each client. Then, the server randomly selects part of the clients and sends the global model. Then, each client trains the model using local data and uploads the trained model parameters to the server. Subsequently, the server weights and aggregates the received model parameters to obtain a new global model. Finally, the server tests the global model to determine whether the learning process is stopped. If so, the global model is distributed and the previous steps are repeated. Otherwise, the communication is ended and the global model is broadcast. The application is suitable for multiple scenarios, such as anomaly detection in industrial internet security and patient data sharing in smart medical treatment, and can effectively alleviate the problems of insufficient global model training and large communication overhead caused by unbalanced training data distribution.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] This invention relates to the field of machine learning technology, and in particular to a heterogeneous data federated learning method based on an improved aggregation algorithm. Background Technology

[0002] With the widespread adoption of terminal devices such as mobile phones, tablets, and home appliances, network data traffic has exploded. Extracting key features from massive amounts of data has become a major challenge. Machine learning, especially deep learning, can effectively process and analyze large-scale data, offering a glimmer of hope for solving the problem of feature extraction from big data. However, using large-scale data for machine learning faces the problem of "data silos," where data is monopolized by a few industry giants, making data sharing difficult both within and outside the industry.

[0003] To address this issue, McMahan et al. proposed the concept of federated learning in 2016. Federated learning is a machine learning framework based on distributed datasets that allows clients to collaboratively train models under the coordination of a central server. Relevant information about the model can be exchanged among the participants, but local training data remains on-premises, thus enabling effective machine learning modeling while protecting privacy.

[0004] While federated learning addresses data privacy concerns, it still faces challenges posed by the non-identical distribution (non-IID) of data. Due to varying user habits and preferences, the types and distribution of data differ significantly across devices, leading to non-identical distribution of training data. In this scenario, local models can only fit local data well, but perform poorly on the global dataset, impacting the accuracy of the overall model. Therefore, mitigating the impact of non-identical data on federated learning has become a crucial issue.

[0005] To address this issue, researchers have proposed solutions such as FedProx and SCAFFOLD. FedProx addresses data heterogeneity by adding a regularization term, but it requires fine-tuning and has limited effectiveness in extreme cases. SCAFFOLD reduces training bias by controlling variables, but it has high maintenance costs and significant communication overhead. Therefore, solving the extreme non-independent and identically distributed problem while maintaining low communication overhead remains a challenge. Summary of the Invention

[0006] The technical problem to be solved by the present invention is to provide a heterogeneous data federated learning method based on an improved aggregation algorithm, which addresses the shortcomings of the prior art.

[0007] To solve the above-mentioned technical problems, the technical solution adopted by the present invention is as follows:

[0008] A heterogeneous data federated learning method based on an improved aggregation algorithm includes the following steps:

[0009] S1. Construct a federated learning network; a single server and Several clients collaborate to build a federated learning network, with both the server and clients possessing local datasets. The server initializes a neural network model. As a global network model, it assigns a unique identifier to each client. ;

[0010] S2. Select a client; the server selects randomly. One client, and the current global network model parameters Send to the selected clients;

[0011] S3, Training Client; Selected Client Receive the global network model from the server and use the local private dataset. This global network model The local network model is obtained through training. The trained local network model parameters Training dataset Size and client identifier Upload them to the server together;

[0012] S4. Server Aggregation: The server receives model parameters from each client. Then, the model parameters from each client are weighted and aggregated to obtain new global network model parameters. Based on this, a new global model is obtained. ;

[0013] S5. Testing the model's accuracy: The server uses the local dataset to test the accuracy of the current global model and compares it with a threshold. If the accuracy is less than the threshold, it returns to S2 and continues training; if the accuracy is greater than or equal to the threshold, learning ends, and the final global model is generated. The broadcast was sent to all participating clients.

[0014] As a further preferred embodiment of the present invention, S3 includes S31, calculating the loss value during the model training process; the client receives the global network model from the server. Then, it was used as the current local network model. Client Local private dataset Represented as:

[0015] (1)

[0016] in, Represents the dataset The Middle Feature vectors of data, Represents the dataset The Middle The actual labels corresponding to each data point; Represents the dataset Length;

[0017] Dataset The Feature vector of data Input to local model In the process, the predicted label value is calculated. , represented as:

[0018] (2)

[0019] in, Indicates client The The predicted output value of each data point. This represents the activation function. Indicates client The network model parameter vector, Indicates client The Feature vectors of data;

[0020] Using the cross-entropy function as the loss function, the client... No. Label prediction loss value of data , represented as:

[0021] (3)

[0022] in, Represents the parameters of a given network model and data samples Label prediction loss value, Indicates the client No. The actual label value of each data point Indicates the client No. The predicted label value for each data point.

[0023] As a further preferred embodiment of the present invention, S3 further includes S32, updating local network model parameters; client. By batch inputting training data into the local model, the local model parameters are iteratively optimized using the gradient descent algorithm. The formula for a single update of the local model parameters is:

[0024] (4)

[0025] in, Indicates the learning rate. Represents the loss function Regarding model parameters In data samples gradient on, This indicates the batch size of data in each training session.

[0026] As a further preferred embodiment of the present invention, in step S4, the FedPW algorithm is used to weighted aggregate the model parameters of each client; the FedPW algorithm includes the following steps:

[0027] S41. Calculate the overlap of the selected clients;

[0028] S42. Calculate and normalize the differences in client model parameters;

[0029] S43. Calculate the aggregate weights of the parameters of each client model;

[0030] S44. Weighted aggregation of model parameters for each client.

[0031] As a further preferred embodiment of the present invention, the specific steps of S41 are as follows:

[0032] Calculate the current number The selected client in the learning round and the first Overlap of selected clients in each learning round , represented as:

[0033] (5)

[0034] in, The number of clients selected in the current round that are the same as those selected in the previous round. , The total number of clients participating in federated learning.

[0035] As a further preferred embodiment of the present invention, the specific steps of S42 are as follows:

[0036] Calculate the number of clients participating in federated learning in the current round. Network model parameter difference vector for:

[0037] (6)

[0038] in, Indicates the client in the current learning round The network model parameter vector uploaded to the server Indicates the client In the network model, the input layer of the first The nth neuron and the output layer The connection rights of each neuron; This represents the local network model parameter vector on the server during the current learning round. This represents the input layer of the global network model. The nth neuron and the output layer The connection rights of each neuron;

[0039] Based on the network model parameter difference vector of the client, calculate each parameter in the client network model. Corresponding importance score , represented as:

[0040] (7)

[0041] in, , Indicates the client In the network model, the first enter, Output link parameter differences, Represented as the first The size of the training dataset for each client.

[0042] As a further preferred embodiment of the present invention, the specific step of S43 is as follows: traverse k to obtain all clients. Network model parameter importance score Calculate the network model parameters for each client. Aggregate weights for:

[0043] ;(8)

[0044] As a further preferred embodiment of the present invention, the specific step of S44 is as follows: the server aggregates weights using client network model parameters. According to equation (9), the parameters of each client model are... We perform weighted calculations to obtain the aggregated global network model parameters. ,in This represents the input layer of the global network model. The nth neuron and the output layer The connection weights of each neuron are calculated as follows:

[0045] (9)

[0046] The server is based on the aggregated network model parameters. Update the global neural network R is the number of randomly selected clients, e is a natural number, and α represents the degree of overlap.

[0047] The present invention has the following beneficial effects:

[0048] This invention proposes a heterogeneous data federated learning scheme based on an improved aggregation algorithm, suitable for federated learning in scenarios where data is severely non-independent and identically distributed. This scheme is built upon the traditional Federated Average (FedAVG) algorithm. First, it evaluates the overlap between clients participating in the current learning round and those participating in the previous round. Then, it calculates the importance score of each parameter based on the update differences of the models uploaded by each client, and obtains weighting coefficients by normalizing the scores at the same parameter positions among clients. Finally, it weights the models of each client based on the weighting coefficients and model overlap to obtain the updated global network model. This approach ensures model convergence while reducing the computational complexity of the client's local computation, effectively alleviating the problems of insufficient global model training and high communication overhead caused by imbalanced training data distribution. Attached Figure Description

[0049] Figure 1 This is the system model in the example of the present invention;

[0050] Figure 2 This is the FedPW algorithm flow in the example of this invention;

[0051] Figure 3 This describes the data distribution of each client in this invention example;

[0052] Figure 4 , Figure 5 , Figure 6 , Figure 7 , Figure 8 , Figure 9 , Figure 10 and Figure 11 These are simulation results from examples of this invention. Detailed Implementation

[0053] The present invention will now be described in further detail with reference to the accompanying drawings and specific preferred embodiments.

[0054] The specific dimensions used in this embodiment are merely illustrative of the technical solution and do not limit the scope of protection of this invention.

[0055] set up Figure 1 System scenario shown:

[0056] like Figure 1 As shown, consider a typical federated learning system framework, which includes a server and Each client maintains its own private dataset locally, and these datasets may have different data distributions.

[0057] Combination Figure 1 System model, Figure 2 The FedPW algorithm flow and the specific implementation steps of this scheme are described in detail below:

[0058] S1. Construct a federated learning network; a single server and Several clients collaborate to build a federated learning network; the server and clients each have their own local datasets, and the datasets of different clients follow different distributions; the server initializes a neural network model. As a global network model, it assigns a unique identifier to each client. .

[0059] The model initialization is described in detail below:

[0060] Set the initial learning rounds The server generates an initial neural network model. As the current global network model, a unique identifier is assigned and distributed to all clients, and the clients store the unique identifier after receiving it.

[0061] S2. Select Client: The server selects randomly. Client and the current global network model parameters Send to the selected clients;

[0062] The specific description of the clients that choose to participate is as follows:

[0063] from Randomly selected from clients Each client participated in the learning process. and the current global network model parameters Send it to the selected client.

[0064] S3, Training Client: Selected Client Receive the global network model from the server and use the local private dataset. This network model Train to obtain a local model The trained network model parameters Training dataset Size Client identifier Upload it to the server together; here This is the client's ID.

[0065] The local model training is described in detail below:

[0066] S31. Calculate the loss value during model training:

[0067] Set the current learning cycle The client receives the global network model from the server. Then, it was used as the current local network model. .

[0068] Client Local private dataset Represented as:

[0069] (1)

[0070] in, Represents the dataset The Middle Feature vectors of data, Represents the dataset The Middle The actual labels corresponding to each data point; Represents the dataset The length.

[0071] Dataset The Feature vector of data Input to local model In the process, the predicted label value is calculated. , represented as:

[0072] (2)

[0073] in, Indicates the client The The predicted output value of each data point. This represents the activation function (usually the sigmoid function). Indicates the client The network model parameter vector, Indicates the client The The feature vector of each data point.

[0074] Using the cross-entropy function as the loss function, the client... No. Label prediction loss value of data , represented as:

[0075] (3)

[0076] in, Represents the parameters of a given network model and data samples Label prediction loss value, Indicates the client No. The actual label value of each data point Indicates the client No. The predicted label value for each data point.

[0077] S32. Update local network model parameters:

[0078] Client By batch inputting training data into the local model, the local model parameters are iteratively optimized using the gradient descent algorithm. To minimize the loss function, the formula for a single update of the local model parameters is:

[0079] (4)

[0080] in, Indicates the learning rate. Represents the loss function Regarding model parameters In data samples gradient on, This indicates the batch size of data in each training session.

[0081] S4. Server Aggregation: The server receives model parameters from each client. Then, the improved FedPW algorithm is used to weighted aggregate the model parameters of each client to obtain new global model parameters. Based on this, a new global model is obtained. ;

[0082] The improved FedPW aggregation algorithm is described in detail below:

[0083] S41. Calculate the overlap of the selected clients;

[0084] Calculate the current number The selected client in the ( ) learning round and the ( ) Overlap of selected clients in each learning round , represented as:

[0085] (5)

[0086] in, The number of clients selected in the current round that are the same as those selected in the previous round. , The total number of clients participating in federated learning.

[0087] S42. Calculate and normalize the differences in client model parameters:

[0088] Calculate the number of clients participating in federated learning in the current round. Network model parameter difference vector for:

[0089] (6)

[0090] in, Indicates the client in the current learning round The network model parameter vector uploaded to the server Indicates the client In the network model, the input layer of the first The nth neuron and the output layer The connection rights of each neuron; This represents the local network model parameter vector on the server during the current learning round. This represents the input layer of the global network model. The nth neuron and the output layer The connection rights of each neuron.

[0091] Based on the network model parameter difference vector of the client, calculate each parameter in the client network model. Corresponding importance score , represented as:

[0092] (7)

[0093] in, , Indicates the client In the network model, the first enter, Output link parameter differences, Represented as the first The size of the training dataset for each client.

[0094] S43. Calculate the aggregate weights of the parameters for each client model:

[0095] Iterate through k to get all clients Network model parameter importance score Calculate the network model parameters for each client according to equation (8). Aggregate weights for:

[0096] (8)

[0097] S44. Weighted aggregation of model parameters for each client:

[0098] The server aggregates weights using improved client network model parameters. According to equation (9), the parameters of each client model are... We perform weighted calculations to obtain the aggregated global network model parameters. ,in This represents the input layer of the global network model. The nth neuron and the output layer The connection weights of each neuron are calculated as follows:

[0099] (9)

[0100] The server is based on the aggregated network model parameters. Update the global neural network ; The number of clients is randomly selected. It is a natural number. The overlap is represented by the piecewise function calculated earlier.

[0101] S5, Model Testing and Judgment: The server uses the local dataset to test the accuracy of the current global model and determines whether to continue to the next round of learning. If yes, return to S2; otherwise, the learning ends, and the final global model is released. The broadcast was sent to all participating clients.

[0102] Accuracy is calculated by inputting server data into the global model, comparing the predicted label and the true label for each data point, and if they match, the data is considered correct. The accuracy is then calculated as the percentage of all correct data points out of the total number of correct data points.

[0103] The specific description of the network model testing and evaluation scheme is as follows:

[0104] S51, Server Model Testing;

[0105] The server uses local test data to test the updated global network model. The calculation model estimates metrics such as loss value, accuracy, precision, recall, and F1 score.

[0106] S52. Determine whether to proceed to the next learning round;

[0107] Within the set training epochs, if the accuracy of the global model on the server test dataset reaches a pre-set threshold for three consecutive epochs, the system will consider the model to have converged, terminate training, and release the final global model. Broadcast to all participating clients. Otherwise, return to S2. Additionally, if the set number of training epochs is reached, training will terminate, and the final global model will be shared. The broadcast is sent to all participating clients. The algorithm has ended.

[0108] The proposed scheme is simulated using PyTorch, and the performance of the FedPW algorithm is compared with similar algorithms such as FedAVG, FedProx, and SCAFFOLD. Simulations are performed on four dimensions: accuracy, precision, recall, and F1 score. The results are as follows: Figure 4 , Figure 5 , Figure 6 , Figure 7 , Figure 8 , Figure 9 , Figure 10 and Figure 11 As shown.

[0109] The simulation parameters are set as follows:

[0110] 1) Model: The network model is a three-layer feedforward neural network with 41 inputs and 5 outputs;

[0111] 2) Clients: Total number of clients K=10, randomly select the number of participating clients considering two cases: R=10 and R=5;

[0112] 3) Training parameters: global_epochs=100, local_epochs=2, batch_size=32, learning ratelr=0.001;

[0113] 4) Training data: NSL-KDD dataset, data distribution Dir(0.2).

[0114] observe Figure 4 , Figure 5 , Figure 6 and Figure 7 It can be seen that when selecting all clients to participate in learning, the convergence time of the FedPW algorithm in this scheme is better than that of the traditional FedAVG algorithm, as well as the improved FedProx algorithm and SCAFFOLD algorithm, across the four dimensions of accuracy, precision, recall, and F1 score.

[0115] observe Figure 8 , Figure 9 , Figure 10 and Figure 11 It can be seen that when selecting a subset of clients to participate in learning, the FedPW algorithm outperforms the traditional FedAVG algorithm in convergence time and accuracy across the four dimensions of accuracy, precision, recall, and F1 score. It is comparable to the improved FedProx algorithm and slightly inferior to the SCAFFOLD algorithm. However, it surpasses both the FedProx and SCAFFOLD algorithms in terms of client-side computational complexity and has lower single-communication overhead than the improved SCAFFOLD algorithm.

[0116] This invention proposes a heterogeneous data federated learning scheme based on an improved aggregation algorithm, suitable for scenarios where data is not independent and identically distributed. The scheme is built upon the traditional Federated Average (FedAVG) algorithm. First, it evaluates the overlap between clients participating in the current learning round and those participating in the previous round. Then, it calculates the importance score of each parameter based on the update differences of the models uploaded by each client, and obtains weighting coefficients by normalizing the scores at the same parameter positions among clients. Finally, it weights the models of each client based on the weighting coefficients and model overlap to obtain the updated global network model. This approach ensures model convergence while reducing the computational complexity of the client's local processing and the communication overhead during transmission.

[0117] The preferred embodiments of the present invention have been described in detail above. However, the present invention is not limited to the specific details of the above embodiments. Within the scope of the technical concept of the present invention, various equivalent transformations can be made to the technical solutions of the present invention, and these equivalent transformations all fall within the protection scope of the present invention.

Claims

1. A heterogeneous data federated learning method based on an improved aggregation algorithm, characterized in that: Includes the following steps: S1. Construct a federated learning network; A federated learning network is constructed through collaboration between a single server and K clients. The server and clients each possess local datasets, and the server initializes a neural network model θ. global As a global network model, and assigning a unique identifier U to each client. k ; S2. Select clients; the server randomly selects R clients, R≤K, and sets the current global network model θ. global The parameter w gobal Send to the selected clients; S3, training client; Selected client U k Receive the global network model from the server, using the local private dataset D. k For this global network model θ global The local network model θ is obtained through training. k The trained local network model θ k The parameter w k Training dataset D k Size N k and client identifier U k Upload them to the server together; S4. Server Aggregation: The server receives model parameters w from each client. k Then, the model parameters of each client are weighted and aggregated to obtain the new global network model parameters w. gobal Based on this, a new global model θ is obtained. global ; S5. Testing the model's accuracy: The server uses the local dataset to test the accuracy of the current global model and compares it with a threshold. If the accuracy is less than the threshold, it returns to S2 and continues training; if the accuracy is greater than or equal to the threshold, learning ends, and the final global model θ is generated. global The broadcast was sent to all participating clients.

2. The heterogeneous data federated learning method based on an improved aggregation algorithm according to claim 1, characterized in that: S3 includes S31, calculating the loss value during model training; the client receives the global network model θ from the server. global Then, it is used as the current local network model θ k , Client U k Local private dataset D k Represented as: Where, x m Represents dataset D k The feature vector of the m-th data point, y m Represents dataset D k The actual label corresponding to the m-th data item; N k Represents dataset D k Length; Dataset D k The feature vector x of the m-th data m Input to local model θ k In the process, the predicted label value is calculated. Represented as: in, Indicates client U k The predicted output value of the m-th data point, where σ represents the activation function, w k Indicates client U k The network model parameter vector, Indicates client U k The feature vector of the m-th data; Using the cross-entropy function as the loss function, calculate the client U. k Label prediction loss value of the m-th data point Represented as: in, This represents the given network model parameters w k and data samples Label prediction loss value Indicates client U k The true label value of the m-th data item. Indicates client U k The predicted label value for the m-th data point.

3. The heterogeneous data federated learning method based on an improved aggregation algorithm according to claim 2, characterized in that: S3 also includes S32, updating local network model parameters; client U k By batch inputting training data into the local model, the gradient descent algorithm is used to iteratively optimize the local model parameters w. k The formula for a single update of the local model parameters is: Where η represents the learning rate. The loss function L represents the loss function with respect to the model parameters w. k In data samples The gradient on the training plane, where n represents the batch size of data in each training session.

4. The heterogeneous data federated learning method based on an improved aggregation algorithm according to claim 1, characterized in that: In step S4, the FedPW algorithm is used to weighted aggregate the model parameters of each client; the FedPW algorithm includes the following steps: S41. Calculate the overlap of the selected clients; S42. Calculate and normalize the differences in client model parameters; S43. Calculate the aggregate weights of the parameters of each client model; S44. Weighted aggregation of model parameters for each client.

5. A heterogeneous data federated learning method based on an improved aggregation algorithm according to claim 4, characterized in that: The specific steps of S41 are as follows: The overlap α between the client selected in the current learning round t and the client selected in the (t-1)th learning round is calculated as follows: Where r is the number of clients selected in the current round that are the same as those selected in the previous round, r≤R, and K is the total number of clients participating in federated learning.

6. A heterogeneous data federated learning method based on an improved aggregation algorithm according to claim 5, characterized in that: The specific steps of S42 are as follows: Calculate the number of clients U participating in federated learning in the current round. k Network model parameter difference vector Δw k for: Δw k =in gobal -In k (6) in, Indicates the current learning round for client U k The network model parameter vector uploaded to the server Indicates client U k The connection weights between the j-th neuron in the input layer and the i-th neuron in the output layer in the network model; This represents the local network model parameter vector on the server during the current learning round. This represents the connection weight between the j-th neuron in the input layer and the i-th neuron in the output layer in the global network model; Based on the network model parameter difference vector of the client, calculate each parameter in the client network model. Corresponding importance score Represented as: in, Indicates client U k In the network model, the parameter difference values ​​between the j-th input and i-th output links, N k This represents the size of the training dataset for the k-th client.

7. A heterogeneous data federated learning method based on an improved aggregation algorithm according to claim 6, characterized in that: The specific steps of S43 are as follows: Traverse k to obtain all client Us k Network model parameter importance score Calculate the network model parameters for each client Aggregate weights for:

8. A heterogeneous data federated learning method based on an improved aggregation algorithm according to claim 7, characterized in that: The specific steps of S44 are as follows: The server aggregates weights using client network model parameters. Based on equation (9), the parameters of each client model are... We perform weighted calculations to obtain the aggregated global network model parameters. in The connection weight between the j-th neuron in the input layer and the i-th neuron in the output layer in the global network model is represented by: The server is based on the aggregated network model parameters w global Update the global neural network θ global R is the number of randomly selected clients, e is a natural number, and α represents the degree of overlap.

Citation Information

Patent Citations

  • Federal continuous learning training method based on memory playback and differential privacy

    CN115081532A

  • Asynchronous federal learning method and system based on T-Step aggregation algorithm

    CN115374853A