Federal learning method based on partial label mask weighted distillation

By using partial label mask weighted distillation method in federated learning, the predicted values ​​of majority and minority classes are decoupled and processed, and the teacher model is constructed based on the knowledge of the global models before and after two times, the client forgetting problem caused by data heterogeneity in federated learning is solved, and the performance and robustness of the model in heterogeneous data scenarios are improved.

CN120163260APending Publication Date: 2025-06-17NORTHWESTERN POLYTECHNICAL UNIV
View PDF 0 Cites 3 Cited by

Patent Information

Application Number
CN202411870411.6
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2024-12-18
Publication Date
2025-06-17

AI Technical Summary

Technical Problem

In federated learning, due to the catastrophic forgetting problem of client locality caused by data heterogeneity, existing methods fail to fully consider the local data distribution, especially the impact of a few samples, resulting in the degradation of the model's performance in data heterogeneous scenarios.

Method used

A federated learning method based on partial label mask weighted distillation is proposed. By decoupling the logits output from the global model, the original predicted value of the majority class samples are blocked and the prediction value of the minority class samples are retained, thereby strengthening the learning of local minority class samples. In addition, a teacher model weighted by global model was constructed, and the knowledge of the global model was used before and after two times was enhanced to enhance the learning ability of the student model.

Benefits of technology

It significantly alleviates the client forgetting problem caused by data heterogeneity, improves the performance and robustness of the model in data heterogeneous scenarios, and ensures the stability and accuracy of the model during training.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN120163260A_ABST
    Figure CN120163260A_ABST
Patent Text Reader

Abstract

The invention discloses a federated learning method based on partial label mask weighted distillation, which comprises the following steps of: constructing a federated learning system comprising a central server and a plurality of clients, and broadcasting a global model to each client by the central server; the client performs calculation by using the received global model to obtain a teacher model, regards a local model as a student model, performs distillation training on the local model by using local data, and uploads trained model parameters to the central server; and the central server performs federal aggregation on the updated local model parameters, and issues the aggregated and updated local model to each client for next round of local training. According to the method, the client side learns knowledge of the global data set in the teacher model through partial label mask distillation, learning of minority class samples in the local data set is enhanced, learning of the local model on the knowledge of the global data set is enhanced by constructing the teacher model with more complete knowledge, and the learning efficiency of the local model is improved. And the risk of disastrous forgetting caused by data isomerism is reduced.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention belongs to the field of application of federated learning technology, and relates to a federated learning method based on partial label mask weighted distillation. Background Art

[0002] Nowadays, with the rise of artificial intelligence technology, privacy security has become a key challenge faced by current AI algorithms and models. Traditional centralized learning requires centralized management of all training data, which greatly increases the risk of data leakage, and people are often reluctant to upload their private data to commercial servers for supporting AI training. In such a scenario, Federated Learning (FL) emerged as a novel distributed machine learning paradigm, which can ensure that each distributed node cooperatively trains a global model for tasks such as image classification, object recognition, natural language processing, etc. without revealing local data.

[0003] The biggest difficulty faced by FL is data heterogeneity, which is manifested in the fact that the data distributions between clients are non-independent and identically distributed (non-iid). This data heterogeneity makes it difficult for the global model to learn features with strong generalization ability, and the differences in label distributions among different clients will also lead to contradictions when the model is globally aggregated, resulting in a slowdown in the model convergence speed and a significant decline in model performance. In addition, the characteristic that only a small part of clients participate in training in each communication round of FL will also exacerbate the local forgetting of clients. Specifically, in the non-iid scenario, the category test accuracy shown by the global model during the training process fluctuates greatly, and there is often such a phenomenon: the categories that can be well predicted in the first few communication rounds will have a significant reduction in their category test accuracy after several rounds of training, and even cannot be recognized, indicating that the forgetting phenomenon is occurring.

[0004] Regarding the challenge of local catastrophic forgetting of clients caused by data heterogeneity, the existing methods have the following problems: (1) Insufficient learning of global knowledge outside the local data distribution of local clients; (2) When masking the majority class labels of the local dataset, the influence of minority class samples is not considered, and directly masking the original majority class prediction values of the minority class samples in the output of the global model will cause the minority class samples to not be fully learned, thus affecting the effect of knowledge distillation; (3) Usually directly using the global model as the teacher model cannot provide more comprehensive global knowledge to local clients.

[0005] Due to the above problems, regarding the challenge of local catastrophic forgetting of clients caused by data heterogeneity, the effects of existing federated learning methods have not yet reached the best, and there is still a large room for performance improvement. How to alleviate and improve the adverse effects of client forgetting on model training still has great research value. Summary of the Invention

[0006] This paper proposes a novel federated learning algorithm based on partial label mask weighted distillation to improve the performance of the model under data heterogeneity by reducing the forgetting risk in FL. We decouple the logits output by the global model according to the true labels and local distributions of the samples, mask the majority-class logits of the global model for majority-class samples, and retain the minority-class samples, so as to strengthen the learning of local minority-class samples while learning knowledge outside the local data distribution, significantly alleviating the catastrophic impact brought by forgetting. In addition, we also propose a method for constructing a teacher model with weighted global models, using the partial participation characteristics of FL to enhance the knowledge reserve of the teacher model and strengthen the learning of the student model about the knowledge of non-self majority-class samples.

[0007] In view of the above technical problems, the present invention proposes a federated learning method based on partial label mask weighted distillation to improve the performance of the model under data heterogeneity by reducing the forgetting risk of the client. First, the present invention decouples the original predicted values output by the global model according to the true labels and local distributions of the samples, masks the majority-class original predicted values of the global model for majority-class samples, and retains the minority-class samples, so as to strengthen the learning of local minority-class samples while learning knowledge outside the local data distribution. In addition, the present invention constructs a teacher model based on the global models before and after, enhances the knowledge reserve of the teacher model, and strengthens the learning of the student model about the knowledge of non-self majority-class samples, thus significantly alleviating the catastrophic impact brought by forgetting. To achieve the above object, the technical solution adopted by the present invention is as follows:

[0008] Step 1. Construct a typical federated learning system including a central server and N clients:

[0009] Step 101. In the t-th communication round, the central server broadcasts the global model w t to all clients;

[0010] Step 102. The central server selects a subset S t of some active clients with probability ε to participate in the training of this communication round;

[0011] Step 2. The selected clients use the received global model to calculate the teacher model, regard the local model as the student model, then use the local data to perform distillation training on the local model, and upload the trained model parameters to the central server:

[0012] Step 201. The local client i ∈ S t uses the global models w t-1 and w t, the teacher model \(w\) is calculated dg , that is:

[0013] \(w\) dg =\(\alpha w\) t-1 +(1 - \(\alpha\))\(w\) t

[0014]

[0015] where \(\alpha\) is the weight coefficient, \(d\) is the distance between the two consecutive global models \(w\) t-1 and \(w\) t , and \(E\) is the number of times the local model is trained and updated;

[0016] Step 202, for the local client \(i\in S\) t Perform partial label masking and calculate the outputs \(q'\) and \(q'\) of the sample \((x, y)\) in the local model and the teacher model, that is: i \(q'\) dg \(q'\)

[0017]

[0018] where \(k\) is the label index of the sample, \(C\) is the total number of sample labels in the dataset, \(k = 1, 2, \cdots, C\), \(z\) i,k is the original predicted value of the \(k\)-th category output by the sample \((x, y)\) in the local model of the \(i\)-th client, \(z\) dg,k is the original predicted value of the \(k\)-th category output by the sample \((x, y)\) in the teacher model, \(\tau\) is the distillation temperature, \(M\) i is the label set of the majority class samples of the \(i\)-th client;

[0019] Step 203, for the local client \(i\in S\) t Calculate the partial label masking weighted distillation loss using the outputs \(q'\) and \(q'\) of the local model and the teacher model i \(q'\) dg \(q'\) That is:

[0020]

[0021] Step 204, for the local client \(i\in S\) t Perform \(E\) local updates on the local model using the stochastic gradient descent method, that is:

[0022]

[0023] where \(\eta\) is the learning rate, is the local loss function, is the classical cross-entropy loss, \(q\) i and 1 yThey are the output of the sample (x, y) in the local model and the one-hot label of the sample, respectively, and β is the weight coefficient of the distillation loss. ;

[0024] Step 205: The client uploads the updated local model to the central server.

[0025] Step 3: The central server performs federated aggregation on the parameters of the updated local model and distributes the aggregated and updated local model to each client for the next round of local training:

[0026] Step 301: The central server performs weighted averaging on the received local models, that is:

[0027]

[0028] where n is the total number of local samples of all clients, is the number of local samples of client i;

[0029] Step 302: The central server distributes the updated global model w t+1 to each client for the next round of local training;

[0030] Step 4: After the central server and the clients communicate for T rounds, the final model is output:

[0031] Step 401: After the central server and the clients communicate for T rounds, the final model w T of this training is calculated, that is:

[0032]

[0033] A federated learning method based on partial label mask weighted distillation provided by the present invention has the following characteristics compared with the prior art:

[0034] (1) The client uses partial label mask distillation to learn the knowledge of the global dataset in the teacher model, while strengthening the learning of minority class samples in the local dataset, reducing the risk of catastrophic forgetting caused by data heterogeneity;

[0035] (2) Considering that the global models before and after contain the knowledge of different clients, the present invention constructs a teacher model with more complete knowledge, enhancing the learning of the local model about the knowledge of the global dataset;

[0036] (3) The federated learning method proposed by the present invention is still robust in the face of strong data heterogeneity and can maintain good model performance. Description of the Drawings

[0037] Figure 1It is a flowchart of the present invention;

[0038] Figure 2 It is a system architecture diagram of the present invention; Specific implementation manners

[0039] The method of the present invention will be further described in detail below in conjunction with the accompanying drawings and the embodiments of the present invention. Obviously, the described embodiments are only a part of the embodiments of the present invention, rather than all of the embodiments. All other embodiments obtained by those of ordinary skill in the art based on the embodiments of the present invention without making creative efforts fall within the scope of protection of the present invention.

[0040] Figure 1 It represents a flowchart of a federated learning method based on partial label mask weighted distillation.

[0041] Figure 2 It represents a system architecture diagram of a federated learning method based on partial label mask weighted distillation.

[0042] To facilitate the understanding of the embodiments of the present invention, the rationality and effectiveness of the present invention will be described below in conjunction with the accompanying drawings, including the following specific steps:

[0043] Step 1: Construct a typical federated learning system including a central server and N = 100 clients:

[0044] Step 101: In the t-th communication round, the central server broadcasts the global model w t to all clients;

[0045] Step 102: The central server selects a subset S of partial active clients with a probability ε = 10% t to participate in the training of this communication round;

[0046] Step 2: The selected clients use the received global model to calculate the teacher model, regard the local model as the student model, then use the local data to perform distillation training on the local model, and upload the trained model parameters to the central server:

[0047] Step 201: The local client i ∈ S t uses the global models w t-1 and w t received before and after to calculate the teacher model w dg , that is:

[0048] w dg = αw t-1 +(1 - α)w t

[0049]

[0050] Among them, α is the weight coefficient, and d is the distance between two consecutive global models w t-1 and w t , and E = 5 is the number of times of local model training and update;

[0051] Step 202: The local client i ∈ S t Performs partial label masking and calculates the outputs q i ' and q dg ' of the sample (x, y) in the local model and the teacher model, that is:

[0052]

[0053] Among them, k is the label index of the sample, C = 10 is the total number of sample labels in the dataset, k = 1, 2,..., C, and z i,k is the original predicted value of the k-th category output by the sample (x, y) in the local model of the i-th client, and z dg,k is the original predicted value of the k-th category output by the sample (x, y) in the teacher model. τ = 1 is the distillation temperature, and M i is the label set of the majority class samples of the i-th client;

[0054] Step 203: The local client i ∈ S t Calculates the partial label masking weighted distillation loss using the outputs q i ' and q dg ' of the local model and the teacher model, that is: That is:

[0055]

[0056] Step 204: The local client i ∈ S t Performs E = 5 local updates on the local model using the stochastic gradient descent method, that is:

[0057]

[0058] Among them, η = 0.1 is the learning rate, is the local loss function, is the classical cross-entropy loss, q i and 1 y are the output of the sample (x, y) in the local model and the one-hot label of the sample respectively. β = 1 is the weight coefficient of the distillation loss ;

[0059] Step 205: The client uploads the updated local model to the central server;

[0060] Step 3. The central server performs federated aggregation on the updated local model parameters and distributes the aggregated and updated local model to each client for the next round of local training:

[0061] Step 301. The central server performs weighted averaging on the received local models, i.e.,

[0062]

[0063] where n is the total number of local samples of all clients, is the number of local samples of client i;

[0064] Step 302. The central server distributes the updated global model w t+1 to each client for the next round of local training;

[0065] Step 4. After the central server and the clients communicate for T = 200 rounds, the final model is output:

[0066] Step 401. After the central server and the clients communicate for T = 200 rounds, the final model w T of this training is calculated, i.e.,

[0067]

Claims

1. A federated learning method based on partial label mask weighted distillation, characterized in that: The following steps are involved: Step 1: Build a typical federated learning system consisting of a central server and N clients: Step 101: In the tth communication round, the central server sends the global model w t Broadcast to all clients; Step 102: The central server selects a subset of active clients S with probability ε t Training for participation in this communication round; Step 2: The selected client uses the received global model to calculate the teacher model, and regards the local model as the student model. It then uses the local data to perform distillation training on the local model and uploads the trained model parameters to the central server: Step 201: local client i∈S t Using the global model w received twice before and after t-1 and w t , calculate the teacher model w dg ,Right now: Among them, α is the weight coefficient, d is the global model w before and after t-1 and w t The distance between them, E is the number of local model training updates; Step 202: Local client i∈S t Perform partial label masking and calculate the output q of the sample (x, y) in the local model and the teacher model i ′ and q dg ',Right now: Where k is the label index of the sample, C is the total number of sample labels in the dataset, k = 1, 2, ..., C, z i,k is the original prediction value of the kth category output by the local model of the i-th client, z dg,k is the original prediction value of the kth category output by the sample (x, y) in the teacher model, τ is the distillation temperature, M i is the label set of the majority class samples of the i-th client; Step 203: local client i∈S t Using the output q of the local model and the teacher model i ′ and q dg ' Calculate the partial label mask weighted distillation loss Right now: Step 204: local client i∈S t The local model is updated E times locally using the stochastic gradient descent method, namely: Where η is the learning rate, is the local loss function, is the classic cross entropy loss, q i and 1 y are the output of the sample (x, y) in the local model and the unique hot label of the sample, and β is the distillation loss The weight coefficient of Step 205: The client updates the local model Upload to the central server; Step 3: The central server aggregates the updated local model parameters and sends the aggregated updated local model to each client for the next round of local training: Step 301: The central server performs weighted averaging on the received local models, namely: Where n is the total number of local samples of all clients, is the number of local samples of client i; Step 302: The central server updates the global model w t+1 Send it to each client for the next round of local training; Step 4: After the central server and the client communicate for T rounds, the final model is output: Step 401: After T rounds of communication between the central server and the client, the final model w of this training is calculated. T ,Right now:

Citation Information

Cited By

  • Field adaptive large model fine tuning method and system based on enterprise private data

    CN120892792A

  • Federal learning method based on comparative learning and knowledge distillation

    CN120952111A

  • Electromagnetic signal forgetting learning method based on mask distillation erasing

    CN121328644A