A Data Heterogeneity Adaptive Federated Learning Method, Medium and Device
The data-heterogeneous adaptive federated learning method addresses non-IID issues by quantifying heterogeneity, training invariant representations, and adjusting aggregation weights, improving model accuracy and convergence in heterogeneous data environments.
Patent Information
- Application Number
- CN202510600778.4
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2025-05-12
- Publication Date
- 2025-07-15
- Estimated Expiration
- 2045-05-12
AI Technical Summary
In federated learning, due to data heterogeneity between clients, local optimization goals are inconsistent with global optimization goals, resulting in low accuracy and slow convergence speed in data heterogeneity scenarios.
By training an encoder that can extract client-side invariant characterization on the client, combining the quantitative indicators of data heterogeneity and the size of the local data set, the aggregate weights are adaptively determined, and the historical global model is injected later in the training stage to optimize the global model update.
It effectively alleviates the problems of poor model accuracy and slow convergence speed caused by data heterogeneity, improves the overall training performance, especially in the early stage of training, and accelerates the convergence speed of the model in the later stage.
Smart Images

Figure CN120124779B_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to the technical field of federated learning, and specifically to a data heterogeneous adaptive federated learning method, medium, and device. Background Art
[0002] The core of supporting artificial intelligence (AI) training is data, especially accurate and representative high-quality data with distribution. It has become a general trend for data to flow freely on the premise of security and compliance. Facing the huge potential value of data owned by commercial companies, even between two companies or departments within the same company, they have to consider the exchange of interests. Often, these institutions will not provide their respective data for direct aggregation with other companies, resulting in data often appearing in the form of isolated islands even within the same company.
[0003] Based on the above three points of insufficient support for implementation, not allowing rough exchange, and unwillingness to contribute value, it has led to a large number of existing data islands and privacy protection problems, and federated learning has emerged as the times require. Federated learning (FL) can utilize the private data of different devices without compromising privacy, and reduce the overhead caused by direct data interaction through local training models. Federated learning is essentially a distributed machine learning framework, and the devices participating in federated learning are regarded as different clients, and multiple clients can perform joint machine learning modeling under the requirement of protecting data privacy.
[0004] However, in reality, a major obstacle to deploying FL applications is the data heterogeneity between clients, that is, the data across clients is non-independent and identically distributed (non-IID). The non-IID problem in FL can be divided into two categories: label shift and feature shift. On the one hand, when clients have different class label distributions, even if the input feature space is the same, label shift will occur. On the other hand, although clients have a common label space, when the underlying distribution of input features varies between different clients, the phenomenon of feature shift will occur. Due to the existence of non-IID situations, the local optimization objectives of clients are inconsistent with the global optimization objective. Therefore, data heterogeneity may cause local models to converge in different directions, reaching local optima rather than the global optimum, thereby reducing the performance of federated learning, and may even be worse than local learning without federated communication. Summary of the Invention
[0005] The present invention provides a data heterogeneous adaptive federated learning method, medium, and device that can effectively alleviate the problems of low model accuracy and slow convergence speed caused by the two shift phenomena in data heterogeneity, and improve the overall training performance, and can at least solve one of the above technical problems.
[0006] To solve the above technical problems, the present invention adopts the following technical solutions:
[0007] A data heterogeneous adaptive federated learning method, comprising the following steps:
[0008] S1. The server initializes the global model parameters, broadcasts the global model and the assumed global label distribution to all clients participating in the training. The clients calculate the corresponding local label distribution according to their own local datasets, and calculate the difference between the global label distribution and the local label distribution as a quantization index of data heterogeneity;
[0009] S2. After receiving the global model parameters, the clients participating in the training perform local client model training locally, and at the same time train an encoder capable of extracting client-invariant representations;
[0010] S3. The server jointly determines the aggregation weights based on the quantization index of data heterogeneity and the size of the local dataset, so as to adaptively aggregate the client models according to the size of data heterogeneity;
[0011] S4. If the client model has reached the later stage of the specified training phase, inject the historical global model and perform weighted averaging with the global model aggregated in the current round as the global model updated in this round. If the client model has not reached the later stage of the specified training phase, use the aggregated client model as the global model updated in the current round;
[0012] S5. The server broadcasts the updated global model to each client participating in the next round of training and enters the next round of local training until the predetermined number of communication rounds is reached, and the global training is completed.
[0013] Further, the S1 further includes:
[0014] S11. The global label distribution and the local label distribution are respectively represented as Dis global and Dis local ;
[0015] The global label distribution Dis global treats all C label categories equally. For each category c, it is denoted as:
[0016]
[0017] The local label distribution of the kth client is defined as , denoted as:
[0018]
[0019] where the cth element represents the proportion of the data volume of the cth category. B k is the dataset of the kth client, x i and yi They are respectively the input of the i-th sample and the corresponding label;
[0020] S12, the local label distribution Dis evaluated by each client local and the global label distribution Dis sent by the server global The difference d k between them, the difference d k is defined as:
[0021]
[0022] This is used to capture the data label heterogeneity of the client side and serves as a quantitative indicator of data heterogeneity;
[0023] Among them, discrepancy(.) is a predefined metric function, and the KL divergence is used as the metric function, then the difference d k is defined as:
[0024]
[0025] As a supplementary indicator of the aggregation weight, it is uploaded to the server together with the client model trained in this round.
[0026] Furthermore, in S2, the client model training process undergoes multiple rounds of iteration. In each round of iteration, the client model parameters are updated to prompt the client model to extract client-invariant representations that only contain the required non-redundant feature information.
[0027] Furthermore, in S2, in each communication round t, where t = 1, 2,..., T, the client subset K t from all client sets K is active and all download the global model from the server to participate in the training of the current round;
[0028] In each local step of the t-th communication round, client k randomly samples a batch of training data k from the corresponding local dataset B , selects the CIRL algorithm based on federated learning to learn the client-invariant representation, reduces the difference between the global dataset and the local dataset in the feature space, reduces the difference between the global model and the local client model, and achieves a lower loss in local training. Thus, the objective function f is designed and the gradient is used to update the local client model , and the updated local client model is denoted as:
[0029]
[0030] Among them, η is the learning rate;
[0031] After E steps of local training, each client will send its trained local client model to the server.
[0032] Furthermore, in S2, the CIRL algorithm based on federated learning includes the following steps:
[0033] S2-a. Each client maps the input x of the sample to a representation z through an encoder network, and the distribution of the representation z is denoted as p(z|x), which is parameterized by the mean and standard deviation parameters. Learn a classifier from the representation z, which uses the prediction distribution parameterized by w to predict (where the parameter w is omitted);
[0034] S2-b. Taking L as the loss function, the supervised loss function on client k is defined as:
[0035]
[0036] where y is the label corresponding to the input x;
[0037] S2-c. To learn a representation z with a consistent conditional distribution p k (z|x) across all clients, use the conditional mutual information I k (x, z|y) between the data input with the given label and the representation to limit the amount of information that the representation can contain, defined as follows:
[0038]
[0039] Minimize to force the representation z to contain only the non-redundant feature information required to predict the label y, and not to contain other information about x that is irrelevant to the label;
[0040] where r(z|y i ) is the conditional distribution of z given y, used as the reference distribution. For the classification task, set r(z|y i ) to a Gaussian distribution with mean and variance , defined as follows:
[0041]
[0042] where μ y and σ yis the parameter to be optimized, y = 1, 2, ..., C, where C is the number of categories;
[0043] S2-d. Obtain the local objective function for each client k:
[0044]
[0045] where α CMI is the hyperparameter of the conditional mutual information regularization term;
[0046] Each client performs gradient descent in multiple rounds of iterative training for the above local objective function, continuously updating the local client model to gradually approach the optimal solution.
[0047] Furthermore, in S3 - S4, while collecting the client models, the server also collects the quantization metrics of the data heterogeneity of each client as supplementary metrics for the aggregation weights. When the communication round reaches the startup round t - startup, the historical global model is introduced based on the current global model for update.
[0048] Furthermore, S3 further includes:
[0049] S31. The server uses the relative size n k of the local dataset and the quantization metric d k to determine more discriminative aggregation weights p k for each client model:
[0050]
[0051] where ReLU(.) is the relu function for handling negative values, a is the hyperparameter for balancing n k and d k , and b is another hyperparameter for adjusting the weights;
[0052] S32. Aggregate the client models and update to obtain the global model for this round as:
[0053]
[0054] where K t is all the clients participating in the training in the t-th round, which is a subset of the client set K.
[0055] Furthermore, S4 further includes:
[0056] S41. If the current communication round t is greater than or equal to the set startup round t-startup, then introduce the historical global models of the previous S rounds including the current global model, and perform weighted averaging with the current global model to update the global model of this round as:
[0057]
[0058] Among them, S is the window size of the historical global model, indicating the number of historical global models to be injected;
[0059] S42. If the current communication round t is less than or equal to the set startup round t-startup, then do not introduce the historical global model, and the update of the global model of this round is the model obtained after aggregating the client models:
[0060] .
[0061] A computer-readable storage medium stores a computer program. When the computer program is executed by a processor, the processor is caused to execute the steps of the above data heterogeneous adaptive federated learning method.
[0062] A computer device includes a memory and a processor. The memory stores a computer program. When the computer program is executed by the processor, the processor is caused to execute the steps of the above data heterogeneous adaptive federated learning method.
[0063] The beneficial effects of the present invention are embodied in:
[0064] In the present invention, on the one hand, in the early stage of training, all clients capture the local data heterogeneity as a quantization index. On the other hand, in the early stage of the training phase, since the local training adopts the client-invariant representation learning method, the difference between the global and local data is reduced in the feature space, prompting the client model to extract client-invariant representations that only contain the required non-redundant feature information. Therefore, the accuracy rate rises rapidly in the early stage of the training phase. On the other hand, in the later stage of the training phase, due to the injection of the historical global model, the degree of the non-IID phenomenon caused by the missing data samples can be reduced, preventing the global model update from deviating from the global optimal direction. At the same time, it also provides a better initialization for the client model to perform the next round of local training, thereby accelerating the convergence speed of the model. In this way, the problems of poor model accuracy and slow convergence speed caused by label shift and feature shift in the data heterogeneous scenario can be alleviated simultaneously, thus effectively improving the overall training performance. BRIEF DESCRIPTION OF THE DRAWINGS
[0065] The drawings described herein are used to provide a further understanding of the present application, and constitute a part of the present application. The illustrative embodiments of the present application and their descriptions are used to explain the present application and do not constitute an improper limitation to the present application.
[0066] Figure 1 It is a schematic diagram of the overall process of the data heterogeneous adaptive federated learning method according to an embodiment of the present invention.
[0067] Figure 2 It is a schematic diagram of the specific process of the data heterogeneous adaptive federated learning method according to an embodiment of the present invention.
[0068] Figure 3 It is a simulation diagram of the model accuracy of the comparative experiment in the test set according to an embodiment of the present invention.
[0069] Figure 4 It is a structural block diagram of a computer device according to an embodiment of the present invention. Detailed implementation manners
[0070] Next, the technical solutions in the embodiments of the present invention will be clearly and completely described in conjunction with the accompanying drawings in the embodiments of the present invention. Obviously, the described embodiments are only a part of the embodiments of the present invention, rather than all the embodiments. Without conflict, the embodiments in the present application and the features in the embodiments can be combined with each other. All other embodiments obtained by those of ordinary skill in the art based on the embodiments of the present invention without creative efforts shall fall within the protection scope of the present invention.
[0071] It should be noted that the meaning of "and / or" appearing throughout the text includes three parallel solutions. Taking "A and / or B" as an example, it includes solution A, or solution B, or a solution that satisfies both A and B at the same time. In addition, "a plurality of" means two or more. In addition, the technical solutions between the embodiments can be combined with each other, but it must be based on the fact that those of ordinary skill in the art can implement it. When the combination of technical solutions appears to be contradictory or unable to be implemented, it should be considered that such a combination of technical solutions does not exist and is not within the protection scope required by the present invention.
[0072] See Figure 1 - Figure 2 , an embodiment of the present invention provides a data heterogeneous adaptive federated learning method, including the following steps:
[0073] S1. The server initializes the global model parameters, broadcasts the global model and the assumed global label distribution to all participating clients for training. The clients calculate the corresponding local label distribution according to their own local datasets, and calculate the difference between the global label distribution and the local label distribution as a quantization index of data heterogeneity;
[0074] S2. After receiving the global model parameters, the participating clients perform client model training locally, and at the same time train an encoder that can extract client-invariant representations;
[0075] S3. The server jointly determines the aggregation weights based on the quantization index of data heterogeneity among servers and the size of the local dataset, so as to adaptively aggregate the client models according to the degree of data heterogeneity.
[0076] S4. If the client model has reached the later stage of the specified training phase, inject the historical global model and perform weighted averaging with the globally aggregated model in the current round as the globally updated model for this round. If the client model has not reached the later stage of the specified training phase, use the aggregated client model as the globally updated model for the current round.
[0077] S5. The server broadcasts the updated global model to each client participating in the training in the next round and enters the next round of local training until the predetermined number of communication rounds is reached, completing the global training.
[0078] In this embodiment, in S1, each client needs to calculate the difference between its local label distribution and the assumed global label distribution. Intuitively, the global label distribution naturally promotes fairness across all classes and the generalization ability of the global model, so it follows a uniform distribution. In addition, each local client can calculate its difference without sharing additional data, thus preventing the leakage of information about the class distribution.
[0079] Specifically, S1 further includes:
[0080] S11. The global label distribution and the local label distribution are respectively denoted as Dis global and Dis local ;
[0081] The global label distribution Dis global treats all C label classes equally. For each class c, it is denoted as:
[0082]
[0083] The local label distribution of the k-th client is defined as , denoted as:
[0084]
[0085] where the c-th element represents the proportion of the data volume of the c-th class. B k is the dataset of the k-th client, and x i and y i are respectively the input and the corresponding label of the i-th sample;
[0086] S12. Each client evaluates the difference between the local label distribution Dis local and the global label distribution Dis issued by the serverglobal The difference d between k , the difference d k is defined as:
[0087]
[0088] to capture the data label heterogeneity on the client side locally and serve as a quantitative metric for data heterogeneity;
[0089] where discrepancy(.) is a predefined metric function, and the KL divergence is used as the metric function, then the difference d k is defined as:
[0090]
[0091] as a supplementary metric for the aggregation weight and uploaded to the server together with the client model trained in this round;
[0092] In the present invention, when selecting the metric function for discrepancy(.), functions that can measure vector distances such as L1 distance, L2 distance, KL divergence, etc. can be selected.
[0093] In this embodiment, in S2, the training process of the client model can be iterated multiple times, and the client model parameters are updated in each iteration to prompt the client model to extract client-invariant representations that only contain the required non-redundant feature information.
[0094] In this embodiment, in S2, in each communication round t, t = 1, 2,..., T, the client subset K t from all client sets K is in an active state and all download the global model from the server to participate in the training of the current round;
[0095] In each local step of the t-th communication round , client k randomly samples a batch of training data from the corresponding local dataset B k and selects the CIRL algorithm based on federated learning to learn the client-invariant representation, reducing the difference between the global dataset and the local dataset in the feature space to reduce the difference between the global model and the local client model and achieving a lower loss in local training. Thus, the objective function f is designed and the gradient is used to update the local client model The updated local client model is denoted as:
[0096]
[0097] where η is the learning rate;
[0098] After E steps of local training, each client sends its respective trained local client model to the server.
[0099] However, heterogeneous data distributions (including label heterogeneity and feature heterogeneity) still affect the training performance in terms of convergence and accuracy. Therefore, during local training, the feature heterogeneity problem will be addressed, where clients have different input distributions (domains) for a given label class y, i.e., c with:
[0100]
[0101] where, and .
[0102] Due to differences in the natural conditions and processing methods of edge devices, feature skew problems often occur in edge AI applications. For example, images taken by smartphones in rural areas may be different from those taken by smartphones in urban areas. As the feature distributions vary among clients, traditional federated learning methods (such as FedAvg) attempt to learn all possible feature knowledge during local training, including client-invariant knowledge and client-specific knowledge. However, client-specific knowledge leads to representational divergence, i.e.,
[0103]
[0104] resulting in a decrease in accuracy, where z is the intermediate feature trained by the encoder network and is called the representation.
[0105] Therefore, the present invention adopts the CIRL algorithm based on federated learning to learn the invariant representation of clients, reduce the difference between the global and local datasets in the feature space, thereby further reducing the difference between the global and local client models, and achieving lower losses in local training.
[0106] Specifically, in S2, the CIRL algorithm based on federated learning includes the following steps:
[0107] S2-a. Each client maps the input x of the sample to a representation z through an encoder network, and the distribution of the representation z is denoted as p(z|x), which is parameterized by the mean and standard deviation . Learn a classifier from the representation z, which, given the representation z, uses the prediction distribution parameterized by w to predict (where the parameter w is omitted);
[0108] S2-b. With \(L\) as the loss function, the supervised loss function on client \(k\) is defined as:
[0109]
[0110] where \(y\) is the label corresponding to the input \(x\);
[0111] S2-c. To learn a representation \(z\) with a consistent conditional distribution \(p\) k \((z|x)\) across all clients, the conditional mutual information \(I\) k \((x, z|y)\) between the data input with the given label and the representation is used to limit the amount of information that the representation can contain, which is defined as follows:
[0112]
[0113] Minimize to force the representation \(z\) to contain only the non-redundant feature information required for predicting the label \(y\), and not to contain other information about \(x\) that is irrelevant to the label;
[0114] where \(r(z|y\) i ) is the conditional distribution of \(z\) given \(y\), serving as the reference distribution. For the classification task, \(r(z|y\) i ) is set to be a Gaussian distribution with as the mean and as the variance, which is defined as follows:
[0115]
[0116] where \(\mu\) y and \(\sigma\) y are parameters to be optimized, \(y = 1, 2, \ldots, C\), and \(C\) is the number of classes;
[0117] S2-d. Obtain the local objective function for each client \(k\):
[0118]
[0119] where \(\alpha\) CMI is the hyperparameter of the conditional mutual information regularization term;
[0120] Each client performs gradient descent in multiple rounds of iterative training for the above local objective function, continuously updating the local client model to gradually approach the optimal solution.
[0121] Furthermore, in the above S2-c, for client \(k\), the conditional mutual information between the data input \(x\) with the given label and the representation \(z\) can be defined as:
[0122]
[0123] where p k (z|y) is the conditional distribution of the representation z corresponding to the label y on the client k. This term is difficult to integrate with respect to x. If we directly introduce I k (x,z|y) as the conditional mutual information term into the objective function, during iterative training with the gradient descent algorithm, the presence of p k (z|y) will make the entire mutual information term difficult to handle. Therefore, an upper bound is derived to minimize this conditional mutual information term:
[0124]
[0125] Here, the upper bound can be calculated and used as a regularizer when training the representation network, optimizing p(z|x) and r(z|y) to minimize . Since both p(z|x) and r(z|y) follow Gaussian distributions, with parameters μ w , σ w and μ y , σ y respectively, so can be calculated according to the following formula:
[0126]
[0127] .
[0128] In this embodiment, in S3 - S4, while collecting the client models, the server also collects the quantization metrics of the data heterogeneity of each client as supplementary metrics for the aggregation weights. When the communication round reaches the startup round t - startup, the historical global model is introduced on the basis of the current global model for update.
[0129] In this embodiment, S3 further includes:
[0130] S31. The server determines more discriminative aggregation weights p k for each client model using the relative size n k of the local dataset and the quantization metric d k of the local data heterogeneity:
[0131]
[0132] where ReLU(.) is the relu function for handling negative values, a is a hyperparameter for balancing n k and d k , and b is another hyperparameter for adjusting the weights;
[0133] S32. Aggregate the client models and update the global model for this round to be:
[0134]
[0135] where K t is all the clients participating in training in the t-th round, which is a subset of the client set K.
[0136] In this embodiment, the S4 further includes:
[0137] S41. If the current communication round number t is greater than or equal to the set startup round number t-startup, introduce the historical global models of the previous S rounds including the current global model, and perform weighted averaging with the current global model to update the global model for this round to be:
[0138]
[0139] where S is the window size of the historical global models, indicating the number of historical global models to be injected;
[0140] S42. If the current communication round number t is less than or equal to the set startup round number t-startup, do not introduce the historical global models, and the update of the global model for this round is the model obtained after aggregating the client models:
[0141] .
[0142] To verify the application of the present invention, a set of specific experiments will be provided below to further illustrate this data heterogeneous adaptive federated learning method:
[0143] Experimental environment:
[0144] The experiment considers training in a distributed framework consisting of 1 central server and 100 clients to be involved in training. In each communication round, 10 clients will be randomly selected to participate in training, and a total of 400 communication rounds will be carried out.
[0145] In local training, a convolutional neural network model with the ResNet18 architecture was considered and validated on the CIFAR-10 dataset. The model initialized multiple convolutional layers, batch normalization layers, ReLU activation functions, max pooling layers, and a series of residual blocks, followed by an average pooling layer and a fully connected layer. The model first passed through a 7x7 convolutional layer with 3 input channels, 64 output channels, a kernel size of 7, accompanied by a max pooling layer with a stride of 2 and a ReLU activation function. Then, the data entered a residual learning module consisting of four stages. In the first stage, there were two residual blocks, each with 64 output channels. The second stage also had two residual blocks, but the output channels increased to 128. The third stage further deepened the network, containing four residual blocks, each with 256 output channels. The last stage also contained two residual blocks, but the number of output channels was 512. When the data passed through all these convolutional and residual processes, it was flattened into a 512-dimensional feature vector by a global average pooling layer. Then, this feature vector passed through a fully connected layer to return the output result.
[0146] In this experiment, 10,000 images were randomly selected from a dataset of 60,000 labeled images and allocated to the central server for testing the accuracy of the global model, and the remaining 50,000 images were allocated to 100 clients.
[0147] To simulate data heterogeneity in this experiment, a two-class data partitioning strategy of only allocating data samples with two class labels to each client and a Dirichlet data partitioning strategy with a concentration parameter α of 0.1 were adopted. The two-class data partitioning strategy would allocate 250 data samples corresponding to two randomly selected class labels to each client; while the Dirichlet data partitioning strategy would randomly allocate different numbers of data samples with different label classes to each client. Here, α is a parameter related to the degree of heterogeneity. The smaller the value of α, the greater the degree of data heterogeneity.
[0148] In the main experiment, unless otherwise specified, the local iteration number was generally set to 5, the batch size was 50, there were 100 clients in total, and 400 communication rounds were run at a sampling rate of 0.1. During local training, the standard configuration in the FL benchmark was followed, and SGD with a learning rate of 0.01 and a momentum of 0.9 was used as the local optimizer. At the same time, the window size of the introduced historical global model was 5, and the starting round was set to three-quarters of the total number of communication rounds, that is, 300. All algorithms used the top-1 accuracy as the metric to evaluate the algorithms.
[0149] Comparison experiment settings:
[0150] · Method of only optimizing local training on the client (control)
[0151] Throughout the communication rounds, the window size of the historical global model is set to 0. When the server aggregates, it weights and aggregates the client models according to the relative size of the local datasets. In terms of the model structure, the last fully connected layer of ResNet18 is replaced by a new fully connected layer with 512 input nodes (corresponding to the dimension of the feature map of the last layer of ResNet18) and 128 output nodes. Then, a classifier is defined, which is a fully connected layer that takes the 128-dimensional feature vector from the previous layer as input and outputs the number of classes required for the CIFAR-10 dataset, which is 10.
[0152] · Method for improving the model aggregation process only on the server side (control)
[0153] Throughout the communication rounds, the window size of the historical global model is set to 0. The clients perform local training using the classical Stochastic Gradient Descent (SGD). In terms of the model structure, the last fully connected layer of ResNet18 is replaced by a new fully connected layer with 512 input nodes (corresponding to the dimension of the feature map of the last layer of ResNet18) and 128 output nodes. Then, a classifier is defined, which is a fully connected layer that takes the 128-dimensional feature vector from the previous layer as input and outputs the number of classes required for the CIFAR-10 dataset, which is 10.
[0154] · Data Heterogeneity Adaptive Federated Learning Method Based on Invariant Representations
[0155] The window size of the historical global model is set to 5, and the starting round is set to 300. The server adopts a data heterogeneity adaptive aggregation strategy to aggregate the client models. When the number of communication rounds reaches the starting round, the historical global model is introduced to update the current global model. In terms of the model structure, the last fully connected layer of ResNet18 is replaced by a fully connected layer with 512 as the input and 1024 as the output. Then, a classifier is defined, which is a fully connected layer that takes the 512-dimensional feature vector of the representation as input and outputs the number of classes required for the CIFAR-10 dataset, which is 10.
[0156] Conclusion analysis:
[0157] The experimental results are as Figure 3 shown. Since the Dirichlet data partitioning strategy with a concentration parameter α of 0.1 is used to simulate data heterogeneity, it causes instability in the training process, resulting in obvious oscillations in the curve. Figure 3 In it, the abscissa is the number of communication rounds of training, and the ordinate is the accuracy of the global model on the test set at the corresponding round.
[0158] From Figure 3It can be seen that the method provided by the present invention is superior to the other two types of methods both in terms of accuracy and convergence speed. Since the local training adopts the client invariant representation learning method, which reduces the difference between the global and local data in the feature space, the accuracy rapidly increases in the early stage of the training phase. In the later stage of the training phase, due to the injection of the historical global model, the degree of the non-IID phenomenon caused by the missing data samples can be reduced, preventing the global model update from deviating from the global optimal direction. At the same time, it also provides a better initialization for the client model to perform the next round of local training, thus accelerating the convergence speed of the model. Therefore, at the start of the 300th round in the accuracy simulation diagram, i.e., Figure 3 the model converges rapidly. Although introducing the historical global model can accelerate convergence, introducing it too early is not a good thing because the early global model contains too much information noise, which will instead reduce the accuracy of the model.
[0159] The method that only optimizes local training on the client side often indirectly aligns local features with global features by sharing local features and aggregating them into global features, thus alleviating the feature drift phenomenon. However, the local features are still easily affected by overfitting of heterogeneous data distributions, resulting in biases in the aggregated global features. Therefore, the accuracy is lower than that of the method provided by the present invention throughout the training phase. In addition, the method that only improves the model aggregation process on the server side cannot effectively alleviate the significant impact of data heterogeneity on the model performance because it cannot make corresponding adjustments according to specific client scenarios. As long as the heterogeneity is large, it is difficult to significantly improve the overall performance of the model by simply improving the global aggregation process. Therefore, it has the worst effect in the simulation experiment.
[0160] The embodiment of the present invention also provides a computer-readable storage medium storing a computer program, which when executed by a processor, causes the processor to execute the steps of the data heterogeneous adaptive federated learning method as described above.
[0161] See Figure 4 , the embodiment of the present invention also provides a computer device including a memory and a processor, where the memory stores a computer program, and when the computer program is executed by the processor, it causes the processor to execute the steps of the data heterogeneous adaptive federated learning method as described above.
[0162] The embodiment of the present invention also provides a computer program product containing instructions, which when running on a computer, causes the computer to execute the steps of the data heterogeneous adaptive federated learning method as described above.
[0163] It is understandable that the systems, devices, and storage media provided in the embodiments of the present invention correspond to the methods provided in the embodiments of the present invention. For the explanations, examples, and beneficial effects of related content, reference can be made to the corresponding parts in the above data heterogeneous adaptive federated learning method.
[0164] It should be noted that those of ordinary skill in the art can understand that all or part of the steps implemented in the embodiments of the present invention can be fully or partially implemented by software, hardware, firmware, or any combination thereof. When implemented using hardware, it can be fully or partially implemented in the form of purchasing standard parts or modified parts. When implemented using software, it can be fully or partially implemented in the form of a computer program product. The computer program product includes one or more computer instructions. When the computer program instructions are loaded and executed on a computer, all or part of the processes or functions described in the embodiments of this application are generated. The computer can be a general-purpose computer, a dedicated computer, a computer network, or other programmable devices. The computer instructions can be stored in a computer-readable storage medium, or transmitted from one computer-readable storage medium to another computer-readable storage medium. For example, the computer instructions can be transmitted from one website, computer, server, or data center to another website, computer, server, or data center via wired (such as coaxial cable, optical fiber, digital subscriber line (DSL)) or wireless (such as infrared, wireless, microwave, etc.) means. The computer-readable storage medium can be any available medium that a computer can access or a data storage device such as a server, data center, etc. that contains one or more integrated available media. The available medium can be a magnetic medium (such as a floppy disk, hard disk, magnetic tape), an optical medium (such as a DVD), or a semiconductor medium (such as a solid-state drive Solid State Disk (SSD)), etc.
[0165] In summary, the present invention provides a data heterogeneous adaptive federated learning method based on invariant representations to alleviate the problems of poor model accuracy and slow convergence speed caused by label shift and feature shift in the data heterogeneous scenario. Specifically, before the start of the first round of training, all clients capture the local data heterogeneity. During training, the clients participating in this round of training adopt the training method of invariant representations to prompt the client models to extract client invariant representations that only contain the required non-redundant feature information. At the end of training, the client models are uploaded to the server. The server uses the data heterogeneous quantization metrics uploaded by each client to determine the specific aggregation weights, aggregates these client models, and injects the historical global model into the current global model in the later stage of the training phase for model update. In this way, the overall training effect is improved.
[0166] It should be understood that the examples and embodiments described herein are for illustrative purposes only and are not intended to limit the present invention. Those skilled in the art can make various modifications or changes based on it. Any modification, equivalent replacement, improvement, etc. made within the spirit and principle of the present invention shall be included within the protection scope of the present invention.
Claims
1. A data heterogeneous adaptive federated learning method, characterized in that It includes the following steps: S1. The server initializes the global model parameters, broadcasts the global model and the assumed global label distribution to all participating clients for training. The clients calculate the corresponding local label distribution based on their own local datasets, and calculate the difference between the global label distribution and the local label distribution as a quantization index of data heterogeneity; S2. After receiving the global model parameters, the participating clients perform client model training locally and simultaneously train an encoder capable of extracting client-invariant representations; S3. The server jointly determines the aggregation weights based on the quantization index of data heterogeneity and the size of the local dataset to adaptively aggregate the client models according to the size of the data heterogeneity; S4. If the client model has reached the later stage of the specified training phase, inject the historical global model and perform a weighted average with the global model aggregated in the current round as the global model updated in this round. If the client model has not reached the later stage of the specified training phase, use the aggregated client model as the global model updated in the current round; S5. The server broadcasts the updated global model to each client participating in the next round of training and enters the next round of local training until the predetermined number of communication rounds is reached to complete the global training.
2. The data heterogeneous adaptive federated learning method according to claim 1, wherein The S1 further includes: S11. The global label distribution and the local label distribution are respectively denoted as Dis global and Dis local ; Global label distribution Dis global Treat all C label categories equally. For each category c, denote as: The local label distribution for the k-th client is defined as , denoted as: Among them, the c-th element represents the proportion of the data volume of the c-th category, B k is the dataset of the k-th client, x i and y i are the input of the i-th sample and the corresponding label, respectively; S12, the local label distribution Dis evaluated by each client local and the global label distribution Dis global issued by the server k The difference d k is defined as: Thereby capturing the data label heterogeneity of the client locally and using it as a quantization index of data heterogeneity; where, discrepancy(.) is a predefined metric function, and the KL divergence is used as the metric function, then the difference d k is defined as: As a supplementary index for the aggregation weight and upload it to the server together with the client model trained in this round.
3. The data heterogeneous adaptive federated learning method according to claim 1, wherein, In the S2, the client model training process is iterated multiple times, and the client model parameters are updated in each iteration to prompt the client model to extract client-invariant representations that only contain the required non-redundant feature information.
4. The data heterogeneous adaptive federated learning method according to claim 3, wherein In S2, in each communication round t, where t = 1, 2,..., T, a client subset K from the set of all clients K t is active and all download the global model from the server to participate in the training of the current round; At each local step of the t-th communication round Client k randomly samples a batch of training data from the corresponding local dataset B k and selects the CIRL algorithm based on federated learning to learn the invariant representation of the client, reducing the difference between the global dataset and the local dataset in the feature space to reduce the difference between the global model and the local client model and achieving lower losses in local training. Thus, the objective function f is designed and the gradient is used to update the local client model . The updated local client model is denoted as : where η is the learning rate; After E steps of local training, each client sends its respective trained local client model to the server.
5. The data heterogeneous adaptive federated learning method according to claim 4, wherein In the S2, the CIRL algorithm based on federated learning includes the following steps: S2-a. Each client maps the input x of the sample to a representation z through an encoder network, and the distribution of the representation z is denoted as p(z|x). A classifier is learned from the representation z, and given the representation z, the prediction distribution parameterized by w prediction ; S2-b. Taking L as the loss function, the supervised loss function on client k is defined as: where y is the label corresponding to the input x; S2-c. To learn a representation z with a consistent conditional distribution p k (z|x) across all clients, the conditional mutual information I k (x,z|y) between the data input with given labels and the representation is used to limit the amount of information that the representation can contain, which is defined as follows: Minimize to enforce that the representation z contains only the non-redundant feature information required to predict the label y, and no other label-independent information about x; where r(z|y i ) is the conditional distribution of z given y. Taking r(z|y i ) as the reference distribution to replace the intractable p(z|y), for the classification task, r(z|y i ) is set to be a Gaussian distribution with mean μ y and variance as defined below: Among them, μ y and σ y are parameters to be optimized, y = 1, 2, ..., C, where C is the number of categories; S2-d. Obtain the local objective function of each client k: where α CMI is the hyperparameter of the conditional mutual information regularization term; Each client performs gradient descent in multiple rounds of iterative training for the above local objective function and continuously updates the local client model to gradually approach the optimal solution.
6. The data heterogeneous adaptive federated learning method according to claim 1, wherein In the S3 - S4, while collecting the client models, the server also collects the quantization indexes of data heterogeneity of each client as supplementary indexes for the aggregation weights. When the communication round reaches the startup round t-startup, introduce the historical global model on the basis of the current global model for update.
7. The data heterogeneous adaptive federated learning method according to claim 6, wherein The S3 further includes: S31. The server uses the relative size n of the local dataset k and the quantization metric d of the local data heterogeneity k , and determines a more discriminative aggregation weight p for each client model k : Among them, ReLU(.) is the relu function used to process negative values, a is the hyperparameter used to balance n k and d k and b is another hyperparameter used to adjust the weight; S32. Aggregate the client models, and update the global model in this round to be: Among them, K t is all the clients participating in the training in the t-th round and is a subset of the client set K.
8. The data heterogeneous adaptive federated learning method according to claim 6, wherein The S4 further includes: S41. If the current communication round t is greater than or equal to the set startup round t-startup, introduce the historical global models of the previous S rounds including the current global model and perform a weighted average with the current global model to update the global model in this round to be: where S is the window size of the historical global model, indicating the number of historical global models to be injected; S42. If the current communication round t is less than or equal to the set startup round t-startup, the historical global model is not introduced, and the update of the global model in this round is the model obtained after aggregating the client models: 。 9. A computer-readable storage medium, characterized in that A computer program is stored, and when the computer program is executed by a processor, the processor is caused to execute the steps of the data heterogeneous adaptive federated learning method according to any one of claims 1-8.
10. A computer device, characterized in that, It includes a memory and a processor. The memory stores a computer program, and when the computer program is executed by the processor, the processor is caused to execute the steps of the data heterogeneous adaptive federated learning method according to any one of claims 1-8.
Citation Information
Patent Citations
Power terminal multi-task federal learning method for power internet of things
CN115049522A
Federal positioning method based on adaptive aggregation and feature alignment
CN119893434A