Federal learning method based on distribution perception regularization
By using generative adversarial networks in federated learning for data augmentation and divergence measurement, combined with the updated parameter aggregation within the time window, the overfitting and training accuracy problems caused by device heterogeneity are solved, and the convergence speed and generalization performance of the global model are improved.
Patent Information
- Application Number
- CN202510084861.0
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-01-20
- Publication Date
- 2025-05-13
- Estimated Expiration
- 2045-01-20
AI Technical Summary
In a federated learning environment, device heterogeneity leads to differences in computing power and data set size, increasing the risk of overfitting, affecting the training accuracy and the convergence speed of the learning process. At the same time, the data distribution differences between different clients lead to slowing down the convergence speed of the global model.
By generating adversarial networks, data augmentation of clients whose computing power and data set sizes do not match, and using divergence to measure the differences between global data distribution and local data distribution, update parameters within the time window are collected to alleviate the problem of inconsistent upload parameters of the device.
It effectively alleviates the risk of overfitting and training accuracy problems caused by device heterogeneity, improves the convergence speed and generalization performance of the global model, and solves the problem of slowing down the convergence speed of the global model caused by the difference in client data distribution.
Smart Images

Figure CN119988019A_ABST
Abstract
Description
Technical Field
[0001] The present invention belongs to the field of application of federated learning technology, and specifically relates to a federated learning method based on distribution-aware regularization. Background Art
[0002] In recent years, researchers have conducted extensive research in the field of device heterogeneity. According to different research focuses, they can be roughly divided into three categories: client selection, aggregation improvement, and asynchronous federation. Among them, the research on client selection is mainly used to solve the limitations of device heterogeneity on training models. Such methods include considering the resources and computing power of the device when selecting the client, or adjusting the complexity and size of the model according to the client's situation. The research on aggregation improvement is mainly used to solve the problem of device unreliability, and reduce the impact of the device by changing the aggregation method. Then the research on asynchronous federation is mainly used to alleviate the problem of stragglers in heterogeneous environments. Through asynchronous methods, the server does not have to wait for the device to respond, which improves the flexibility of participating devices.
[0003] Although there are various solutions to the device heterogeneity problem in federated learning, the following problems still exist:
[0004] 1) A federated environment usually involves a large number of devices with different computing capabilities and dataset sizes, which will increase the risk of overfitting;
[0005] 2) Since some clients will train an old version of the global model due to device heterogeneity, and the server has updated the global model multiple times with other clients, this obsolescence will have a negative impact on the training accuracy and slow down the convergence of the learning process;
[0006] 3) When the data distribution on different clients is quite different, the local models trained by each client may conflict with each other when aggregated in the global model, resulting in slower convergence of the global model. Summary of the invention
[0007] To solve the above problems, the present invention provides a federated learning method based on distribution-aware regularization, comprising the following steps:
[0008] S1. Initialize the historical training data of all clients, set the number of clients participating in each round of federated learning M; set the maximum number of federated learning rounds, and initialize the parameter t = 1;
[0009] S2. In the tth round of federated learning, the server sends the global model to all clients, and the server determines the M clients that participate in this training;
[0010] S3. The server obtains the data processing method of each client according to the computing power and the amount of local data, and each client obtains training data according to its corresponding data processing method;
[0011] S4. Each client participating in this training trains the global model with the training data to obtain the local model parameters and uploads them to the server. During the training process, each client calculates the local loss according to the number of samples of different categories in the training data.
[0012] S5. The server calculates the time window based on the historical training data, collects local model parameters within the time window, and aggregates them to obtain a new global model;
[0013] S6. Determine whether the maximum number of federated learning rounds has been reached. If so, end the iteration. If not, set t=t+1 and return to step S2.
[0014] Beneficial effects of the present invention:
[0015] The present invention generates adversarial networks to expand data for clients whose computing power and data set size do not match, and uses divergence to measure the difference between global data distribution and local data distribution. At the same time, the update parameters within the time window are collected during model aggregation to alleviate the problem of inconsistent device upload parameters. BRIEF DESCRIPTION OF THE DRAWINGS
[0016] Figure 1 is a flow chart of the method of the present invention;
[0017] Figure 2 This is a schematic diagram of data balancing of the client of the present invention;
[0018] Figure 3 Schematic diagram of divergence calculation of the client of the present invention. DETAILED DESCRIPTION
[0019] The following will be combined with the drawings in the embodiments of the present invention to clearly and completely describe the technical solutions in the embodiments of the present invention. Obviously, the described embodiments are only part of the embodiments of the present invention, not all of the embodiments. Based on the embodiments of the present invention, all other embodiments obtained by ordinary technicians in this field without creative work are within the scope of protection of the present invention.
[0020] The present invention provides a federated learning method based on distribution-aware regularization, such as Figure 1 As shown, the following steps are included:
[0021] S1. Initialize the historical training data of all clients, set the number of clients participating in each round of federated learning M, set the maximum number of federated learning rounds, and initialize the parameter t=1.
[0022] Specifically, the historical training data includes: the time when the client receives the global model, the time when the client uploads the local model, the transmission power when the client uploads the local model, the channel gain between the client and the server, the signal interference when the client uploads the local model, the number of local samples of the client, the number of CPU cycles required for the client to train the global model once using a local sample, and the number of model parameters of the global model.
[0023] S2. In the tth round of federated learning, the server sends the global model to all clients, and the server determines the M clients that participate in this training.
[0024] S3. The server obtains the data processing method of each client according to the computing power and the amount of local data, and each client obtains training data according to its corresponding data processing method.
[0025] Specifically, step S3 includes:
[0026] S31. Calculate the preliminary score based on the computing power and local data volume of each client, expressed as
[0027]
[0028] Among them, D k represents the amount of local data of client k, represents the computing power of client k, m k represents the preliminary score of client k;
[0029] In this embodiment, the computing power of client k is the time required for client k to perform a local iterative training using one local sample. The specific calculation formula is:
[0030]
[0031] Among them, f k represents the computing resources of client k, that is, the CPU cycle frequency of client k; b k It indicates the number of CPU cycles required for client k to train a local model using a local sample. The local sample can be an image, statistical data, etc.
[0032] S32. Since the initial score may be greater than 1, the softmax function is used to normalize the preliminary score of each client, so that the server can integrate the data of each client more balanced, thereby improving the robustness and accuracy of the global model. The device score is expressed as
[0033]
[0034] Among them, ε krepresents the device score of client k, and N represents the total number of clients;
[0035] S33. Calculate the data threshold p according to the device score, expressed as
[0036]
[0037] S34. For each client participating in this training, if its device score is greater than the data threshold, the client uses the generative adversarial network (DCGAN) to perform data enhancement to obtain training data; if its device score is not greater than the data threshold, the client randomly samples the local data to obtain training data, such as Figure 2 shown.
[0038] S4. Each client participating in this training trains the global model through the training data to obtain the local model parameters and uploads them to the server; during the training process, each client calculates the local local loss based on the number of samples of different categories in the training data.
[0039] Specifically, in step S4, any client participating in this training calculates the local loss according to the number of samples of different categories in the training data, including:
[0040] S41. Calculate the normalized vector of each category sample of client k, expressed as
[0041]
[0042] Among them, p(c k ) represents the normalized vector corresponding to the sample of category c on client k, N represents the number of samples of category c in the training data of client k, k represents the total number of samples of training data of client k;
[0043] S42. Calculate the KL divergence of client k based on the normalized vector, expressed as
[0044]
[0045] Among them, KL k (p k ||q) represents the KL divergence of client k, q(c) represents the sum of normalized vectors corresponding to samples of category c, C represents the number of categories, and M represents the number of clients participating in the training;
[0046] By comparing the normalized vector of each class sample of the client with the sum of the normalized vectors of all participants on the class samples, the server is able to evaluate the difference between each client's data and the overall data set distribution. This evaluation is critical for optimizing the global model because it helps determine which clients' data contributes most to the training of the overall model. By quantifying the differences between different data sources, we can better understand the diversity of client data and adjust the model training strategy accordingly to ensure that the model can better generalize to different data distributions.
[0047] S43. Figure 3 As shown, the local loss of client k is calculated according to the KL divergence Expressed as
[0048]
[0049]
[0050] Among them, ω represents the initial weight parameter of the model, ω t represents the weight parameter updated by the local model in the tth round of federated learning, F k (ω) represents the initial loss function of client k, x k represents the sample extracted from the training data of client k, f k (ω; x k ) is the initial weight parameter ω of the model in sample x k The loss function on represents the expected value of the prediction loss of client k on the local dataset.
[0051] S5. The server calculates the time window based on the historical training data, collects local model parameters within the time window, and aggregates them to obtain a new global model.
[0052] Specifically, step S5 includes:
[0053] S51. Record the training time of each client participating in the previous round of federated learning, extract the latest training time of all clients for DBSCAN cluster analysis, and select the maximum value of the cluster diameter as the size of the time window of the current round of federated learning, that is, calculate the time window T t It can be expressed as
[0054]
[0055] in, represents the diameter of the i-th cluster in the t-1th round of federated learning; t a and t b Represents cluster C i The training time of any two points in represents the training time of client k, β k represents the batch size of client k, N represents the time required for client k to execute a single batch. k represents the training data size of client k;
[0056] The time window is calculated before each round of federated learning starts, and the time is counted when federated learning starts. The time window is used to limit the number of clients collected during each round of federated learning aggregation. The global model aggregates as many clients as the local model parameters of the clients collected during this time.
[0057] S52. Calculate the delay function S corresponding to the client k participating in this training k , expressed as
[0058]
[0059] in, It represents the sum of the number of local batches collected by the server in the tth round of federated learning until the client k returns the local model parameters. represents the sum of the number of local batches collected by the server in the t-1th round of federated learning, represents the number of local batches used by client k to train the local model in the tth round of federated learning, and a represents the hyperparameter that controls the decay speed of the function (a>0);
[0060] S53. The delay function is used to measure the delay degree of the clients participating in the training within the time window, and the mixed hyperparameters of the client k in the current federated learning round are calculated, which is expressed as
[0061]
[0062] in, represents the hybrid hyperparameters of client k in the tth round of federated learning;
[0063] S54. Start timing at the start time of the current federated learning round, and perform asynchronous aggregation every time a local model parameter is collected until the timing period reaches the time window corresponding to the current federated learning round to obtain the global model parameters of the current federated learning round, expressed as
[0064]
[0065] Among them, w t-1 represents the global model parameters after the t-1th round of federated learning, represents the initial aggregation parameters of the tth round of federated learning, represents the global model parameters after aggregating the local model parameters of the kth client in the tth round of federated learning, represents the local model parameters of the kth client.
[0066] S6. Determine whether the maximum number of federated learning rounds has been reached. If so, end the iteration. If not, set t=t+1 and return to step S2.
[0067] This paper proposes a learning model with distribution-aware regularization under federated learning to alleviate the problem of device heterogeneity among multiple clients. By proposing data balancing and aggregating the overall distribution of local data of different clients, the convergence direction of the global model does not deviate from the optimal direction, which improves the convergence speed and generalization performance of the global model and solves the client drift problem caused by data heterogeneity.
[0068] In the present invention, unless otherwise clearly stipulated and limited, the terms such as "installation", "setting", "connection", "fixation" and "rotation" should be understood in a broad sense. For example, it can be a fixed connection, a detachable connection, or an integral one; it can be a mechanical connection or an electrical connection; it can be directly connected or indirectly connected through an intermediate medium; it can be the internal connection of two elements or the interaction relationship between two elements. Unless otherwise clearly defined, ordinary technicians in this field can understand the specific meanings of the above terms in the present invention according to the specific circumstances.
[0069] Although embodiments of the present invention have been shown and described, it will be appreciated by those skilled in the art that various changes, modifications, substitutions and variations may be made to the embodiments without departing from the principles and spirit of the present invention, and that the scope of the present invention is defined by the appended claims and their equivalents.
Claims
1. A federated learning method based on distribution-aware regularization, characterized in that: The following steps are involved: S1. Initialize the historical training data of all clients and set the number of clients participating in each round of federated learning M; Set the maximum number of federated learning rounds and initialize the parameter t=1; S2. In the tth round of federated learning, the server sends the global model to all clients, and the server determines the M clients that participate in this training; S3. The server obtains the data processing method of each client participating in this training according to the computing power and the amount of local data, and each client obtains training data according to its corresponding data processing method; S4. Each client participating in this training trains the global model with the training data to obtain local model parameters, and uploads them to the server; During the training process, each client calculates the local loss based on the number of samples of different categories in the training data for training; S5. The server calculates the time window based on the historical training data, collects local model parameters within the time window, and aggregates them to obtain a new global model; S6. Determine whether the maximum number of federated learning rounds has been reached. If so, end the iteration. If not, set t=t+1 and return to step S2.
2. The method for federated learning based on distribution-aware regularization according to claim 1, characterized in that: Step S3 specifically includes: S31. Calculate the preliminary score based on the computing power and local data volume of each client, expressed as Among them, D k represents the amount of local data of client k, represents the computing power of client k, m k represents the preliminary score of client k; S32. Normalize the preliminary score of each client to obtain the device score, expressed as Among them, ε k represents the device score of client k, and N represents the total number of clients; S33. Calculate the data threshold p according to the device score, expressed as S34. For each client participating in this training, if its device score is greater than the data threshold, the client uses the generative adversarial network to perform data enhancement to obtain training data; if its device score is not greater than the data threshold, the client randomly samples the local data to obtain training data.
3. The method for federated learning based on distribution-aware regularization according to claim 1, characterized in that: Step S4: Any client participating in this training calculates the local loss according to the number of samples of different categories in the training data, including: S41. Calculate the normalized vector of each category sample of client k, expressed as Among them, p(c k ) represents the normalized vector corresponding to the sample of category c on client k, N represents the number of samples of category c in the training data of client k, k represents the total number of samples of training data of client k; S42. Calculate the KL divergence of client k based on the normalized vector, expressed as Among them, KL k (p k ||q) represents the KL divergence of client k, q(c) represents the sum of normalized vectors corresponding to samples of category c, and C represents the number of categories; S43. Calculate the local partial loss of client k according to the KL divergence.
4. The method for federated learning based on distribution-aware regularization according to claim 3, characterized in that: Local partial loss Expressed as Among them, ω represents the initial weight parameter of the model, ω t represents the weight parameter updated by the local model in the tth round of federated learning, F k (ω) represents the initial loss function of client k, x k represents the sample extracted from the training data of client k, f k (ω; x k ) is the initial weight parameter ω of the model in sample x k The loss function on represents the expected value of the prediction loss of client k on the local dataset.
5. The method for federated learning based on distribution-aware regularization according to claim 1, characterized in that: Step S5 specifically includes: S51. Calculate the time window T of the tth round of federated learning t , expressed as in, represents the diameter of the i-th cluster; S52. Calculate the delay function S corresponding to the client k participating in this training k , expressed as in, It represents the sum of the number of local batches collected by the server in the tth round of federated learning until the client k returns the local model parameters. represents the sum of the number of local batches collected by the server in the t-1th round of federated learning, represents the number of local batches used by client k to train the local model in the tth round of federated learning, and a represents the hyperparameter that controls the decay speed of the function; S53. Calculate the mixed hyperparameters of client k in the current federated learning round, expressed as in, represents the hybrid hyperparameter of client k in the tth round of federated learning, and α represents the initial hybrid hyperparameter value; S54. Start timing at the start time of the current federated learning round, and perform asynchronous aggregation each time a local model parameter is collected until the timing period reaches the time window corresponding to the current federated learning round to obtain the global model parameters of the current federated learning round.
6. The method for federated learning based on distribution-aware regularization according to claim 5, characterized in that: The aggregation process of step S54 is expressed as Among them, w t-1 represents the global model parameters after the t-1th round of federated learning, represents the initial aggregation parameters of the tth round of federated learning, represents the global model parameters after aggregating the local model parameters of the kth client in the tth round of federated learning, represents the local model parameters of the kth client.
Citation Information
Patent Citations
Federal learning algorithm based on diffusion model and weight adaptive knowledge distillation
CN116665000A
Federal learning method for non-independent identically distributed and unbalanced data set
CN117829270A
Federal learning method for distributed data resources under privacy protection constraint
CN118428491A
Heterogeneous data federal learning method based on improved aggregation algorithm
CN118821909A
Federated learning scheduling method, device, and system
WO2022116323A1
Cited By
Distributed cross-border storage network inventory management method and system based on federated learning
CN121616198A
Federal learning model training method based on distributed random convex difference optimization
CN121787516A
Bayesian federal learning method based on double prior perception mechanism
CN121920571A
Federal learning optimization method and system for unbalanced data set
CN122065089A