A multi-cell federated learning model training method

By introducing iterative synchronization and cumulative update mechanisms in the training of federated learning models in multi-cells, the impact of different base station models update speed differences and wireless network bandwidth resources on model parameter transmission performance is solved, efficient model synchronization and update are achieved, and the performance of federated learning model training is improved.

CN113850397BActive Publication Date: 2025-05-06CHINA SOUTHERN POWER GRID INTERNET SERVICE CO LTD
View PDF 2 Cites 0 Cited by

Patent Information

Application Number
CN202111191808.9
Authority / Receiving Office
CN · China
Patent Type
Patents(China)
Current Assignee / Owner
Filing Date
2021-10-13
Publication Date
2025-05-06
Estimated Expiration
2041-10-13

AI Technical Summary

Technical Problem

The existing federated learning algorithms rarely consider model sharing between multiple cellular base stations, and fail to effectively handle the differences in model update speeds of different base stations, the difference values ​​of multiple iterative model updates, and the impact of wireless network bandwidth resources on model parameter transmission performance, resulting in limited algorithm performance.

Method used

A multi-cell federated learning model training method is proposed. By training the model locally on each cell user and sending the model parameters to the associated base station for global aggregation, each base station records the number of iterations and broadcasts periodically, judges the difference in the number of iterations and the difference in model updates, and decides whether to send model parameters to the neighboring base station based on the available bandwidth threshold value to realize the synchronization and update of the model.

Benefits of technology

By introducing an iterative synchronization mechanism and a cumulative update mechanism, the efficiency of model synchronization between multiple cells is improved, and the utilization rate of wireless link resources is effectively improved, and the performance of federated learning model training is improved.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN113850397B_ABST
    Figure CN113850397B_ABST
Patent Text Reader

Abstract

The present invention relates to a multi-cell federated learning model training method, which belongs to the field of machine learning. The specific steps of the method are: a user trains a local model based on local data and sends it to the associated cellular base station; the base station receives the user's local model and aggregates it into a global model, and records the number of iterations; if the difference in the number of iterations of the base station exceeds a threshold value, the iteration is stopped, otherwise, it is determined whether the model update difference is lower than the threshold value; if so, the sending of model parameters is suspended, otherwise, it is determined whether the current available bandwidth is lower than the threshold value; if so, the base station only sends the corresponding weights whose model weight difference values ​​are higher than the threshold value; otherwise, the base station sends all model parameters; each base station updates the model until the model converges.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention belongs to the field of machine learning and relates to a multi-cell federated learning model training method. Background Art

[0002] The rapid development of communication technology and Internet of Things technology has generated a huge amount of data, including sensitive and security information of users. However, due to the instability of wireless networks and highly limited resources, important information may be leaked or lost during transmission. Federated learning technology, as a typical distributed learning technology, allows devices to use data collected by individuals to train learning models locally. Devices only need to interact with model parameters, thus avoiding the transmission of large amounts of data through wireless links. While protecting user privacy and ensuring information security, it can enable multiple users to share the same model.

[0003] Existing federated learning algorithms mostly consider model training within a cell, and rarely involve model sharing between multiple cellular base stations; in addition, existing research rarely considers the differences in model update speeds of different base stations, the difference values ​​of multiple iterative model updates, and the impact of wireless network bandwidth resources on model parameter transmission performance, resulting in severe limitations on algorithm performance. Summary of the invention

[0004] In view of this, an object of the present invention is to provide a multi-cell federated learning model training method.

[0005] In order to achieve the above object, the present invention provides the following technical solutions:

[0006] A multi-cell federated learning model training method comprises the following steps:

[0007] S1: Each cell user trains a user local model based on local data and sends the local model parameters to the associated base station;

[0008] S2: Each base station receives the user local model parameters sent by the associated user, globally aggregates the model, and generates a base station global model;

[0009] S3: Each base station records the number of iterations of the base station global model and periodically broadcasts it to neighboring base stations;

[0010] S4: Each base station calculates the difference between the number of iterations of its global model and the number of iterations of the neighboring base station model, and determines whether the difference value is higher than a preset threshold value of the number of iterations difference. If so, execute S5, otherwise, execute S6;

[0011] S5: The base station suspends model iteration and returns to S4;

[0012] S6: each base station calculates the model update difference;

[0013] S7: Determine whether the model update difference value is lower than the model update difference threshold value, if so, execute S8, otherwise, execute S9;

[0014] S8: The base station continues to iterate and update the model, calculates the model update difference value, and returns to S7;

[0015] S9: The base station detects whether the available bandwidth is lower than the available bandwidth threshold value. If so, S10 is executed; otherwise, S11 is executed.

[0016] S10: The base station determines the update difference value of each weight in the model and compares it with the threshold value. If the update difference value of the weight is higher than the threshold value, the corresponding weight is sent to the neighboring base station and the process goes to S12;

[0017] S11: The base station sends the global model parameters at the current moment to the neighboring base stations and resets the model update difference;

[0018] S12: Each base station updates the local model according to the received global model parameters and determines whether the model has converged. If so, the model training is completed; otherwise, the process returns to S1.

[0019] Further, in step S1, each cell user trains a user local model based on local data and sends the local model parameters to the associated base station. Specifically, N cells are included, each cell consists of a cellular base station and multiple users, and the number of users in cell n is U n , 1≤n≤N. Each user collects data from the external environment as input for model training. The input of user i in cell n is expressed as where x n,i,m is the mth sample collected by user i in cell n, 1≤m≤M i , M i is the number of samples collected by user i in cell n. The output of user i in cell n can be expressed as User based on X n,i and Y n,i Train the user local model, let represents the local model parameter set corresponding to the tth iteration of user i in cell n, where is the kth parameter of the local model corresponding to the tth iteration of user i in cell n, 1≤k≤K, K is the number of model parameters, and user i in cell n transmits Send to base station n.

[0020] Further, in step S2, the base station receives the user local model parameters sent by the associated user, performs global aggregation on the model, and generates a base station global model, specifically: base station n receives the local model parameters sent by the associated user Perform global model parameter aggregation based on weighted average, let represents the global model parameters determined by base station n in the tth iteration, where The kth parameter of the global model corresponding to the tth iteration of base station n, 1≤k≤K, is modeled for Among them, g n,i is the weight of user i in base station n, and the base station further Sent to the user it is associated with for subsequent iterations of the user.

[0021] Further, in step S3, each base station records the number of iterations of the base station global model and periodically broadcasts it to neighboring base stations. Specifically, let t n represents the number of iterations of the aggregation model at base station n. Base station n will n The value is periodically sent to its neighboring base stations.

[0022] Further, in step S4, each base station calculates the difference between the number of iterations of its global model and the number of iterations of the neighboring base station model, and determines whether the difference value is higher than a preset threshold value of the number of iterations difference, specifically: let Ψ n represents the set of neighboring base stations of base station n. The number of iterations of the aggregation model received by base station n from each neighboring base station is represents the neighboring base station n of base station n 1 ∈Ψ n The number of iterations of the current aggregation model. Δt represents the maximum difference in the number of iterations between base station n and its neighboring base stations, that is: If Δt n >Δt th , that is, if the current iteration number difference of base station n is higher than the iteration number difference value threshold, base station n suspends model iteration until the difference value is lower than the threshold value.

[0023] Further, in step S6, each base station calculates the model update difference, specifically: each base station calculates the difference between the global model corresponding to the current iteration and the previous iteration, and sets Δω n,t It represents the difference between the model of the t-th iteration of base station n and the model of the previous iteration, which can be modeled as

[0024] Further, in step S7, it is determined whether the model update difference value is lower than the model update difference threshold value, specifically: let Δω th To update the difference value threshold of the model, base station n compares Δω n,t With Δω th , if Δω n,t <Δω th , the base station continues to iterate and update the model, otherwise, it detects whether the available bandwidth is lower than the available bandwidth threshold.

[0025] Furthermore, in step S8, the base station continues to perform iterative model updates until the update difference value exceeds the model update difference threshold value.

[0026] Further, in step S9, the base station detects whether the available bandwidth is lower than the available bandwidth threshold value, specifically: let B n,t is the available bandwidth of base station n at the tth iteration, is the available bandwidth threshold of base station n, base station n compares B n,t and like The base station then determines the update difference value of each weight in the model, otherwise, the base station sends the global model parameters at the current moment to the neighboring base stations.

[0027] Further, in step S10, the base station determines the update difference value of each weight in the model and compares it with the threshold value, specifically: let Δω n,t,k represents the updated difference value of the kth weight in the model of base station n, which is modeled as Let Δω′ th Represents the model weight update difference threshold, if Δω n,t,k ≥Δω′ th , then base station n will Send to neighboring base stations, otherwise, base station n does not send

[0028] Further, in step S11, the base station sends the global model parameters at the current moment to the neighboring base stations, specifically: base station n aggregates all the parameters of the model, namely Send to neighboring base stations, 1≤n≤N.

[0029] In step S12, each base station updates the local model according to the received global model parameters and determines whether the model converges. Specifically, base station n receives the global model parameters from its neighboring base stations. n 1 ∈Ψ n , the model parameters of base station n are updated as: Among them, h n is the weight of base station n; the global model convergence criterion is n 1 ∈Ψ n , where ε is the convergence threshold of the global model of each cell.

[0030] The beneficial effects of the present invention are as follows: the present invention is based on a distributed learning framework, realizes federated learning model training between multiple cells, and by introducing an iterative synchronization mechanism, can realize model synchronization between multiple cells; by introducing a cumulative update mechanism and a partial parameter update mechanism, can effectively improve the utilization rate of wireless link resources.

[0031] Other advantages, objectives and features of the present invention will be described in the following description to some extent, and to some extent, will be obvious to those skilled in the art based on the following examination and study, or can be taught from the practice of the present invention. The objectives and other advantages of the present invention can be realized and obtained through the following description. BRIEF DESCRIPTION OF THE DRAWINGS

[0032] In order to make the purpose, technical solutions and advantages of the present invention more clear, the present invention will be described in detail below in conjunction with the accompanying drawings, wherein:

[0033] Figure 1 This is a framework diagram of the multi-cell federated learning model training system;

[0034] Figure 2 This is a flowchart for training a multi-cell federated learning model. DETAILED DESCRIPTION

[0035] The following describes the embodiments of the present invention by specific examples, and those skilled in the art can easily understand other advantages and effects of the present invention from the contents disclosed in this specification. The present invention can also be implemented or applied through other different specific embodiments, and the details in this specification can also be modified or changed in various ways based on different viewpoints and applications without departing from the spirit of the present invention. It should be noted that the illustrations provided in the following embodiments only illustrate the basic concept of the present invention in a schematic manner, and the following embodiments and features in the embodiments can be combined with each other without conflict.

[0036] Among them, the drawings are only used for illustrative explanations, and they only represent schematic diagrams rather than actual pictures, and should not be understood as limitations on the present invention. In order to better illustrate the embodiments of the present invention, some parts of the drawings may be omitted, enlarged or reduced, and do not represent the size of actual products. For those skilled in the art, it is understandable that some well-known structures and their descriptions in the drawings may be omitted.

[0037] The same or similar numbers in the drawings of the embodiments of the present invention correspond to the same or similar parts; in the description of the present invention, it should be understood that if the terms "upper", "lower", "left", "right", "front", "rear", etc. indicate the orientation or position relationship, they are based on the orientation or position relationship shown in the drawings, which is only for the convenience of describing the present invention and simplifying the description, rather than indicating or implying that the device or element referred to must have a specific orientation, be constructed and operate in a specific orientation. Therefore, the terms describing the position relationship in the drawings are only used for illustrative purposes and cannot be understood as limiting the present invention. For ordinary technicians in this field, the specific meanings of the above terms can be understood according to specific circumstances.

[0038] Figure 1 For multi-cell federated learning model training, this architecture includes multiple cells, each of which consists of a cellular base station and multiple users, where:

[0039] Cell: It is composed of a cellular base station and multiple users, and the users in each cell are associated with the cellular base station. The transmission between users and base stations is based on wireless links.

[0040] Cellular base station: collects the local models of the associated users and performs global aggregation based on weighted average, and then returns the aggregated base station global model to the user to update the user's local model; the base station also sends the global model to neighboring base stations to update the global model of neighboring base stations.

[0041] Cell users: mainly used to collect local data, train user local models based on local data, and send the local models to the base station through wireless links to achieve global aggregation of base stations.

[0042] Dataset: A collection of local data collected for users. Users train local models based on data samples in the dataset.

[0043] Figure 2 This is a multi-cell federated learning model training flow chart in the method of the present invention, which specifically includes the following steps:

[0044] Step 1: The user trains the local model;

[0045] Assume there are N cells, each of which consists of a cellular base station and multiple users. Let the number of users in cell n be U n , 1≤n≤N. Each user collects data from the external environment as input for model training. The input of user i in cell n is expressed as where x n,i,m is the mth sample collected by user i in cell n, 1≤m≤M i , M i is the number of samples collected by user i in cell n, and the output of user i in cell n is expressed as User based on X n,i and Y n,i Train the user local model, let represents the local model parameter set corresponding to the tth iteration of user i in cell n, where is the kth parameter of the local model corresponding to the tth iteration of user i in cell n, 1≤k≤K, K is the number of model parameters, and user i in cell n transmits Send to base station n;

[0046] Step 2: The base station performs global aggregation to generate a base station global model; base station n receives local model parameters sent by its associated users Perform global model parameter aggregation based on weighted average, let represents the global model parameters determined by base station n in the tth iteration, where The kth parameter of the global model corresponding to the tth iteration of base station n, 1≤k≤K, is modeled for Among them, g n,i is the weight of user i in base station n, and the base station further Sent to the user it is associated with for subsequent iterations of the user;

[0047] Step 3: Each base station records the number of model iterations and periodically broadcasts it to neighboring base stations. Let t n represents the number of iterations of the aggregation model at base station n. Base station n will n The value is periodically sent to its neighboring base stations;

[0048] Step 4: Each base station determines whether the model iteration number difference value is higher than the preset iteration number difference threshold value; let Ψ n represents the set of neighboring base stations of base station n. The number of iterations of the aggregation model received by base station n from each neighboring base station is represents the neighboring base station n of base station n 1 ∈Ψ n The number of iterations of the current aggregation model. Δt represents the maximum difference in the number of iterations between base station n and its neighboring base stations, that is:

[0049] Step 5: If Δt n >Δt th , that is, if the current iteration number difference of base station n is higher than the iteration number difference value threshold, base station n will suspend model iteration until the difference value is lower than the threshold value;

[0050] Step 6: If the current iteration number difference is lower than the iteration number difference value threshold, each base station determines whether the model update difference value is higher than the preset model update difference threshold value; each base station calculates the difference between the global model corresponding to the current iteration and the previous iteration, and sets Δω n,t It represents the difference between the model of the t-th iteration of base station n and the model of the previous iteration, which can be modeled as Let Δω th To update the difference value threshold of the model, base station n compares Δω n,t With Δω th . ;

[0051] Step 7: If Δω n,t <Δω th, that is, the model update difference value is lower than the model update difference threshold value, the base station does not send the global model, and accumulates the model update difference value until the difference value is higher than the threshold value.

[0052] Step 8: If the model update difference value is higher than the model update difference threshold value, the base station determines whether the current available bandwidth is higher than the preset available bandwidth threshold value; let B n,t is the available bandwidth of base station n at the tth iteration, is the available bandwidth threshold of base station n, base station n compares B n,t and like If the current available bandwidth is lower than the available bandwidth threshold, then the current available bandwidth is higher than the available bandwidth threshold; otherwise, the current available bandwidth is higher than the available bandwidth threshold;

[0053] Step 9: If the current available bandwidth is lower than the available bandwidth threshold, let Δω n,t,k represents the updated difference value of the kth weight in the model of base station n, which is modeled as Let Δω′ th Represents the model weight update difference threshold, if Δω n,t,k ≥Δω′ th , then base station n will Send to neighboring base stations, otherwise, base station n does not send

[0054] Step 10: If the current available bandwidth is higher than the available bandwidth threshold, the base station sends the global model to the neighboring base stations and resets the model update difference value; the base station n aggregates all the parameters of the model, that is, Send to neighboring base stations, 1≤n≤N;

[0055] Step 11: Each base station updates the model according to the received global model parameters; base station n receives the global model parameters from its neighboring base stations. n 1 ∈Ψ n , the model parameters of base station n are updated as: Among them, h n is the weight of base station n;

[0056] Step 12: The base station determines whether the model has converged; the global model convergence criterion is Among them, ε is the convergence threshold of the global model of each cell. If the convergence criterion is met, the training of the federated learning model is completed, otherwise return to step 1.

[0057] Finally, it should be noted that the above embodiments are only used to illustrate the technical solution of the present invention rather than to limit it. Although the present invention has been described in detail with reference to the preferred embodiments, those skilled in the art should understand that the technical solution of the present invention can be modified or replaced by equivalents without departing from the purpose and scope of the technical solution, which should be included in the scope of the claims of the present invention.

Claims

1. A multi-cell federated learning model training method, characterized by: The method comprises the following steps: S1: Each cell user trains a user local model based on local data and sends the local model parameters to the associated base station; S2: Each base station receives the user local model parameters sent by the associated user, globally aggregates the model, and generates a base station global model; S3: Each base station records the number of iterations of the base station global model and periodically broadcasts it to neighboring base stations; S4: Each base station calculates the difference between the number of iterations of its global model and the number of iterations of the neighboring base station model, and determines whether the difference value is higher than a preset threshold value of the number of iterations difference. If so, execute S5, otherwise, execute S6; S5: The base station suspends model iteration and returns to S4; S6: each base station calculates the model update difference; S7: The base station determines whether the model update difference value is lower than the model update difference threshold value. If so, execute S8; otherwise, execute S9; S8: The base station continues to perform iterative model updates, calculates the model update difference value, and returns to S7; S9: The base station detects whether the available bandwidth is lower than the available bandwidth threshold value. If so, S10 is executed; otherwise, S11 is executed. S10: The base station determines the update difference value of each weight in the model and compares it with the threshold value. If the update difference value of the weight is higher than the threshold value, the corresponding weight is sent to the neighboring base station and the process goes to S12; S11: The base station sends the global model parameters at the current moment to the neighboring base stations and resets the model update difference; S12: Each base station updates the local model according to the received global model parameters and determines whether the model has converged. If so, the model training is completed; otherwise, the process returns to S1.

2. A multi-cell federated learning model training method according to claim 1, characterized in that: In S1, each cell user trains a user local model based on local data and sends the local model parameters to the associated base station. Specifically, N cells are included, each cell consists of a cellular base station and multiple users, and the number of users in cell n is U n , 1≤n≤N; each user collects data from the external environment as the input for model training. The input of user i in cell n is expressed as where x n,i,m is the mth sample collected by user i in cell n, 1≤m≤M i , M i is the number of samples collected by user i in cell n, then the output of user i in cell n can be expressed as User based on X n,i and Y n,i Train the user local model, let represents the local model parameter set corresponding to the tth iteration of user i in cell n, where is the kth parameter of the local model corresponding to the tth iteration of user i in cell n, 1≤k≤K, K is the number of model parameters, and user i in cell n transmits Send to base station n.

3. A multi-cell federated learning model training method according to claim 2, characterized in that: In S2, the base station receives the user local model parameters sent by the associated user, performs global aggregation on the model, and generates a base station global model. Specifically, the base station n receives the local model parameters sent by the associated user. Perform global model parameter aggregation based on weighted average, let represents the global model parameters determined by base station n in the tth iteration, where is the kth parameter of the global model corresponding to the tth iteration of base station n, 1≤k≤K, modeling for Among them, g n,i is the weight of user i in base station n, and the base station further Sent to the user it is associated with for subsequent iterations of the user.

4. A multi-cell federated learning model training method according to claim 3, characterized in that: In S3, each base station records the number of iterations of the base station global model and periodically broadcasts it to neighboring base stations. Specifically, let t n represents the number of iterations of the aggregation model at base station n. Base station n will n The value is sent to its neighboring base stations.

5. A multi-cell federated learning model training method according to claim 4, characterized in that: In S4, each base station calculates the difference between its global model iteration number and the neighbor base station model iteration number, and determines whether the difference value is higher than a preset iteration number difference threshold value, specifically: let Ψ n represents the set of neighboring base stations of base station n. The number of iterations of the aggregation model received by base station n from each neighboring base station is Denotes the neighboring base station n1∈Ψ of base station n n The number of iterations of the current aggregation model, Δt represents the maximum difference in the number of iterations between base station n and its neighboring base stations, that is: If Δt n >Δt th , that is, if the current iteration number difference of base station n is higher than the iteration number difference threshold, base station n suspends model iteration until the difference value is lower than the threshold value.

6. A multi-cell federated learning model training method according to claim 5, characterized in that: In S6, each base station calculates the model update difference, specifically: each base station calculates the difference between the global model corresponding to the current iteration and the previous iteration, and sets Δω n,t It represents the difference between the model of the t-th iteration of base station n and the model of the previous iteration, which can be modeled as 7. A multi-cell federated learning model training method according to claim 6, characterized in that: In S7, it is determined whether the model update difference value is lower than the model update difference threshold value, specifically: let Δω th To update the difference value threshold of the model, base station n compares Δω n,t With Δω th , if Δω n,t <Δω th , then execute S8, otherwise, execute S9.

8. A multi-cell federated learning model training method according to claim 7, characterized in that: In S9, the base station detects whether the available bandwidth is lower than the available bandwidth threshold value, specifically: let B n,t is the available bandwidth of base station n at the tth iteration, is the available bandwidth threshold of base station n, base station n compares B n,t and like Then execute S10, otherwise, execute S11.

9. A multi-cell federated learning model training method according to claim 8, characterized in that: In S10, the base station determines the update difference value of each weight in the model and compares it with the threshold value, specifically: let Δω n,t,k represents the updated difference value of the kth weight in the model of base station n, which is modeled as Let Δω′ th Represents the model weight update difference threshold, if Δω n,t,k ≥Δω′ th , then base station n will Send to neighboring base stations, otherwise, base station n does not send 10. A multi-cell federated learning model training method according to claim 9, characterized in that: In S11, the base station sends the global model parameters at the current moment to the neighboring base stations. Specifically, base station n aggregates all the parameters of the model, i.e. Send to neighboring base stations, 1≤n≤N; In S12, each base station updates the local model according to the received global model parameters and determines whether the model converges. Specifically, base station n receives the global model parameters from its neighboring base stations. The model parameters of base station n are updated as follows: Among them, h n is the weight of base station n; the global model convergence criterion is Among them, ε is the convergence threshold of the global model of each cell.

Citation Information

Patent Citations

  • Reliable federated learning method and system based on terminal reputation in wireless network

    CN112153650A

  • Wireless service traffic prediction method based on weighted federated learning

    WO2021169577A1