Data heterogeneous adaptive federal learning method, medium and equipment

By using data heterogeneity quantization indicators in federated learning for adaptive model aggregation and injecting historical global models later in the training period, the problems of low model accuracy and slow convergence speed caused by data heterogeneity are solved, and more efficient training performance is achieved.

CN120124779AActive Publication Date: 2025-06-10NANJING UNIV OF POSTS & TELECOMM

Patent Information

Application Number
CN202510600778.4
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2025-05-12
Publication Date
2025-06-10
Estimated Expiration
2045-05-12

AI Technical Summary

Technical Problem

The label offset and feature offset problems caused by data heterogeneity lead to low accuracy and slow convergence speed of federated learning models.

Method used

By initializing the global model and label distribution on the server side, the client calculates the difference in local label distribution as a quantitative indicator, the server determines the aggregation weight based on the quantitative indicator and data set size, performs adaptive model aggregation, and injects historical global models later in the training to accelerate convergence.

Benefits of technology

It effectively alleviates the problems of low model accuracy and slow convergence speed caused by data heterogeneity, and improves the overall training performance.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN120124779A_ABST
    Figure CN120124779A_ABST
Patent Text Reader

Abstract

The invention discloses a data isomerism adaptive federal learning method, a medium and equipment, and the method comprises the steps: S1, a server initializes a global model, and a client calculates a quantitative index of data isomerism; s2, locally executing client model training by the client, and training an encoder capable of extracting the invariant representation of the client at the same time; s3, the server performs adaptive aggregation on the client model according to the size of the data isomerism by combining the quantitative index of the data isomerism and the size of the local data set; s4, injecting a historical global model, and adding and averaging the historical global model and the global model obtained by aggregation of the current round to serve as an updated global model of the current round; and S5, completing global training until a preset number of communication rounds is reached. According to the method, the problems of poor model precision and low convergence speed caused by label offset and feature offset in a data heterogeneous scene can be relieved at the same time, and the overall training effect is improved.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention relates to the technical field of federated learning, and in particular to a 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 high-quality data with distribution representativeness. It has become a general trend for data to flow freely on the premise of security and compliance. Facing the huge potential value 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 being unwilling to contribute value, it has led to the existing problems of a large number of data islands and privacy protection. 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. 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 the input features varies between different clients, the feature shift phenomenon 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] A data heterogeneous adaptive federated learning method, medium and device provided by the present invention 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: A data heterogeneous adaptive federated learning method, comprising 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 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; S2. After receiving the global model parameters, the participating clients perform client model training locally, and at the same time 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, so as to adaptively aggregate the client models according to the size of data heterogeneity; 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 obtained by aggregation 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.

[0007] Further, S1 further includes: S11. The global label distribution and the local label distribution are respectively denoted as Dis global and Dis local ; The global label distribution Dis global treats all C label categories equally. For each category c, it is denoted as:

[0008] The local label distribution of the k-th client is defined as , denoted as:

[0009] where 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 respectively the input and the corresponding label of the i-th sample; S12. The difference d local between the local label distribution Dis global evaluated by each client and the global label distribution Dis k issued by the server, and the difference dk It is defined as:

[0010] In this way, the data label heterogeneity of the client side is captured locally and used as a quantitative index for data heterogeneity; Among them, discrepancy(.) is a predefined metric function, and the KL divergence is used as the metric function, then the difference d k It is defined as:

[0011] As a supplementary index for the aggregation weight, it is uploaded to the server together with the client model trained in this round.

[0012] Furthermore, in S2, the training process of the client model is iterated multiple times, and the client model parameters are updated in each iteration, prompting the client model to extract client-invariant representations that only contain the required non-redundant feature information.

[0013] Furthermore, in S2, in each communication round t, where t = 1, 2,..., T, the client subset K from all client sets K t is active and all download the global model from the server to participate in the training of the current round; 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:

[0014]

[0014] where η is the learning rate; After E steps of local training, each client sends its own trained local client model to the server.

[0015] Furthermore, in S2, the CIRL algorithm based on federated learning includes the following steps: S2-a. Each client maps the input x of the sample to the representation z through the encoder network, and the distribution of the representation z is denoted as p(z|x), and this representation distribution is obtained by the expectation and standard deviation Parameterized. Learn a classifier from the representation z, which, given the representation z, uses a prediction distribution parameterized by w to predict (where the parameter w is omitted); S2-b. Taking L as the loss function, the supervised loss function on client k is defined as:

[0016] 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 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:

[0017] Minimize to force the representation z to contain only the non-redundant feature information required to predict the label y, and not other information about x that is irrelevant to the label; 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 be a Gaussian distribution with as the mean and as the variance, defined as follows:

[0018] where μ y and σ y are parameters to be optimized, y = 1, 2,..., C, and C is the number of classes; S2-d. Obtain the local objective function for each client k:

[0019] 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, continuously updating the local client model to gradually approach the optimal solution.

[0020] Further, 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 updating.

[0021] Further, S3 further includes: S31. The server uses the relative size n k of the local dataset k and the quantization metric d k of the local data heterogeneity

[0022] to determine a more discriminative aggregation weight p k for each client model k : 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 weight; S32. Aggregate the client models and update to obtain the global model of this round as:

[0023] 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.

[0024] Further, S4 further includes: S41. If the current communication round number 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:

[0025] 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 number 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: .

[0026] 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 - mentioned data - heterogeneity - adaptive federated learning method.

[0027] 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 performs the steps of the above-mentioned data heterogeneous adaptive federated learning method.

[0028] The beneficial effects of the present invention are as follows: 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 third 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, thereby effectively improving the performance of the overall training. Description of the Drawings

[0029] The drawings described herein are used to provide a further understanding of the present application and constitute a part of the present application. The schematic 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.

[0030] Figure 1 It is a schematic diagram of the overall process of the data heterogeneous adaptive federated learning method in an embodiment of the present invention.

[0031] Figure 2 It is a schematic diagram of the specific process of the data heterogeneous adaptive federated learning method in an embodiment of the present invention.

[0032] Figure 3 It is a simulation diagram of the model accuracy rate of the comparative experiment in an embodiment of the present invention in the test set.

[0033] Figure 4 It is a structural block diagram of the computer device in an embodiment of the present invention. Detailed Embodiments

[0034] 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 embodiments. Without conflict, the embodiments in the present application and the features in the embodiments can be combined with each other. Based on the embodiments of the present invention, all other embodiments obtained by those of ordinary skill in the art without creative efforts shall fall within the protection scope of the present invention.

[0035] 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 various embodiments can be combined with each other, but it must be based on the ability of those of ordinary skill in the art to implement. 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.

[0036] See Figure 1 - Figure 2 , the embodiments of the present invention provide a data heterogeneous adaptive federated learning method, including 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 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. S2. After receiving the global model parameters, the participating clients perform local client model training and simultaneously train an encoder capable of extracting client-invariant features. 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 data heterogeneity. 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. 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.

[0037] 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 among all classes and the generalization ability of the global model. Therefore, it satisfies 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; Specifically, S1 further includes: S11. The global label distribution and the local label distribution are respectively denoted as Dis global and Dis local ; The global label distribution Dis global treats all C label classes equally. For each class c, it is denoted as:

[0038] The local label distribution of the k-th client is defined as , and is denoted as:

[0039] 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; S12. The difference d local between the local label distribution Dis global evaluated by each client and the global label distribution Dis k issued by the server, and the difference d k is defined as:

[0040] This is used to capture the data label heterogeneity of the client local and serves as a quantitative index for 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:

[0041] As a supplementary index for the aggregation weight, it is uploaded to the server together with the client model trained in this round; In the present invention, when selecting the metric function for discrepancy(.), functions that can measure vector distances such as the L1 distance, L2 distance, KL divergence, etc. can be selected.

[0042] In this embodiment, in S2, the client model training process can be iterated multiple times. Each iteration updates the client model parameters, prompting the client model to extract client invariant representations that only contain the required non-redundant feature information.

[0043] In this embodiment, in S2, in each communication round t, where t = 1, 2,..., T, a subset of clients K t from the set of all clients K is in an active state and all download the global model from the server to participate in the training of the current round; 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 , and the updated local client model is denoted as:

[0044] where η is the learning rate; After E steps of local training, each client sends its trained local client model to the server.

[0045] 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 solved. Among them, clients have different input distributions (domains) for a given label class y c , that is:

[0046] where , and .

[0047] Due to the 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 between clients change, traditional federated learning methods (such as FedAvg) will try to learn all possible feature knowledge during local training, including client invariant knowledge and client-specific knowledge. However, client-specific knowledge will cause the representation to diverge, that is:​

[0048] This leads to a decrease in accuracy, where z is the intermediate feature trained by the encoder network and is called the representation.

[0049] Therefore, the present invention adopts the CIRL algorithm based on federated learning to learn the invariant representation of the client, reduce the difference between the global and local data sets in the feature space, thereby further reducing the difference between the global and local client models, and achieving a lower loss in local training.

[0050] Specifically, in S2, the CIRL algorithm based on federated learning includes the following steps: S2-a. Each client maps the input x of the sample to the representation z through the encoder network, and the distribution of the representation z is denoted as p(z|x), and this representation distribution is parameterized by the expectation and the standard deviation A classifier is learned from the representation z, and this classifier uses the prediction distribution parameterized by w to predict (where the parameter w is omitted); S2-b. Taking L as the loss function, the supervised loss function on client k is defined as:

[0051] where y is the label corresponding to the input x; S2-c. In order to learn a representation z with a consistent conditional distribution p k (z|x) on each client, 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, and is defined as follows:

[0052] Minimize to force the representation z to contain only the non-redundant feature information required to predict the label y, and not contain other information about x that is irrelevant to the label; where r(z|y i ) is the conditional distribution of z given y and serves 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, and is defined as follows:

[0053] where μ y and σ yis the parameter to be optimized, y = 1, 2, ..., C, where C is the number of categories; S2-d. Obtain the local objective function for each client k:

[0054] 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, continuously updating the local client model to gradually approach the optimal solution.

[0055] Furthermore, in S2-c, for client k, the conditional mutual information between the data input x with given label y and the representation z can be defined as:

[0056] where p k (z|y) is the conditional distribution of the representation z corresponding to label y on client k. This term is difficult to integrate with respect to x. If I k (x, z|y) is directly introduced as the conditional mutual information term into the objective function, the presence of p k (z|y) will make the entire mutual information term difficult to handle during iterative training using the gradient descent algorithm. Therefore, an upper bound is derived to minimize this conditional mutual information term:

[0057] Here, the upper bound can be calculated and used as a regularizer when training the representation network to optimize 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:

[0058] .

[0059] 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 based on the current global model for update.

[0060] In this embodiment, S3 further includes: S31. The server uses the relative size n of the local dataset k and the quantization index d of local data heterogeneity k to determine a more discriminative aggregation weight p for each client model k :

[0061] where ReLU(.) is the relu function used to handle negative values, and a is a 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 of this round to:

[0062] where K t is all the clients participating in training in the t-th round, which is a subset of the client set K.

[0063] In this embodiment, the S4 further includes: 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 first S rounds including the current global model, and perform weighted averaging with the current global model, and update the global model of this round to:

[0064] 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 number t is less than or equal to the set startup round number t-startup, 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: .

[0065] To verify the application of the present invention, a set of specific experiments will be provided below to further illustrate this data heterogeneity adaptive federated learning method: Experimental environment: 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.

[0066] 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. In the second stage, there were also two residual blocks, but the output channels increased to 128. In the third stage, the network was further deepened, containing four residual blocks, each with 256 output channels. The last stage also contained two residual blocks, but the output channels were 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.

[0067] In this experiment, 10,000 data samples with labels were randomly selected from 60,000 labeled datasets and allocated to the central server for testing the accuracy of the global model, and the remaining 50,000 were allocated to 100 clients.

[0068] In this experiment, to simulate data heterogeneity, 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, the greater the degree of data heterogeneity.

[0069] 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-fourths of the total communication rounds, that is, 300. All algorithms used the top-1 accuracy as the metric to evaluate the algorithms.

[0070] Comparison experiment settings: · Method of only optimizing local training on the client (control) Throughout the communication rounds, the window size of the historical global model is set to 0, and when the server performs aggregation, 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 with a new fully connected layer with 512 input nodes (corresponding to the dimension of the feature map of the last layer of ResNet18), and the number of output nodes is set to 128; then a classifier is defined, which is a fully connected layer that receives the 128-dimensional feature vector from the previous layer as input and outputs the number of classes 10 required by the CIFAR-10 dataset.

[0071] · Method for improving the model aggregation process only on the server side (control) Throughout the communication rounds, the window size of the historical global model is set to 0, and the clients use the classical Stochastic Gradient Descent (SGD) for local training. In terms of the model structure, the last fully connected layer of ResNet18 is replaced with a new fully connected layer with 512 input nodes (corresponding to the dimension of the feature map of the last layer of ResNet18), and the number of output nodes is set to 128; then a classifier is defined, which is a fully connected layer that receives the 128-dimensional feature vector from the previous layer as input and outputs the number of classes 10 required by the CIFAR-10 dataset.

[0072] · Data Heterogeneity Adaptive Federated Learning Method Based on Invariant Representations 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 with a fully connected layer with 512 as input and 1024 as 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 10 required by the CIFAR-10 dataset.

[0073] Conclusion Analysis: 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 will cause the phenomenon of unstable training process, resulting in obvious oscillations in the curve. Figure 3 In, 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.

[0074] 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.

[0075] 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 offset phenomenon. However, the local features are still prone to being overfitted by the heterogeneous data distribution, resulting in biases in the aggregated global features. Therefore, the accuracy is lower than 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 model performance by simply improving the global aggregation process. Therefore, it has the worst effect in the simulation experiment.

[0076] 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.

[0077] 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.

[0078] The embodiment of the present invention also provides a computer program product containing instructions, which when run on a computer, causes the computer to execute the steps of the data heterogeneous adaptive federated learning method as described above.

[0079] 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.

[0080] 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 realized by software, hardware, firmware, or any combination thereof. When implemented using hardware, it can be fully or partially realized in the form of purchasing standard parts or modified parts. When implemented using software, it can be fully or partially realized 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 the present application are generated. The computer can be a general-purpose computer, a special-purpose 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. For example, the computer instructions can be transmitted from one website, computer, server, or data center to another website, computer, server, or data center by wire (such as coaxial cable, optical fiber, digital subscriber line (DSL)) or wireless (such as infrared, wireless, microwave, etc.). The computer-readable storage medium can be any available medium that a computer can access or a data storage device such as a server or data center that includes 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)).

[0081] 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 data heterogeneous scenarios. 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 heterogeneity 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.

[0082] 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: The following steps are involved: 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, and the clients calculate the corresponding local label distribution based on their own local data sets, and calculates the difference between the global label distribution and the local label distribution as a quantitative indicator of data heterogeneity. S2. After receiving the global model parameters, the client participating in the training performs client model training locally and trains an encoder that can extract the client-invariant representation; S3, the server combines the quantitative index of data heterogeneity and the size of the local data set to jointly determine the aggregation weight to adaptively aggregate the client model according to the size of data heterogeneity; S4. If the client model has reached the late stage of the specified training phase, the historical global model is injected and averaged with the global model obtained by the current round of aggregation to serve as the updated global model of this round. If the client model has not reached the late stage of the specified training phase, the aggregated client model is used as the updated global model of the current round. S5. The server broadcasts the updated global model to each client that wants to participate 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.

2. The data heterogeneous adaptive federated learning method according to claim 1, characterized in that: Said S1 further comprises: S11, global label distribution and local label distribution are represented as Dis global and Dis local ; Global label distribution Dis global Treat all C label categories equally, for each category c, record it as: The local label distribution for the kth client is defined as , recorded as: Among them, the cth element represents the proportion of data in the cth category, B k is the dataset of the kth client, x i and i are the input and corresponding label of the i-th sample respectively; S12, local label distribution Dis evaluated by each client local The global label distribution sent by the server global The difference between k , difference d k Defined as: This captures the local data label heterogeneity of the client and serves as a quantitative indicator of data heterogeneity; Where discrepancy (.) is a predefined metric function, and KL divergence is used as the metric function, then the difference d k Defined as: As a supplementary indicator of the aggregation weight, it is uploaded to the server together with the client model trained in this round.

3. The data heterogeneous adaptive federated learning method according to claim 1, characterized in that: In S2, the client model training process is iterated for multiple rounds, and each round of iteration updates the client model parameters, prompting the client model to extract the client invariant representation that only contains the required non-redundant feature information.

4. The data heterogeneous adaptive federated learning method according to claim 3, characterized in that: In S2, in each communication round t, t=1, 2, ..., T, a subset of clients K from the set of all clients K t In active state, they download the global model from the server to participate in the current round of training; At each local step of the tth communication round In the example, client k is in the corresponding local dataset B k Randomly extract a batch of training data from , the CIRL algorithm based on federated learning is selected to learn the invariant representation of the client, narrow the difference between the global dataset and the local dataset in the feature space, reduce the difference between the global model and the local client model, and achieve lower loss in local training. The objective function f is designed and the gradient is used To update the local client model , the updated local client model is recorded as: Where η is the learning rate; After E steps of local training, each client The trained local client models Send to server.

5. The data heterogeneous adaptive federated learning method according to claim 4, characterized in that: In S2, the CIRL algorithm based on federated learning includes the following steps: S2-a, each client maps the sample input x to the representation z through the encoder network, and the distribution of the representation z is recorded as p(z|x), and learns a classifier from the representation z. The classifier uses the prediction distribution parameterized by w given the representation z. predict ; S2-b, taking L as the loss function, the supervision loss function on client k is defined as: Among them, y is the label corresponding to the input x; S2-c, in order to learn a conditional distribution p that is consistent across all clients k The representation z of (z|x) is obtained by using the conditional mutual information I between the input data and the representation given the label k (x,z|y) to limit the amount of information that the representation can contain, defined as follows: minimize To force the representation z to contain only the non-redundant feature information needed to predict the label y, without containing other information about x that is not related to the label; Among them, r (z|y i ) is the conditional distribution of z given y, and r(z|y i ) as the reference distribution, replacing the difficult-to-solve p(z|y), and for classification tasks, r(z|y i ) is set to μ y is the mean, is a Gaussian distribution with variance defined as follows: Among them, μ y and σ y is the parameter to be optimized, y=1,2,...,C, C is the number of categories; S2-d, get the local objective function of each client k: Among them, α 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, continuously updating the local client model to gradually approach the optimal solution.

6. The data heterogeneous adaptive federated learning method according to claim 1, characterized in that: In S3-S4, while collecting the client models, the server also collects quantitative indicators of data heterogeneity of each client as a supplementary indicator of the aggregation weight. 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 updating.

7. The data heterogeneous adaptive federated learning method according to claim 6, characterized in that: The S3 further comprises: S31, the server uses the relative size of the local data set n k and a quantitative indicator of local data heterogeneity d k , determine a more discriminative aggregation weight p for each client model k : Among them, ReLU (.) is the relu function used to handle negative values, and a is used to balance n k and d k is a hyperparameter, and b is another hyperparameter used to adjust the weight; S32, aggregate the client models, and update the global model of this round to: Among them, K t It is all the clients participating in the training in the tth round, which is a subset of the client set K.

8. The data heterogeneous adaptive federated learning method according to claim 6, characterized in that: The S4 further comprises: S41. If the current communication round number t is greater than or equal to the set startup round number t-startup, the historical global model of the previous S rounds including the current global model is introduced, and the historical global model is summed and averaged with the current global model, and the global model of this round is updated as follows: Among them, 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 number t is less than or equal to the set startup round number t-startup, the historical global model is not introduced, and the update of the global model of this round is the model obtained after aggregating the client model: 。 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 executes the steps of the data heterogeneous adaptive federated learning method as described in any one of claims 1 to 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 executes the steps of the data heterogeneous adaptive federated learning method as described in any one of claims 1 to 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

  • Disentangled personalized federated learning method via consensus representation extraction and diversity propagation

    US20240320513A1

Cited By

  • Generalized time-frequency positioning method for distributed training scene

    CN121692206A