Method and system for solving data heterogeneity problem in federated learning based on empty distillation and large class suppression
By adopting empty distillation and large-class suppression techniques in federated learning, the model optimization offset caused by inconsistent data distribution among users is solved, the performance of the global model is significantly improved, and the classification accuracy of empty class data is enhanced without increasing communication and computing overhead and the ability to reduce misclassification of small class samples is enhanced.
Patent Information
- Application Number
- CN202510008150.5
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-01-03
- Publication Date
- 2025-05-06
- Estimated Expiration
- Not applicable · inactive patent
AI Technical Summary
Inconsistent data distribution among users in federated learning leads to model optimization offset. The existing methods have limitations in extreme data heterogeneity scenarios, limiting the convergence speed of local models or failing to fully consider the existence of empty categories.
Using a method based on empty class distillation and large class suppression, we use the empty class distillation loss function to constrain the output of the local model to ensure that the model maintains the classification ability of empty classes during the update process. In addition, by counting the category distribution probability of the local training data set and introducing a large-class suppression loss function, the large-class logit value in the small-class sample output is punished, and the risk of small-class samples being misclassified into large-class categories is reduced.
It effectively alleviates the model optimization offset caused by inconsistent data distribution, significantly improves the performance of the global model, enhances the accuracy of the classification of empty-class data by the user model, and reduces the risk of misclassification of small-class samples into large-class categories, and does not introduce additional communication costs and calculation overhead.
Smart Images

Figure CN119939461A_ABST
Abstract
Description
Technical Field
[0001] The present invention belongs to the field of federated learning, and specifically relates to a method and system for solving the data heterogeneity problem in federated learning based on empty class distillation and large class suppression. Background Art
[0002] Federated learning is a joint training paradigm designed to protect user privacy. Its core idea is to complete model training by sharing information such as model parameters between users and central servers without collecting user training data.
[0003] The typical process of federated learning usually includes the following four steps. In the first step, the server distributes the aggregated global model to all users (the randomly initialized global model is distributed in the initial stage), and the user obtains the global model and initializes it as a local model; in the second step, the user trains the local model based on local data; in the third step, the user uploads the trained local model to the server; in the fourth step, the server performs weighted aggregation of the local model according to the proportion of the training data of each user to generate a new global model. The above four steps constitute a training cycle. Usually, federated learning will perform multiple training cycles until the model converges or stops training after reaching the predetermined number of training rounds.
[0004] In practical applications, there is often a phenomenon of inconsistent data distribution among users participating in federated training. For example, in federated learning in the medical field, due to the different patient compositions of different hospitals, the amount of data for specific disease categories in each hospital is different, and the sample characteristics of the same disease category are also different. This difference causes deviations in the convergence direction and speed of each user's local model, which in turn affects the convergence effect and final performance of the global model. At present, the mainstream methods mainly use two ideas to alleviate this problem: one is to use the global model to constrain the update direction of the local model to reduce the deviation between the local model and the global model; the other is to achieve more balanced performance on the local model through the category balancing technology in long-tail learning, thereby accelerating the convergence of the global model. For example, there is a method that incorporates the difference between the parameters of the local model and the global model as a supplement into the loss function to constrain the model update direction during the user's local training process. In addition, there is a method that weights the logit output of the model according to the distribution probability of the user data, trying to alleviate the offset problem of the local model.
[0005] However, in extreme data heterogeneity scenarios, the above methods all have certain limitations. The aforementioned method of using the global model to constrain the update direction of the local model usually limits the convergence speed of the local model, thereby increasing the computational and communication costs; while the method based on the data distribution probability weighted model output logit fails to fully consider the existence of empty categories in user data, and the method of relying solely on the data distribution probability for weighting is still difficult to fully deal with the problem of category imbalance.
[0006] Therefore, how to effectively solve the problem of data heterogeneity in federated learning is of great significance for the application of federated learning in the real world. Summary of the invention
[0007] In order to solve the above technical problems, the present invention proposes a method and system for solving the data heterogeneity problem in federated learning based on empty class distillation and large class suppression, which effectively solves the impact of inconsistent data distribution among users on the global model.
[0008] To achieve the above object, the present invention adopts the following technical solution:
[0009] A method for solving the data heterogeneity problem in federated learning based on empty class distillation and large class suppression includes the following steps:
[0010] Step S100, each user counts the categories that do not contain training data locally, that is, empty categories;
[0011] Step S200, during the local training process, the logit value corresponding to the empty category in each sample output of the global model is used to constrain the logit value corresponding to the empty category in the local model output, ensuring that the local model can maintain the classification ability of the empty category during the update process, thereby alleviating the model update differences between users.
[0012] The categories that each user counts locally and does not include training data can be the same or different, but the sum of all users' training data covers all categories of training data. Even if a user's training data does not contain samples of certain categories, the data of this category will still be included in the training sets of other users.
[0013] The categories that do not contain training data locally counted by each user are initialized once at the beginning of training. This process will not affect subsequent training, nor will it consume additional computing resources, so it has no adverse effect on the operating efficiency of the system.
[0014] The global model is generated on the server side by aggregating the local model parameters of each user. After each aggregation, the server generates an updated global model and distributes it to each user. It is worth noting that during the local training process, the parameters of the global model remain fixed, and users only update the parameters of the local model.
[0015] The local model is initialized based on the global model and updated through several rounds of local training. After a fixed number of epochs of training, the user will upload the updated local model parameters to the server so that the server can aggregate the model parameters of all users and generate a new global model.
[0016] The logit of each sample output by the global model remains unchanged during the local training process, and the logit of the global model and the local model corresponding to the empty category have the same size and shape, which ensures that the constraints between the two can be implemented smoothly.
[0017] In order to ensure that the local model can maintain the classification ability of the empty category, the present invention provides a variety of methods to constrain the difference between the global model and the local model on the empty category logit. Common constraint methods include linear distance, Euclidean distance, and Kullback-Leibler divergence. The specific constraint method can be appropriately adjusted according to task requirements and experimental results to optimize the classification performance of the local model for the empty category.
[0018] Here we take Kullback-Leibler divergence as an example. The loss function corresponding to the above method is Has the following form:
[0019] ,
[0020] in, is expectation, is the training data set of the kth user, represents the set of all empty categories of the k-th user, the superscript g represents the global model, and y represents the label of sample x. is the local model input sample After The output probability of the empty category is as follows:
[0021] ,
[0022] in, represents the local model parameters, Represents the local model input sample After The output logit of the k-th category, t represents the category of the empty category set of the k-th user. and The calculation method is the same, except that the local model is replaced by the global model.
[0023] The constraint between the global model and the local model on the logit corresponding to the empty class in each sample output is defined as the empty class distillation loss function This loss function is designed to ensure that the local model can maintain the classification ability of the empty class during the local training process. The empty class distillation loss function will work together with the original local training cross entropy loss function to guide the local model to not only reduce the classification error during the optimization process, but also enhance the classification accuracy of the empty class.
[0024] In order to balance the impact of the empty class distillation loss function on the local training cross entropy loss function, a weight coefficient needs to be introduced during weighting . Weight coefficient It can be set according to the requirements of specific tasks, or adjusted through experimental experience.
[0025] Step S300, each user counts the category distribution probability of his local training data set;
[0026] Step S400, during the local training process, based on the category distribution probability of the local training data set obtained by the above statistics, the logit output of each training sample is constrained. This constraint is intended to effectively penalize the large-category logit value in the output logit of the small-category sample, thereby reducing the risk of the small-category sample being misclassified as a large-category sample, and ultimately achieving the purpose of alleviating the model drift between users and reducing the differences between different user models.
[0027] The category distribution probability of the local training data set counted for each user only includes the categories actually existing in the user's local data set, and does not consider empty categories.
[0028] The probability of the category distribution of the local training data set counted by each user is performed before the training begins, and the statistical process only needs to be performed once, and no repeated calculation is required during the training process, thereby avoiding additional computing overhead.
[0029] The specific definition of the probability of the category distribution of the user's local training data set is as follows:
[0030] ,
[0031] Among them, c represents the first categories, The local training dataset representing user statistics The probability of a category appearing, Represents the user's The sample size of training data for each category, Represents the total amount of training data for the user, specifically defined as ,in Represents the collection containing all categories.
[0032] The loss function that the user constrains the logit output of each sample is called the large-class suppression loss function , which is defined as follows:
[0033] ,
[0034] in, Represents the total number of samples output by the user model after a batch of samples is input. The sum of the logit values for the categories.
[0035] The said It is not enough to directly optimize the sum of the logit values of the categories, which will affect the convergence stability and the final accuracy. The sum of the logit values of the categories is modified and defined as follows:
[0036] ,
[0037] in, Representative samples The output does not correspond to the correct label. is the model input sample After that, output The logit of each category. are model parameters.
[0038] In the large-category suppression loss function, By optimizing the data distribution probability Weighted, this can effectively suppress the value of the large class logit in the small class sample output, and only slightly suppress the value of the small class logit in the large class sample output, which will not affect the learning of the large class and can alleviate the problem of small classes being misclassified as large classes.
[0039] On the other hand, the present invention provides a system for solving the data heterogeneity problem in federated learning based on empty class distillation and large class suppression, which includes:
[0040] A statistical unit, for each user, to count the categories of which no training samples exist in the local training data set, and record the set consisting of the categories of which no training samples exist as an empty class set;
[0041] The empty class distillation constraint unit is used to constrain the logit value of the empty class corresponding to each sample output of the global model and the local model based on the empty class distillation loss function during the local training process;
[0042] A distribution probability calculation unit is used to calculate the category distribution probability of each user's local training data set;
[0043] The large-category suppression constraint unit is used to directly constrain the logit value of each sample output by the local model based on the large-category suppression loss function based on the statistical category distribution probability of the local training data set during the local training process, so as to penalize the large-category logit value in the small-category sample output.
[0044] The beneficial effects of the present invention are:
[0045] The present invention effectively alleviates the model optimization deviation caused by inconsistent data distribution and significantly improves the performance of the global model. Through the empty class distillation technology, the classification accuracy of the optimized user model for empty category data is increased. The large category suppression method is adopted to reduce the risk of the user model misclassifying small category samples as large category categories. While solving the problem of inconsistent data distribution between users, no additional communication costs are introduced, and the newly added computing overhead is extremely small and can be ignored. BRIEF DESCRIPTION OF THE DRAWINGS
[0046] Figure 1 is the output probability confusion matrix of the user model on the test set in an embodiment of the present invention;
[0047] Figure 2 This is a network schematic diagram of a method for solving the data heterogeneity problem in federated learning based on empty class distillation and large class suppression in the present invention;
[0048] Figure 3 The experimental results of the present invention and the prior art on the test data set are shown;
[0049] Figure 4 This is a comparison chart of the experimental results of the global model of the present invention on the test data set after using different loss functions. DETAILED DESCRIPTION
[0050] The present invention will be further described below in conjunction with the accompanying drawings and embodiments.
[0051] It should be noted that the following detailed description is only an example, and is intended to further illustrate the embodiments of the present disclosure. Unless otherwise specified, all technical terms used herein should be interpreted according to the understanding of ordinary technicians in the field.
[0052] It should be noted that the terms used in the present disclosure are only used to describe specific embodiments and are not intended to limit the scope of implementation of the present disclosure. In this specification, unless the context clearly indicates otherwise, the singular form should also be understood to include the plural form. In addition, when "comprise" and / or "include" are used, their meaning is to indicate the presence of certain features, steps, operations, devices, components and / or their combinations.
[0053] As people's concerns about data privacy continue to increase, federated learning has gradually gained favor because it can train models without accessing user training data. However, in federated learning, the heterogeneity of data between users has a significant negative impact on the final performance of the global model. Therefore, how to solve the problems caused by user data heterogeneity has important application prospects and commercial value for the practical application of federated learning.
[0054] like Figure 1 As shown, based on a series of experimental observations, the present invention finds that the existing federated learning methods have significant defects, which are specifically manifested as follows: due to the lack of empty class (i.e., data that does not contain a specific category) samples in the training data, the performance of the local model on the empty class drops sharply. In addition, large-category data occupies a large proportion in the training process, causing the user model to overfit on the large class, while small-category samples are underfit due to insufficient training data. Therefore, protecting the local model's classification ability for empty classes and effectively suppressing overfitting are key issues in improving the performance of the federated learning model. Based on this, Figure 2 As shown, the present invention proposes a method for solving the data heterogeneity problem in federated learning based on empty class distillation and large class suppression, comprising the following steps:
[0055] Step S100: Each user counts the training data categories that are not included in his / her local data set, which are recorded as empty class sets. Specifically:
[0056] The training data categories that are not included in the local statistics of each user are called empty class sets. . Empty class collections for different users They can be the same or different, but the sum of all users’ training data will cover all categories’ training data.
[0057] The statistics of the empty class set are performed before training begins, and the statistical process only needs to be performed once, and no additional computing resources will be consumed during the training process.
[0058] Step S200, during the local training process, the logit value corresponding to the empty category in each sample output of the global model is used to constrain the logit value corresponding to the local model, ensuring that the classification ability of the empty category is maintained during the local model update process, thereby alleviating the model update differences between different users.
[0059] Since empty category samples are missing in local training data, this will cause the user model to gradually lose classification information on empty categories. To avoid this problem, the present invention proposes to distill the information about empty categories in the global model and inject it into the local model.
[0060] The specific implementation method is to constrain through the empty class distillation loss function, which is specifically defined as follows:
[0061] ,
[0062] in, is expectation, is the training data set of the kth user, represents the set of all empty categories of the k-th user, and The user model and the global model are respectively After The output probability of the empty category is as follows:
[0063] ,
[0064] ,
[0065] in, and Represent the local model and global model parameters respectively, Representative model input samples After The output logit of the empty category, t represents the category of the empty category set of the kth user.
[0066] It is important to note that during the local update process, only the parameters of the user model need to calculate gradients and perform parameter updates, while the parameters of the global model remain fixed during the local update. Therefore, when each user performs local training, the output logit value of the global model for the same training sample remains unchanged.
[0067] In order to reduce the consumption of communication resources, each user usually uploads the updated local model to the server after local training for multiple epochs. In this process, the output logit of the global model for each sample remains constant.
[0068] To further save computing resources, it is recommended to calculate the output logit of the global model for each sample in the first epoch of each user's local training and save it. After that, the user does not need to recalculate the output logit of the global model in subsequent epochs, which effectively reduces the computing burden.
[0069] Here, when the loss function is added to the original local training cross entropy loss function, a weight coefficient α needs to be introduced. The specific value of this coefficient can be set according to different task requirements or adjusted through experience. Experiments show that α has good robustness and can provide stable performance in different tasks.
[0070] Step S300, each user counts the probability of category distribution of the local training data set;
[0071] In this step, users only need to count the distribution probabilities of the categories already included in their local training data, without considering the empty categories. This statistical process should be completed before training begins and only needs to be performed once, without introducing additional computational overhead in subsequent training processes.
[0072] Specifically, the category distribution probability of the user's local training data set is specifically defined as follows:
[0073] ,
[0074] in, Representing user The probability of a category appearing, Represents the user's The sample size of training data for each category, Represents the total amount of training data for the user, specifically defined as ,in Represents the collection containing all categories.
[0075] Step S400, during the local training process, based on the probability of the local training data distribution obtained by the above statistics, the logit output of each sample is constrained to better punish the large-category logit value in the small-category sample output, thereby effectively alleviating the problem of small-category samples being misclassified as large-category samples, and ultimately alleviating the differences in models between users.
[0076] It should be pointed out that existing methods usually use the logit output of the weighted model to deal with the problem of inconsistent user data distribution. The specific form is as follows:
[0077] ,
[0078] in, Represents the user The probability of a category appearing, is the model input sample After that, output The logit value of each class.
[0079] However, the present invention has shown through experimental observation that despite the use of the above loss function, there is still a large proportion of small class samples being misclassified as large class categories. Therefore, directly punishing small class samples that are misclassified as large class categories becomes the key to solving this problem.
[0080] To this end, the present invention proposes an improved method by using the user local training data distribution probability obtained by previous statistics , constrains the logit output of each sample. This constraint is intended to reduce the phenomenon that small class samples are misclassified as large class samples. This loss function is called the "large class suppression loss function", and the specific definition of the function is as follows:
[0081] ,
[0082] in, Represents the number of samples output by the user model after a batch of samples is input. The sum of the logit values for the categories.
[0083] Furthermore, the loss function is explained as follows. For the logit output of the small class sample, the logit value of the large class is usually larger, which causes the small class sample to be misclassified as the large class. Specifically, the probability distribution of the large class is Usually large, while the probability distribution of small classes Therefore, the introduction of To weight the logit output of each category to strengthen the penalty for misclassifying small class samples into large class categories.
[0084] It should be noted that It is not enough to directly optimize the logit value of the category, which may lead to instability in the user model convergence process and thus affect the final performance. The sum of the logit values of the classes is adjusted and the following improved optimization strategy is defined:
[0085] ,
[0086] in, Representative samples The output does not correspond to the correct label. is the model input sample After that, output The logit of each class. is the user model parameter.
[0087] Here, the logit for the correct category is obtained by Directly optimized.
[0088] Therefore, the large-category suppression loss function proposed in the present invention can alleviate the problem of small-category samples being misclassified into large-category categories without affecting the normal learning of other samples.
[0089] In order to control the influence of the large-class suppression loss function on the entire training process, this implementation also introduces a coefficient , which is used to adjust the influence of the large-category suppression loss function under different classification tasks and different degrees of data heterogeneity. It can be set according to specific task requirements or tuned through experiments, thus providing flexible adaptability in different application scenarios.
[0090] Final loss function Has the following form:
[0091] ,
[0092] Usually, here The default setting is 0.1. Also set to 0.1 by default.
[0093] Step S500: Use the trained global model to test the accuracy on a test data set containing all categories.
[0094] On the other hand, the present invention also provides a system for solving the data heterogeneity problem in federated learning based on empty class distillation and large class suppression, wherein each unit included in the system can execute each step of the aforementioned method, specifically including:
[0095] A statistical unit, for each user, to count the categories of which no training samples exist in the local training data set, and record the set consisting of the categories of which no training samples exist as an empty class set;
[0096] The empty class distillation constraint unit is used to constrain the logit value of the empty class corresponding to each sample output of the global model and the local model based on the empty class distillation loss function during the local training process;
[0097] A distribution probability calculation unit is used to calculate the category distribution probability of each user's local training data set;
[0098] The large-category suppression constraint unit is used to directly constrain the logit value of each sample output by the local model based on the large-category suppression loss function based on the statistical category distribution probability of the local training data set during the local training process, so as to penalize the large-category logit value in the small-category sample output.
[0099] In order to verify the effectiveness of the present disclosure and further explain the specific implementation of the present disclosure in detail, the proposed method is applied to a public database, the CIFAR-10 database. The database contains 10 categories, and the data is divided into 50,000 training samples and 10,000 test samples. In order to simulate the data heterogeneity distribution among users in the real world, the Dirichlet sampling method is used to distribute the 50,000 training samples to 10 users to form a training set for each user. The degree of data heterogeneity among users is expressed by the Dirichlet coefficient Control, where The smaller the value, the greater the degree of data heterogeneity among users. It should be noted that each user's training set usually contains major categories, minor categories, and empty categories. Specifically, major categories are categories for which the user has a large amount of training data. Minor categories are categories for which the user has a small amount of training data; and empty categories are categories for which the user has no training data.
[0100] To test the effectiveness of the method proposed in this disclosure, the trained global model was tested on a test set containing 10,000 test images. The specific results are as follows: Figure 3 As shown in the figure, the results show that the method of the present invention significantly improves the performance of the global model. At the same time, the effects of the empty class distillation loss function and the large class suppression loss function on the performance of the final global model are verified. The experimental results are shown in Figure 4 As shown in the figure, the results on CIFAR10, CIFAR10 and TinyImagenet datasets are shown respectively, and the results prove that each loss function in the method of the present invention is crucial to improving the performance of the global model. Therefore, each experiment effectively verifies the effectiveness of the disclosed method in solving data heterogeneity tasks in federated learning.
[0101] The specific embodiments described above further illustrate the objectives, technical solutions and beneficial effects of the present invention in detail. It should be understood that the above description is only a specific embodiment of the present invention and is not intended to limit the present invention. Any modifications, equivalent substitutions, improvements, etc. made within the spirit and principles of the present invention should be included in the scope of protection of the present invention.
Claims
1. A method for solving the data heterogeneity problem in federated learning based on empty class distillation and large class suppression, characterized in that: The following steps are involved: Step S100, each user counts the categories without training samples in the local training data set, and records the set consisting of the categories without training samples as an empty class set; Step S200, during the local training process, the logit value of the global model and the local model corresponding to the empty category in each sample output is constrained based on the empty category distillation loss function; Step S300, each user counts the category distribution probability of the local training data set; Step S400, during the local training process, based on the statistical category distribution probability of the local training data set, directly constrain the logit value of each sample output by the local model based on the large category suppression loss function to penalize the large category logit value in the small category sample output.
2. According to claim 1, a method for solving data heterogeneity problems in federated learning based on empty class distillation and large class suppression is characterized in that: The operation of counting empty classes for each user in step S100 is performed only once during the entire training process; the local empty class sets between the users may be the same or different, but the local training data sets of all users added together contain data of all categories.
3. According to claim 1, a method for solving data heterogeneity problems in federated learning based on empty class distillation and large class suppression is characterized in that: The constraint based on the empty class distillation loss function in step S200 includes distilling the information about the empty class in the global model and injecting it into the local model.
4. According to claim 1, a method for solving data heterogeneity problems in federated learning based on empty class distillation and large class suppression is characterized in that: In step S200, the global model remains unchanged during the local update process, and the local model changes continuously during the training process. The local model and the global model have logit outputs of the same shape on the same input sample.
5. According to claim 1, a method for solving data heterogeneity problems in federated learning based on empty class distillation and large class suppression is characterized in that: In the step S200, the Kullback-Leibler divergence constraint method is adopted based on the empty class distillation loss function.
6. The method for solving the data heterogeneity problem in federated learning based on empty class distillation and large class suppression according to claim 5 is characterized in that: The empty class distillation loss function works together with the local training cross entropy loss function to balance the influence of the empty class distillation loss function on the local training cross entropy loss function based on the first weight coefficient.
7. The method for solving data heterogeneity problem in federated learning based on empty class distillation and large class suppression according to claim 1 is characterized in that: In the step S300, each user counts the category distribution probability of the local training data set before the training starts, and only needs to count once, and does not count empty classes.
8. The method for solving data heterogeneity problem in federated learning based on empty class distillation and large class suppression according to claim 1 is characterized in that: In the step S400, based on the large-category suppression loss function constraint, the logit value is only applied to the non-label category, and the logit value is processed by exponentially calculating the expectation of the logit output by the local model, and then taking the logarithm of the obtained expectation.
9. The method for solving data heterogeneity problem in federated learning based on empty class distillation and large class suppression according to claim 8 is characterized in that: The step S400 also includes using a second weight coefficient to adjust the influence of the large-category suppression loss function under different classification tasks and different degrees of data heterogeneity.
10. A system for solving data heterogeneity problems in federated learning based on empty class distillation and large class suppression, characterized in that: include: A statistical unit, for each user, to count the categories of which no training samples exist in the local training data set, and record the set consisting of the categories of which no training samples exist as an empty class set; The empty class distillation constraint unit is used to constrain the logit value of the empty class corresponding to each sample output of the global model and the local model based on the empty class distillation loss function during the local training process; A distribution probability calculation unit is used to calculate the category distribution probability of each user's local training data set; The large-category suppression constraint unit is used to directly constrain the logit value of each sample output by the local model based on the large-category suppression loss function based on the statistical category distribution probability of the local training data set during the local training process, so as to penalize the large-category logit value in the small-category sample output.
Citation Information
Patent Citations
Federal learning method, device and system for active directional data distillation
CN117669698A