Hierarchical compromise model federal learning method and device based on momentum

By introducing a hierarchical trade-off model federated learning method based on momentum in federated learning, the problems of communication bottlenecks, generalization ability and personalized trade-offs, and low model convergence efficiency in federated learning are solved, and more efficient model training and better adaptability are achieved.

CN120046754AActive Publication Date: 2025-05-27SUN YAT SEN UNIV
View PDF 3 Cites 0 Cited by

Patent Information

Application Number
CN202510060304.5
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2025-01-15
Publication Date
2025-05-27
Estimated Expiration
2045-01-15

AI Technical Summary

Technical Problem

In federated learning in massive data scenarios, the communication bottleneck between computing nodes and central nodes, the trade-off between system model generalization capabilities and personalization, and the low convergence efficiency of model.

Method used

The federated learning method of hierarchical tradeoff model is adopted based on momentum, and the local model, group model and global model are trained by updating local momentum, group momentum and global momentum in the local device, and weighted average aggregation in the edge node and the central server.

Benefits of technology

The overall model adaptability and learning efficiency of federated learning are improved, and the generalization and personalization capabilities of the model are enhanced through two-way collaborative optimization strategies, and communication overhead is reduced.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN120046754A_ABST
    Figure CN120046754A_ABST
Patent Text Reader

Abstract

The invention discloses a hierarchical compromise model federal learning method based on momentum, which comprises the following steps of: updating local momentum, group momentum and global momentum in local equipment according to a group model and a global model issued by a corresponding edge node, correspondingly obtaining first local momentum, first group momentum and first global momentum, training a local model according to the first local momentum; according to the first group of momentum and the first global momentum, a second group of momentum and a second global momentum are obtained through aggregation in the edge node, a group model is trained according to the second group of momentum, and according to the second global momentum, a third global momentum is obtained through aggregation in the central server; and training a global model according to the third global momentum until the global model obtained by the latest round of iteration is converged, and ending the iteration. According to the invention, the learning efficiency of the federal learning method can be improved.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention relates to the field of software engineering, and in particular, to a hierarchical compromise model federated learning method and device based on momentum. Background Art

[0002] With the popularization of technologies such as big data, machine learning, and large models, distributed learning in the scenario of massive data has become the focus of research and the industrial community. Federated learning is an encrypted distributed learning paradigm that allows multiple participants (such as user devices, edge devices, or cloud center servers, etc.) to jointly train a model without sharing the original data. The core idea of federated learning is to decentralize the process of model training to local devices, rather than centralizing the data to a central server for processing. This enables the original data to always remain local, and the privacy of the participants is protected to the greatest extent.

[0003] Although federated learning has significant advantages in protecting data privacy and efficiently utilizing distributed resources, it still faces many challenges in practical applications. Communication bottleneck between computing nodes and the central node: The communication overhead increases. Due to the different data distributions of each node, the required training time and the number of training steps are also different. Some nodes may need more local training iterations to achieve a model performance similar to that of other nodes, which leads to an imbalance in the frequency and quantity of data uploaded by different nodes. Secondly, there is a trade-off between the generalization ability and personalization of the system model. Since the data distributions of each node may vary greatly, a single global model may not be able to well adapt to the specific data characteristics of each node, resulting in limited generalization ability of the global model. Finally, the model convergence efficiency is low: In a data heterogeneous scenario, the gradients of each node may vary greatly, which leads to inconsistent directions of global gradient updates during the aggregation process, thus delaying the convergence speed of the model and even possibly causing the model to fall into a local optimum. There may be significant differences in the data volume and computing power of different nodes. Stronger nodes may hope to update the model faster, while weaker nodes may become the system bottleneck. This asymmetry further reduces the overall convergence efficiency of the model. Summary of the Invention

[0004] The present invention aims to overcome the above-mentioned defects of the prior art and provides a hierarchical compromise model federated learning method and device based on momentum, which can improve the learning efficiency of the federated learning method.

[0005] An embodiment of the present invention provides a hierarchical compromise model federated learning method and device based on momentum, including the following steps:

[0006] Update the local momentum, group momentum, and global momentum in the local device according to the group model and global model sent by the corresponding edge node, and correspondingly obtain the first local momentum, first group momentum, and first global momentum, and train the local model according to the first local momentum; where one edge node corresponds to several local devices;

[0007] Aggregate the first group momentum and the first global momentum in the edge node to obtain the second group momentum and the second global momentum, and train the group model according to the second group momentum, and then send the trained group model and the global model sent by the corresponding central server to the several local devices corresponding to the edge node; where one central server corresponds to several edge nodes;

[0008] Aggregate the third global momentum in the central server according to the second global momentum, train the global model according to the third global momentum, and after sending the trained global model to the several edge nodes corresponding to the central server, start the next round of iteration, and end the iteration until it is determined that the global model obtained in the latest round of iteration converges.

[0009] Further, the updating of the local momentum, group momentum, and global momentum in the local device according to the group model and global model sent by the corresponding edge node, and correspondingly obtaining the first local momentum, first group momentum, and first global momentum specifically includes:

[0010] Update the local momentum according to a preset random subset and a preset local gradient calculated by the local model, and correspondingly obtain the first local momentum; where the update formula of the local momentum is specifically:

[0011]

[0012] where a is the first local momentum, b is the local momentum, γ is a preset update coefficient, is the preset local gradient;

[0013] Update the group momentum and global momentum respectively according to the group model and global model sent by the corresponding edge node, and correspondingly obtain the first group momentum and first global momentum.

[0014] Further, the training of the local model according to the first local momentum specifically includes:

[0015] Update the local model by the momentum gradient descent method according to the first local momentum;

[0016] Correct the updated local model according to the group model sent by the corresponding edge node to obtain the trained local model.

[0017] Preferably, correcting the updated local model according to the group model sent by the corresponding edge node specifically includes:

[0018] Calculating the first squared norm between the updated local model and the group model sent by the corresponding edge node;

[0019] Using the first squared norm as a correction term to correct the updated local model; where the specific formula for the correction is:

[0020]

[0021] where, v t is the updated local model in the t-th iteration round, u t ′ is the group model sent by the corresponding edge node in the t-th iteration round, η is a preset learning rate, λ is a preset correction coefficient, ||v t -u t || 2 is the first squared norm.

[0022] Further, aggregating to obtain a second group of momentum and a second global momentum in the edge node according to the first group of momentum and the first global momentum specifically includes:

[0023] Aggregating to obtain the second group of momentum according to the first group of momentum by the weighted average method; where the aggregation weight for aggregating the first group of momentum is the ratio of the sample data volume of the dataset contained in the local device corresponding to the first group of momentum to the sample data volume of the dataset contained in the edge node;

[0024] Aggregating to obtain the second global momentum according to the first global momentum by the weighted average method; where the aggregation weight for aggregating the first global momentum is the ratio of the sample data volume of the dataset contained in the local device corresponding to the first global momentum to the sample data volume of the dataset contained in the edge node.

[0025] Further, training the group model according to the second group of momentum specifically includes:

[0026] Updating the group model by the momentum gradient descent method according to the second group of momentum;

[0027] Correcting the updated group model according to the global model sent by the corresponding central server to obtain the trained group model.

[0028] Preferably, correcting the updated group model according to the global model sent by the corresponding central server specifically includes:

[0029] Calculate the second squared norm between the updated group model and the global model sent by the corresponding central server;

[0030] Use the second squared norm as a correction term to correct the updated group model; where the specific formula for the correction is:

[0031] u t+1 = u t - d - λ||u t - w t || 2

[0032] where, u t is the updated group model in the t-th iteration round, w t is the global model sent by the corresponding central server in the t-th iteration round, d is the second group momentum, λ is a preset correction coefficient, and ||u t - w t || 2 is the second squared norm.

[0033] Furthermore, the aggregating the second global momentum in the central server to obtain a third global momentum and training the global model according to the third global momentum specifically includes:

[0034] Aggregate the second global momentum to obtain the third global momentum by the weighted average method; where the aggregation weight for aggregating the second group momentum is the ratio of the sample data volume of the dataset contained in the corresponding edge node to the sample data volume of the dataset contained in the central server;

[0035] Train the global model by the momentum gradient descent method according to the third global momentum.

[0036] Furthermore, the convergence condition of the global model is specifically:

[0037] Calculate the loss function value of the global model sent by the corresponding edge node according to the preset loss function in the local device;

[0038] Aggregate the loss function values of the global model calculated by all the local devices to obtain a global loss value;

[0039] When the global loss value reaches the minimum, determine that the global model converges.

[0040] Another embodiment of the present invention provides a momentum-based hierarchical compromise model federated learning device, including: a local module, an edge module, and a central module;

[0041] The local module is used to update the local momentum, group momentum, and global momentum in the local device according to the group model and global model sent by the corresponding edge node, and correspondingly obtain the first local momentum, the first group momentum, and the first global momentum, and train the local model according to the first local momentum; where one edge node corresponds to a plurality of the local devices;

[0042] The edge module is used to aggregate the second group momentum and the second global momentum in the edge node according to the first group momentum and the first global momentum, and train the group model according to the second group momentum, and then send the trained group model and the global model sent by the corresponding central server to the plurality of local devices corresponding to the edge node; where one central server corresponds to a plurality of the edge nodes;

[0043] The central module is used to aggregate the third global momentum in the central server according to the second global momentum, train the global model according to the third global momentum, and start the next round of iteration after sending the trained global model to the plurality of edge nodes corresponding to the central server, until it is determined that the global model obtained in the latest round of iteration converges, and then end the iteration.

[0044] Compared with the prior art, the beneficial effects of the present invention are as follows:

[0045] By introducing a two-way collaborative optimization strategy, each level of the model can not only correct and improve the generalization ability based on the superior model, but also improve the model personalization ability based on momentum aggregation, thereby further improving the overall model adaptation ability and learning efficiency of federated learning. Description of the Drawings

[0046] Figure 1 It is a schematic flowchart of a method for hierarchical compromise model federated learning based on momentum provided by an embodiment of the present invention.

[0047] Figure 2 It is a schematic structural diagram of a hierarchical compromise model federated learning framework based on momentum provided by an embodiment of the present invention.

[0048] Figure 3 It is a schematic structural diagram of a device for hierarchical compromise model federated learning based on momentum provided by another embodiment of the present invention. Detailed Embodiments

[0049] The drawings are only for illustrative purposes and should not be construed as a limitation of this patent;

[0050] For those skilled in the art, it is understandable that some well-known structures and their descriptions in the drawings may be omitted.

[0051] The technical solution of the present invention will be further described below with reference to the accompanying drawings and embodiments.

[0052] Referring to Figure 1 , which is a schematic flowchart of a hierarchical compromise model federated learning method based on momentum provided by an embodiment of the present invention, including the following steps:

[0053] S1: Update the local momentum, group momentum, and global momentum in the local device according to the group model and global model sent by the corresponding edge node, and correspondingly obtain the first local momentum, the first group momentum, and the first global momentum, and train the local model according to the first local momentum; wherein, one edge node corresponds to a plurality of the local devices;

[0054] S2: Aggregate in the edge node to obtain the second group momentum and the second global momentum according to the first group momentum and the first global momentum, and train the group model according to the second group momentum, and then send the trained group model and the global model sent by the corresponding central server to the plurality of local devices corresponding to the edge node; wherein, one central server corresponds to a plurality of the edge nodes;

[0055] S3: Aggregate in the central server to obtain the third global momentum according to the second global momentum, train the global model according to the third global momentum, and after sending the trained global model to the plurality of edge nodes corresponding to the central server, start the next round of iteration, and end the iteration until it is determined that the global model obtained in the latest round of iteration converges.

[0056] For step S1, specifically, the updating of the local momentum, group momentum, and global momentum in the local device according to the group model and global model sent by the corresponding edge node, and correspondingly obtaining the first local momentum, the first group momentum, and the first global momentum specifically includes:

[0057] Update the local momentum according to a preset random subset and a preset local gradient calculated by the local model, and correspondingly obtain the first local momentum; wherein, the update formula of the local momentum is specifically:

[0058]

[0059] where a is the first local momentum, b is the local momentum, γ is a preset update coefficient, is the preset local gradient;

[0060] Update the group momentum and the global momentum respectively according to the group model and global model sent by the corresponding edge node, and correspondingly obtain the first group momentum and the first global momentum.

[0061] In a preferred embodiment, during each round of local iteration, the local device uses a random subset of the local data and the local model to calculate the gradient and update the local momentum. When updating the momentum, the parameter γ can control the influence of the current gradient on the momentum. Generally, γ takes values in the range of [0, 1). The larger γ is, the greater the influence of the historical gradient on the momentum. When γ = 0, the momentum degenerates into the gradient. is the stochastic gradient calculated by the local device through the stochastic gradient descent method.

[0062] The reason for selecting the random subset as the training data is that it can, to a certain extent, alleviate the phenomenon of non-independent and identically distributed data and avoid overfitting. Similarly, the group momentum and the global momentum can also be updated using the group model and the global model sent by the edge node respectively. After completing the momentum update, the local model uploads the updated first local momentum, first group momentum, and the first global momentum to the corresponding edge node, providing a data basis for implementing the reverse model feedback strategy.

[0063] Further, training the local model according to the first local momentum specifically includes:

[0064] Updating the local model through the momentum gradient descent method according to the first local momentum;

[0065] Correcting the updated local model according to the group model sent by the corresponding edge node to obtain the trained local model.

[0066] Preferably, correcting the updated local model according to the group model sent by the corresponding edge node specifically includes:

[0067] Calculating the first squared norm between the updated local model and the group model sent by the corresponding edge node;

[0068] Taking the first squared norm as the correction term to correct the updated local model; where the specific formula for the correction is:

[0069]

[0070] where, v t is the updated local model in the t-th iteration round, u t ′ is the group model sent by the corresponding edge node in the t-th iteration round, η is the preset learning rate, λ is the preset correction coefficient, ||v t -u t || 2 is the first squared norm.

[0071] In a preferred embodiment, after the local device completes the local momentum update, it can use the updated local momentum (i.e., the first local momentum) to replace the gradient and update the local model through the momentum gradient descent method. Since momentum not only contains the information of the current gradient but also, to a certain extent, contains the information of all previous gradients, training the model using momentum can make the convergence direction of model training more stable and the convergence process more efficient.

[0072] Meanwhile, during the training process, the group model passed back by the edge nodes in the previous iteration rounds is introduced as a correction term. By calculating the squared norm of the local model and the group model, the local model is corrected and updated to implement the forward model correction strategy.

[0073] This preferred embodiment uses the parameter λ to control the degree of correction. When λ → 0, the local device tends to personalized learning; when λ → ∞, the local device tends to assimilate with the group model. Among them, the value of λ can be weighed and modified by the local device according to the actual situation.

[0074] Through the above steps, the local device has completed a complete local iteration implementation process. Since the local devices in the same group have inconsistent data update frequencies due to factors such as region and user habits, and the stability between devices is also limited by factors such as hardware conditions, in this preferred embodiment, local devices are allowed to perform local iterations in parallel and asynchronously. That is, in a certain iteration round, when the local devices in a subset within the group have completed the local iteration, it can be regarded as the group has completed a successful local iteration implementation process.

[0075] For step S2, specifically, aggregating the second group momentum and the second global momentum in the edge nodes according to the first group momentum and the first global momentum specifically includes:

[0076] Aggregating the second group momentum through the weighted average method according to the first group momentum; among them, the aggregation weight for aggregating the first group momentum is the ratio of the sample data volume of the data set contained in the local device corresponding to the first group momentum to the sample data volume of the data set contained in the edge node.

[0077] Aggregating the second global momentum through the weighted average method according to the first global momentum; among them, the aggregation weight for aggregating the first global momentum is the ratio of the sample data volume of the data set contained in the local device corresponding to the first global momentum to the sample data volume of the data set contained in the edge node.

[0078] In a preferred embodiment, the edge node is regarded as a special type of node different from the local device. This type of node can only act as the group leader and is only responsible for maintaining the group model through the information fed back by the nodes within the group without having to store any data, which helps with the applicability of the edge node.

[0079] The edge node aggregates the first group of momenta uploaded by the local devices within the group through weighted averaging to obtain the second group of momenta. During the aggregation process, the weight can generally use the ratio of the data volume of the node within the group to the data volume within the group, which means that the node with a larger data volume contributes more to the group model, making the group model perform better on the nodes with a larger data volume and improving the overall generalization ability at the same time.

[0080] At the same time, the edge node also needs to perform an aggregation on the first global momentum to obtain the second global momentum, and upload the second global momentum to the corresponding central server for the central server to train the global model.

[0081] Further, training the group model according to the second group of momenta specifically includes:

[0082] Updating the group model through the momentum gradient descent method according to the second group of momenta;

[0083] Correcting the updated group model according to the global model sent by the corresponding central server to obtain the trained group model.

[0084] Preferably, correcting the updated group model according to the global model sent by the corresponding central server specifically includes:

[0085] Calculating the second squared norm between the updated group model and the global model sent by the corresponding central server;

[0086] Taking the second squared norm as a correction term to correct the updated group model; where the specific formula for the correction is:

[0087] u t+1 =u t -d-λ||u t -w t || 2

[0088] where, u t is the updated group model in the t-th iteration round, w t is the global model sent by the corresponding central server in the t-th iteration round, d is the second group of momenta, λ is a preset correction coefficient, and ||u t -w t || 2 is the second squared norm.

[0089] In a preferred embodiment, after aggregating the second set of momenta, the group model can be trained by the momentum gradient descent method, and the global model sent by the central server is introduced for correction during the training process.

[0090] After training the group model, the edge node also needs to synchronize the trained group model and the global model sent by the central server to all local devices within the group in a broadcast form to further reduce communication overhead. It should be noted that since the periods of intra-group aggregation and global aggregation are inconsistent, the edge node may synchronize the same global model to local devices within multiple intra-group aggregation periods. This preferred embodiment does not force the edge node to synchronize the latest global model to local devices each time, because considering the instability of the edge node, the edge node may not successfully receive the global model in each global aggregation period. Therefore, the edge node only needs to ensure that the global model stored by itself is synchronized to local devices, and this measure can provide a certain degree of fault tolerance for the system.

[0091] Through the above steps, the edge node completes a full round of the intra-group aggregation implementation process. In this preferred embodiment, the edge node only needs to undertake the functions of computing and communication, and does not need to store a large amount of data. This enables when selecting the hardware device of the edge node in the actual scenario, only the computing and communication related metrics need to be considered, without the need to be equipped with a large-capacity storage, reducing the construction cost of the edge node. At the same time, since the edge node does not use any of its own data for computing, the data sent back by the edge node to the central server can better represent the data distribution within the group.

[0092] For step S3, specifically, aggregating the third global momentum in the central server according to the second global momentum, and training the global model according to the third global momentum specifically includes:

[0093] Aggregating the third global momentum by the weighted average method according to the second global momentum; wherein, the aggregation weight for aggregating the second set of momenta is the ratio of the sample data volume of the data set contained in the edge node corresponding to the second set of momenta to the sample data volume of the data set contained in the central server;

[0094] Training the global model by the momentum gradient descent method according to the third global momentum.

[0095] In a preferred embodiment, the central server aggregates the global momentum uploaded by the edge nodes through weighted averaging, using the ratio of the data volume within the edge node group to the total data volume as the weight. Considering the instability of the edge nodes, when the central server receives a certain number of edge node data returns within a single cycle, this step can be executed. After the training is completed, the trained global model is sent to all the edge nodes. Thus, the method of the present invention completes a complete federated learning process. Generally, the method of the present invention continuously repeats the above process until it is determined that the global model obtained from the latest training converges.

[0096] Further, the convergence condition of the global model is specifically:

[0097] Calculate the loss function value of the global model sent to the corresponding edge node according to the preset loss function in the local device;

[0098] Aggregate the loss function values of the global model calculated by all the local devices to obtain a global loss value;

[0099] When the global loss value reaches the minimum, it is determined that the global model converges.

[0100] In a preferred embodiment, after multiple rounds of iteration, the method of the present invention can obtain an optimal global model, such that w * , so that the global loss value obtained by aggregating through the function G(·) can reach the minimum. In the Fedavg (Federated Averaging) framework, the aggregation function where |D t | is the number of dataset samples of the k-th local device, |D| is the number of all dataset samples in the central server, and f(·) is the preset loss function.

[0101] Thus, the optimization objective of federated learning can be modeled as: w * = argminG[f 1 (w),... f k (w)]

[0102] Refer to Figure 2 , which is a schematic structural diagram of a federated learning framework based on a momentum-based hierarchical trade-off model provided by an embodiment of the present invention. As Figure 2 can be seen, on the basis of the existing federated learning algorithm, the present invention introduces the following two technical improvements:

[0103] First is the hierarchical topology. The entire federated learning architecture centers around multiple relay devices such as edge nodes within a certain area. User devices are logically divided into groups according to certain rules (in actual scenarios, factors such as geography, device types, and user habits are often considered for division), forming a three-level topology: local devices – edge nodes – central server. Edge nodes can be regarded as the group leaders of their respective groups. At the same time, it is stipulated that during the entire learning process, local devices only communicate with the edge nodes of their respective groups, and edge nodes synchronize information with nodes within the group through broadcasting, reducing the communication overhead within the group. The central server only communicates with the group leaders of each group, that is, the edge nodes. During the model training process, after local devices complete multiple rounds of local training, they send the updates to the edge nodes within their respective groups, and the edge nodes then weighted aggregate these updates to train the group model maintained by the edge nodes. Similarly, after edge nodes complete several rounds of training, they send the updates to the central server, and the central server also aggregates these updates through weighted averaging and updates the global model accordingly.

[0104] Second is the two-way collaborative optimization strategy. Generally speaking, the two-way collaborative optimization strategy includes a forward model correction strategy from top to bottom, that is, a reverse model feedback strategy from bottom to top. These two strategies work together based on the hierarchical structure during the model training process to dynamically balance the generalization and personalization capabilities of models at all levels. For the forward model correction strategy, during the training of models at all levels, the superior model is introduced from top to bottom as a correction term to correct the objective function, balancing the differences between the local model and the group model, and between the group model and the global model. This enables models at all levels to better fit local data in scenarios with data heterogeneity while also improving the generalization ability of the model to a certain extent and avoiding overfitting.

[0105] Based on the above two improvement strategies, the federated learning architecture provided by the present invention reduces unnecessary communication losses and communication costs through the hierarchical topology. At the same time, the introduction of the two-way collaborative optimization strategy enables models at all levels to improve their generalization ability based on the correction of the superior model and their personalization ability based on momentum aggregation, thereby further improving the overall model adaptation ability of federated learning.

[0106] In addition, in the method of the present invention, parameters such as the learning rate of the local model, the correction coefficient, the threshold for the edge node to perform in-group aggregation when receiving feedback information from a certain number of local devices, and the threshold for the central server to receive the number of edge nodes all need to be flexibly adjusted according to actual needs in the actual scenario.

[0107] Refer to Figure 3, which is a schematic structural diagram of a hierarchical compromise model federated learning device based on momentum provided by another embodiment of the present invention, including: a local module 101, an edge module 102, and a central module 103;

[0108] The local module 101 is used to update the local momentum, group momentum, and global momentum in the local device according to the group model and the global model sent by the corresponding edge node, and correspondingly obtain the first local momentum, the first group momentum, and the first global momentum, and train the local model according to the first local momentum; wherein, one of the edge nodes corresponds to a plurality of the local devices;

[0109] The edge module 102 is used to aggregate the second group momentum and the second global momentum in the edge node according to the first group momentum and the first global momentum, and train the group model according to the second group momentum, and then send the trained group model and the global model sent by the corresponding central server to the plurality of local devices corresponding to the edge node; wherein, one of the central servers corresponds to a plurality of the edge nodes;

[0110] The central module 103 is used to aggregate the third global momentum in the central server according to the second global momentum, train the global model according to the third global momentum, and after sending the trained global model to the plurality of edge nodes corresponding to the central server, start the next round of iteration until it is determined that the global model obtained in the latest round of iteration converges, and then end the iteration.

[0111] Obviously, the above embodiments of the present invention are only examples for clearly illustrating the present invention, and are not intended to limit the implementation manners of the present invention. For those of ordinary skill in the art, other different forms of changes or modifications can be made based on the above description. It is not necessary and impossible to enumerate all the implementation manners here. Any modifications, equivalent replacements, and improvements made within the spirit and principle of the present invention shall be included in the protection scope of the claims of the present invention.

Claims

1. A momentum-based hierarchical compromise model federated learning method, characterized in that: The steps include: According to the group model and the global model sent by the corresponding edge node, the local momentum, the group momentum and the global momentum are updated in the local device, and the first local momentum, the first group momentum and the first global momentum are obtained accordingly, and the local model is trained according to the first local momentum; wherein one edge node corresponds to several local devices; According to the first group of momentum and the first global momentum, a second group of momentum and a second global momentum are aggregated in the edge node, and a group model is trained according to the second group of momentum, and then the trained group model and the global model sent by the corresponding central server are sent to the local devices corresponding to the edge nodes; wherein one central server corresponds to several edge nodes; According to the second global momentum, a third global momentum is aggregated in the central server, a global model is trained according to the third global momentum, and after the trained global model is sent to several edge nodes corresponding to the central server, the next round of iteration is started, and the iteration is terminated after it is determined that the global model obtained in the latest round of iteration converges.

2. The momentum-based hierarchical compromise model federated learning method according to claim 1, characterized in that: The updating of the local momentum, the group momentum and the global momentum in the local device according to the group model and the global model sent by the corresponding edge node, and correspondingly obtaining the first local momentum, the first group momentum and the first global momentum, specifically includes: The local momentum is updated according to the preset random subset and the preset local gradient calculated by the local model, and the first local momentum is obtained accordingly; wherein the update formula of the local momentum is specifically: Wherein, a is the first local momentum, b is the local momentum, γ is a preset update coefficient, Preset local gradient for said method; According to the group model and the global model sent by the corresponding edge node, the group momentum and the global momentum are updated respectively, and the first group momentum and the first global momentum are obtained correspondingly.

3. The momentum-based hierarchical compromise model federated learning method according to claim 1, characterized in that: The training of the local model according to the first local momentum specifically includes: updating the local model by a momentum gradient descent method according to the first local momentum; The updated local model is corrected according to the group model sent by the corresponding edge node to obtain the trained local model.

4. The momentum-based hierarchical compromise model federated learning method according to claim 3, characterized in that: The updating of the local model according to the group model sent by the corresponding edge node specifically includes: Calculating a first square norm between the updated local model and the group model sent by the corresponding edge node; The first square norm is used as a correction term to correct the updated local model; wherein the correction formula is specifically: Among them, v t is the updated local model described in the tth iteration round, u t ′ is the group model sent by the corresponding edge node in the tth iteration round, η is the preset learning rate, λ is the preset correction coefficient, ∥v t -u t ∥ 2 is the first square norm.

5. The momentum-based hierarchical compromise model federated learning method according to claim 1, characterized in that: The step of aggregating the first group of momentum and the first global momentum in the edge node to obtain a second group of momentum and a second global momentum specifically includes: According to the first group of momentums, the second group of momentums are aggregated by weighted average method to obtain; wherein the aggregation weight of the first group of momentums is the ratio of the sample data volume of the data set contained in the local device corresponding to the first group of momentums to the sample data volume of the data set contained in the edge node; According to the first global momentum, the second global momentum is aggregated by weighted average method; wherein the aggregation weight of the first global momentum is the ratio of the sample data volume of the data set contained in the local device corresponding to the first global momentum to the sample data volume of the data set contained in the edge node.

6. The momentum-based hierarchical compromise model federated learning method according to claim 1, characterized in that: The step of training the group model according to the second group of momentum specifically includes: updating the set of models by a momentum gradient descent method according to the second set of momentums; The updated group model is modified according to the global model sent by the corresponding central server to obtain the trained group model.

7. The method according to claim 6, characterized in that The step of modifying the updated group model according to the global model sent by the corresponding central server specifically includes: Calculating a second square norm between the updated group model and the global model sent by the corresponding central server; The updated group model is corrected by using the second square norm as a correction term; wherein the correction formula is specifically: u t+1 =u t -d-λ∥u t -w t ∥ 2 Among them, u t is the updated group model in the tth iteration round, w t is the global model sent by the corresponding central server in the tth iteration round, d is the second group of momentum, λ is the preset correction coefficient, ∥u t -w t ∥ 2 is the second square norm.

8. The method according to claim 1, characterized in that The step of aggregating the second global momentum in the central server to obtain a third global momentum, and training a global model according to the third global momentum specifically includes: According to the second global momentum, the third global momentum is obtained by aggregating by weighted average method; wherein the aggregation weight of the second group of momentums is the ratio of the sample data volume of the data set contained in the edge node corresponding to the second group of momentum to the sample data volume of the data set contained in the central server; The global model is trained according to the third global momentum by a momentum gradient descent method.

9. The method according to claim 1, characterized in that The convergence condition of the global model is specifically: Calculate the loss function value of the global model sent by the corresponding edge node according to the preset loss function in the local device; Aggregating the global model loss function values ​​calculated by all the local devices to obtain a global loss value; When the global loss value reaches a minimum, it is determined that the global model converges.

10. A momentum-based hierarchical compromise model federated learning device, characterized in that: include: local modules, edge modules, and central modules; The local module is used to update the local momentum, the group momentum and the global momentum in the local device according to the group model and the global model sent by the corresponding edge node, and obtain the first local momentum, the first group momentum and the first global momentum accordingly, and train the local model according to the first local momentum; wherein one edge node corresponds to a plurality of local devices; The edge module is used to aggregate the second group of momentum and the first global momentum in the edge node to obtain a second group of momentum and a second global momentum, and train a group model according to the second group of momentum, and then send the trained group model and the global model sent by the corresponding central server to the local devices corresponding to the edge node; wherein one central server corresponds to several edge nodes; The central module is used to aggregate the third global momentum in the central server according to the second global momentum, train the global model according to the third global momentum, and start the next round of iteration after sending the trained global model to several edge nodes corresponding to the central server, and end the iteration after determining that the global model obtained in the latest round of iteration converges.

Citation Information

Patent Citations

  • Network topology construction method and system in hierarchical federated learning scene

    CN114650227A

  • Lateral federated learning fault detection method based on multilayer grouping aggregation

    CN116820816A

  • Group personalized federal learning method

    CN117313834A