Federal learning method based on adaptive optimizer

By adopting the design of adaptive optimizer and local momentum buffer in federated learning, the problems of large communication overhead, slow convergence speed and low accuracy in federated learning are solved, and faster convergence speed and higher model training accuracy are achieved.

CN120218273APending Publication Date: 2025-06-27SHANGHAI JIAOTONG UNIV
View PDF 0 Cites 0 Cited by

Patent Information

Application Number
CN202311828532.X
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2023-12-27
Publication Date
2025-06-27

AI Technical Summary

Technical Problem

Federated learning faces the problems of large communication overhead, slow convergence speed and low model training accuracy caused by data heterogeneity.

Method used

Adaptive optimizer-based federated learning method, using the Adamax optimizer to accelerate model updates by retaining local momentum buffers on the client, computing and weighting model parameters, first and second-order moments, and averaging local buffers and model parameters in each round of communication.

Benefits of technology

The convergence speed of federated learning is improved, the convergence ability of the client model is enhanced, and the problem of low model training accuracy caused by data heterogeneity is partially solved, and the defect of insufficient federated learning communication ability is compensated.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN120218273A_ABST
    Figure CN120218273A_ABST
Patent Text Reader

Abstract

The invention relates to a federal learning method based on an adaptive optimizer, and the method comprises the steps: S1, selecting one or more clients, and transmitting an initial model parameter of a current communication round t of each local model, an initial first moment of a model gradient, and an initial second moment of the model gradient; s2, calculating a weighted sum as a current model parameter, a current first-order moment and a current second-order moment; s3, calculating a model parameter, a first moment and a second moment of a communication round t + 1; s4, updating model parameters, a first moment and a second moment of the global model, storing the model parameters, the first moment and the second moment in a buffer area, and updating the model parameters based on an Adamax method; and S5, taking the communication round t + 1 as a new current communication round, and repeating the steps S1 to S4 until the upper limit of the communication round is reached. Compared with the prior art, the method has the advantages of improving the convergence speed of federal learning and the like.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention relates to the technical field of federated learning, and in particular, to a federated learning method based on an adaptive optimizer. Background Art

[0002] Federated learning has many application fields, but it faces some new challenges. These challenges can be roughly divided into two categories: challenges related to model training and challenges related to system security. Challenges related to model training include communication overhead during multiple training iterations, heterogeneity of participating devices, and heterogeneity of the data used. Challenges related to system security include the risk of user privacy leakage, poisoning attacks on federated learning data and models, and backdoor attacks on federated learning.

[0003] In traditional machine learning modeling, the data required for model training is usually collected in a data center and then the model is trained. In federated learning, it can be regarded as sample-based distributed model training, where all data is distributed to different machines. Each machine downloads the model from the server, then uses local data to train the model, and then returns the parameters that need to be updated to the server, and then repeats the above process. In this process, the models on each machine are the same and complete, and the machines do not communicate or depend on each other. Each machine can also make independent predictions during prediction.

[0004] Federated learning algorithms are particularly sensitive to communication costs. And in the process of distributed data collection, major challenges such as energy efficiency and system latency problems will occur. The most common method for federated learning optimization is FedAvg. The basic idea of FedAvg is to upload the model parameters trained by local clients to the server. The server calculates the average value of all collected client model parameters, and then broadcasts this average value back to all local devices. This process can be iterated multiple times until convergence. Although FedAvg partially solves the challenge of limited communication capabilities of clients, there is still room for improvement in the convergence speed, and the heterogeneity of each client's data, that is, the data is non-independent and identically distributed (Non-IID), will affect the accuracy of federated learning model training. Summary of the Invention

[0005] The purpose of the present invention is to provide a federated learning method based on an adaptive optimizer to improve the convergence speed of federated learning and make up for the insufficient communication ability of federated learning.

[0006] The purpose of the present invention can be achieved by the following technical solutions:

[0007] A federated learning method based on an adaptive optimizer, characterized in that the method includes:

[0008] S1. Select one or more clients, and the server sends the initial model parameters, the initial first moment of the model gradient, and the initial second moment of the model gradient of each local model at the current communication round t to the selected clients;

[0009] S2. The client calculates the weighted sum of the initial model parameters of each model and the model parameters of the global model stored in the buffer at the previous communication round, the weighted sum of the initial first moment and the first moment of the global model stored in the buffer at the previous communication round, and the weighted sum of the initial second moment and the second moment of the global model stored in the buffer at the previous communication round, and takes the three weighted sums as the current model parameters, the current first moment, and the current second moment respectively;

[0010] S3. Set the initial number of iterations to 0, and update the model by one step for the current model parameters, the current first moment, and the current second moment of each client based on the Adamax method, calculate the model parameters, the first moment, and the second moment of the next iteration of each model, and repeat S3 until the iteration upper limit M is reached. At this time, the model parameters, the first moment, and the second moment obtained in the last iteration represent the model parameters, the first moment, and the second moment of communication round t+1;

[0011] S4. Each client transmits back the model parameters, the first moment, and the second moment of communication round t+1. Based on the transmitted parameters, update the model parameters, the first moment, and the second moment of the global model to obtain the model parameters, the first moment, and the second moment of the global model of communication round t+1, and store the model parameters, the first moment, and the second moment of the global model of communication round t+1 in the buffer, where the model parameters are updated based on the Adamax method;

[0012] S5. Take communication round t+1 as the new current communication round, and repeat S1~S4 until the communication round upper limit is reached.

[0013] Further, the specific calculation of the weighted sum of the initial model parameters of each model and the model parameters of the global model stored in the buffer at the previous communication round is as follows:

[0014] Multiply the initial model parameters of each model by the first weight parameter, multiply the model parameters of the global model stored in the buffer at the previous communication round by the second weight parameter, and add the two products to obtain the weighted sum of the initial model parameters of each model and the model parameters of the global model stored in the buffer at the previous communication round.

[0015] Further, the specific calculation of the weighted sum of the initial first moment of each model and the first moment of the global model stored in the buffer at the previous communication round is as follows:

[0016] Multiply the initial first moment of each model by the third weight parameter, multiply the first moment of the global model in the previous communication round stored in the buffer by the fourth weight parameter, add the two products to obtain the weighted sum of the initial first moment of each model and the first moment of the global model in the previous communication round stored in the buffer.

[0017] Furthermore, the specific steps for calculating the weighted sum of the initial second moment of each model and the second moment of the global model in the previous communication round stored in the buffer are as follows:

[0018] Multiply the initial second moment of each model by the fifth weight parameter, multiply the second moment of the global model in the previous communication round stored in the buffer by the sixth weight parameter, add the two products to obtain the weighted sum of the initial second moment of each model and the second moment of the global model in the previous communication round stored in the buffer.

[0019] Furthermore, the weight parameter is a constant between 0 and 1.

[0020] Furthermore, the specific steps for updating the first moment of the global model in S4 are as follows:

[0021] Sum the first moments of all local models of all clients at communication round t + 1, divide the obtained sum by the number of selected clients, and the resulting quotient is the first moment of the global model at communication round t + 1.

[0022] Furthermore, the specific steps for updating the second moment of the global model in S4 are as follows:

[0023] Sum the second moments of all local models of all clients at communication round t + 1, divide the obtained sum by the number of selected clients, and the resulting quotient is the second moment of the global model at communication round t + 1.

[0024] Furthermore, when updating the current model parameters, current first moment, and current second moment of each client by one step based on the Adamax method, the gradients of the local models of each client are used.

[0025] Furthermore, when updating the model parameters of the global model based on the Adamax method, the gradients of the server side are used.

[0026] Furthermore, after updating the model parameters of the global model based on the Adamax method, the model parameters of the global model at communication round t + 1 are:

[0027] w t+1 = w t -Δθ

[0028] where Δθ represents the change between the model parameters of the global model in two communication rounds, and w t+1The model representing the global model of communication round t+1.

[0029] Compared with the prior art, the present invention has the following beneficial effects:

[0030] In the present invention, each client retains a local momentum buffer to record the direction in which each client model should converge, calculates the weighted sum of the initial model parameters of each model and the model parameters of the global model of the previous communication round stored in the buffer, the weighted sum of the initial first moment and the first moment of the global model of the previous communication round stored in the buffer, and the weighted sum of the initial second moment and the second moment of the global model of the previous communication round stored in the buffer, and averages the local buffer and local model parameters in each round of communication. With such a design in each round of communication, the client model has not only the model parameters sent by the server, but also the model gradient directions trained by the previous local clients, which can thus accelerate the convergence of the local model and to a certain extent solve the impact on the accuracy of the federated learning model training caused by the heterogeneity of data among clients (i.e., non-independent and identically distributed data, Non-IID). BRIEF DESCRIPTION OF THE DRAWINGS

[0031] Figure 1 is a flowchart of the present invention;

[0032] Fig. 2(a) is a flowchart of the update of the model gradients of the standard federated learning client and server;

[0033] Fig. 2(b) is a flowchart of the update of the momentum when aggregating the federated learning server model of the present invention;

[0034] Fig. 2(c) is a flowchart of the update of the momentum when the federated learning client of the present invention updates the model. DETAILED DESCRIPTION OF THE INVENTION

[0035] The present invention will be described in detail below with reference to the accompanying drawings and specific embodiments. This embodiment is implemented on the premise of the technical solution of the present invention, and gives the detailed implementation manner and specific operation process, but the protection scope of the present invention is not limited to the following embodiments.

[0036] The present invention proposes a federated learning method based on an adaptive optimizer, and the flowchart of the method is as Figure 1As shown. From the perspective of the server side, each aggregation of the server-side model is regarded as a convergence process. Therefore, the present invention designs an adaptive optimizer AdaMax for this convergence to accelerate the convergence of the global loss function. Such a design enables the update of the server-side model parameters to utilize both the aggregation of local client parameters and the gradient information of the previous global model parameter updates, thereby accelerating the convergence of the federated learning training process and making up for the insufficient communication ability of federated learning. From the perspective of the client side, the present invention introduces a client design, that is, each client retains a local momentum buffer to record the direction in which each client model should converge, and averages the local buffer and local model parameters in each round of communication. With such a design in each round of communication, the client model has not only the model parameters sent by the server side, but also the gradient directions of the models trained by the previous local clients, which can also accelerate the convergence of the local model and, to a certain extent, solve the impact of data heterogeneity among clients (i.e., non-independent and identically distributed data, Non-IID) on the accuracy of federated learning model training, and also achieve the acceleration of the convergence of the federated learning training process.

[0037] Figure 2(a) is a flow chart of the update of the standard federated learning client and server model gradients. In the server side of the present invention, as shown in Figure 2(b), when the system aggregates the federated learning server model, if the global average is updated:

[0038]

[0039] Simple combination and rewriting:

[0040]

[0041] Then the second term of the above formula:

[0042]

[0043] is actually the gradient, where N is the number of aggregated clients, w t represents the parameters of the global model in the t-th round of communication, represents the parameters of the local model uploaded by the i-th client in the t-th round of communication. Therefore, the server-side aggregation is similar to the stochastic gradient descent scenario and can be accelerated using an adaptive optimizer on the server side. Adamax is an effective stochastic optimization method that calculates the adaptive learning rate for different parameters through the first-order moment estimation and second-order moment estimation of the gradient. The Adamax algorithm is a variant of the Adam algorithm based on the infinity norm, making the learning rate update algorithm more stable and simple. The calculation formula for its parameter update is as follows:

[0044]

[0045] r t= max(v * r t-1 , |g t |)

[0046]

[0047] In the client, as shown in Figure 2(c), similar to the momentum on the server side, in the local model local update step, each client maintains a local momentum buffer, which represents the gradient direction of the model during the convergence process. And after each round of communication with the server side, the average value of the local buffer and the local model parameters is calculated, which is equivalent to sharing the convergence direction of the model in addition to sharing the model parameters during the model aggregation stage.

[0048] A federated learning method based on an adaptive optimizer of the present invention includes:

[0049] S1. Select one or more clients, and the server sends the initial model parameters, the initial first moment of the model gradient, and the initial second moment of the model gradient of the current communication round t of each local model to the selected clients;

[0050] S2. The client calculates the weighted sum of the initial model parameters of each model and the model parameters of the global model of the previous communication round stored in the buffer, the weighted sum of the initial first moment and the first moment of the global model of the previous communication round stored in the buffer, and the weighted sum of the initial second moment and the second moment of the global model of the previous communication round stored in the buffer, and takes the three weighted sums as the current model parameters, the current first moment, and the current second moment respectively;

[0051] S3. Set the initial number of iterations to 0, and update the current model parameters, the current first moment, and the current second moment of each client by one step based on the Adamax method, calculate the model parameters, the first moment, and the second moment of the next iteration of each model, and repeat S3 until the iteration upper limit M is reached. At this time, the model parameters, the first moment, and the second moment obtained in the last iteration represent the model parameters, the first moment, and the second moment of the communication round t + 1;

[0052] S4. Each client transmits back the model parameters, the first moment, and the second moment of the communication round t + 1, and updates the model parameters, the first moment, and the second moment of the global model based on the transmitted parameters to obtain the model parameters, the first moment, and the second moment of the global model of the communication round t + 1, and stores the model parameters, the first moment, and the second moment of the global model of the communication round t + 1 in the buffer, where the model parameters are updated based on the Adamax method;

[0053] S5. Take the communication round t + 1 as the new current communication round, and repeat S1 to S4 until the communication round upper limit is reached.

[0054] In S1, randomly select N clients from all clients and denote them as SetN, where N is much smaller than the total number of clients K.

[0055] In S2 and S3, for each selected client i ∈ SetN, perform the following steps in parallel:

[0056] a) The server sends the initial model parameters of the current communication round t The first moment of the model gradient The second moment of the model gradient to each client, and the client calculates the weighted sum of the buffer parameters and the received parameters from the server where M represents the number of local model iterations of the client, and α and β represent constants between 0 and 1.

[0057] b) Each client performs one step of model update and updates the first moment of the model gradient Updates the second moment of the model gradient where g t is determined by the model trained by the local client, and the Adamax equation is used to update the model parameters

[0058] c) Repeat step b) for M times of client iteration.

[0059] d) Each client sends back the updated model parameters The updated first moment of the model gradient The updated second moment of the model gradient to the server.

[0060] In S4, each client sends back the updated model parameters The updated first moment of the model gradient The updated second moment of the model gradient to the server. Then update the first moment and second moment of the global model gradient as:

[0061]

[0062]

[0063] Calculate the accumulated momentum at the server side, where g t is the gradient at the server side, that is Use the Adamax equation to update the model parameter w t+1 = w t -Δθ.

[0064] In S5, repeat S1 - S4 for T times.

[0065] The preferred specific embodiments of the present invention have been described in detail above. It should be understood that those of ordinary skill in the art can make many modifications and variations based on the concept of the present invention without creative efforts. Therefore, all technical solutions that can be obtained by those skilled in the art in the technical field based on the concept of the present invention through logical analysis, reasoning, or limited experiments on the basis of the prior art should fall within the protection scope determined by the claims.

Claims

1. A federated learning method based on an adaptive optimizer, characterized in that, The method includes: S1. Select one or more clients, and the server sends the initial model parameters, the initial first moment of the model gradient, and the initial second moment of the model gradient of the current communication round t of each local model to the selected clients; S2. The client calculates the weighted sum of the initial model parameters of each model and the model parameters of the global model of the previous communication round stored in the buffer, the weighted sum of the initial first moment and the first moment of the global model of the previous communication round stored in the buffer, and the weighted sum of the initial second moment and the second moment of the global model of the previous communication round stored in the buffer, and takes the three weighted sums as the current model parameters, the current first moment, and the current second moment respectively; S3. Set the initial number of iterations to 0, perform one-step model update on the current model parameters, the current first moment, and the current second moment of each client based on the Adamax method, calculate the model parameters, the first moment, and the second moment of the next iteration of each model, and repeat S3 until the iteration upper limit M is reached. At this time, the model parameters, the first moment, and the second moment obtained in the last iteration represent the model parameters, the first moment, and the second moment of communication round t + 1; S4. Each client transmits back the model parameters, the first moment, and the second moment of communication round t + 1. Based on the transmitted parameters, update the model parameters, the first moment, and the second moment of the global model to obtain the model parameters, the first moment, and the second moment of the global model of communication round t + 1, and store the model parameters, the first moment, and the second moment of the global model of communication round t + 1 in the buffer, where the model parameters are updated based on the Adamax method; S5. Take communication round t + 1 as the new current communication round, and repeat S1 - S4 until the communication round upper limit is reached.

2. The federated learning method based on an adaptive optimizer according to claim 1, wherein Specifically, calculating the weighted sum of the initial model parameters of each model and the model parameters of the global model of the previous communication round stored in the buffer is as follows: Multiply the initial model parameters of each model by the first weight parameter, multiply the model parameters of the global model of the previous communication round stored in the buffer by the second weight parameter, and add the two products to obtain the weighted sum of the initial model parameters of each model and the model parameters of the global model of the previous communication round stored in the buffer.

3. The federated learning method based on an adaptive optimizer according to claim 2, wherein, Specifically, calculating the weighted sum of the initial first moment of each model and the first moment of the global model of the previous communication round stored in the buffer is as follows: Multiply the initial first moment of each model by the third weight parameter, multiply the first moment of the global model of the previous communication round stored in the buffer by the fourth weight parameter, and add the two products to obtain the weighted sum of the initial first moment of each model and the first moment of the global model of the previous communication round stored in the buffer.

4. A federated learning method based on an adaptive optimizer according to claim 2, characterized in that Specifically, calculating the weighted sum of the initial second moment of each model and the second moment of the global model of the previous communication round stored in the buffer is as follows: Multiply the initial second moment of each model by the fifth weight parameter, multiply the second moment of the global model of the previous communication round stored in the buffer by the sixth weight parameter, and add the two products to obtain the weighted sum of the initial second moment of each model and the second moment of the global model of the previous communication round stored in the buffer.

5. A federated learning method based on an adaptive optimizer according to claim 2 or 3 or 4, characterized in that The weight parameter is a constant between 0 and 1.

6. The federated learning method based on an adaptive optimizer according to claim 1, wherein The specific steps for updating the first moment of the global model in S4 are as follows: Sum the first moments of all local models of all clients at communication round t + 1, divide the obtained sum by the number of selected clients, and the quotient obtained is the first moment of the global model at communication round t + 1.

7. A federated learning method based on an adaptive optimizer according to claim 1, characterized in that, The specific steps for updating the second moment of the global model in S4 are as follows: Sum the second moments of all local models of all clients at communication round t + 1, divide the obtained sum by the number of selected clients, and the quotient obtained is the second moment of the global model at communication round t + 1.

8. A federated learning method based on an adaptive optimizer according to claim 1, characterized in that When updating the current model parameters, current first moment, and current second moment of each client by one step of the model based on the Adamax method, the gradients of the local models of each client are used.

9. A federated learning method based on an adaptive optimizer according to claim 1, wherein When updating the model parameters of the global model based on the Adamax method, the gradients of the server side are used.

10. A federated learning method based on an adaptive optimizer according to claim 9, characterized in that After updating the model parameters of the global model based on the Adamax method, the model parameters of the global model at communication round t + 1 are: w t+1 = w t -Δθ Among them, Δθ represents the change between the model parameters of the global models in two communication rounds, and w t+1 represents the model of the global model in communication round t + 1.