A Federated Learning Defense Method, Electronic Device, and Storage Medium

By constructing a historical global model storage pool and K-means clustering algorithm, combining knowledge distillation and FedAvg aggregation, the problem of non-independent and homogeneous data and poisoning attacks in federated learning is solved, the model performance and robustness are improved, and client privacy is protected.

CN120030536BActive Publication Date: 2025-07-22SHANDONG COMP SCI CENTNAT SUPERCOMP CENT IN JINAN
View PDF 4 Cites 0 Cited by

Patent Information

Application Number
CN202510496401.9
Authority / Receiving Office
CN · China
Patent Type
Patents(China)
Current Assignee / Owner
Filing Date
2025-04-21
Publication Date
2025-07-22
Estimated Expiration
2045-04-21

AI Technical Summary

Technical Problem

When federated learning faces non-independent homogeneous data and poisoning attacks, the existing technology is difficult to effectively defend, resulting in reduced performance and insufficient robustness of the global model.

Method used

By building a historical global model storage pool, the server uses the K-means clustering algorithm to filter the client model, and performs knowledge distillation and personalized training on the client. Combined with the FedAvg aggregation method, a stable teacher model is constructed to resist poisoning attacks.

Benefits of technology

Effectively identifying and eliminating malicious clients improves the performance and robustness of the global model, ensures the security and overall stability of the federated learning system, and protects the privacy of the client.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN120030536B_ABST
    Figure CN120030536B_ABST
Patent Text Reader

Abstract

The present invention belongs to the technical field of information security, and particularly relates to a federated learning defense method, an electronic device, and a storage medium. The present invention is a defense method for resisting poisoning attacks applicable to a non-independent and identically distributed data environment. By fusing historical global models to construct a teacher model, knowledge distillation and personalized training are carried out on the client side, and the K-means aggregation method is jointly used with the server side to defend against model poisoning attacks. By simulating the attack and defense mechanisms in the training process, this method and system can identify and eliminate malicious clients, not only improving the defense ability of the federated learning framework against model poisoning attacks, but also significantly enhancing the performance of the final global model, ensuring the overall robustness and security of the federated learning system.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention belongs to the technical field of information security, and particularly relates to a federated learning defense method, an electronic device, and a storage medium. Background Art

[0002] Federated learning is a distributed machine learning framework that allows multiple clients to perform model training locally. After local training is completed, the clients send model updates to the server side for aggregation to obtain a global model. The greatest advantage of federated learning is that clients do not need to share local data, thus protecting privacy. In the federated learning scenario, the local data of clients is often non-independent and identically distributed. This data imbalance poses a huge challenge to federated learning because the traditional aggregation method FedAvg performs well when dealing with independent and identically distributed data, but under non-independent and identically distributed data conditions, the convergence and accuracy of the global model will be affected, resulting in a decline in model performance or deviation. At the same time, federated learning also faces many security threats, especially attacks from malicious clients. One of the most common attack types is the poisoning attack. The poisoning attack is divided into a data poisoning attack that manipulates training data and a model poisoning attack that tamper with local model updates. Both of these attack methods will interfere with the training of the global model and even cause the model to output specific biased results. To defend against these attacks, additional mechanisms need to be introduced in the federated learning framework to detect, isolate, and filter malicious updates. Traditional defense methods include model verification based on redundancy, anomaly detection, or robust algorithms based on aggregation, but these methods may not always be effective when facing complex Non-IID data.

[0003] Chinese Patent Document CN118761455A discloses a personalized federated learning method based on knowledge distillation, including the following steps: initializing the global model of the central server and the number of local devices; broadcasting the initialized global model to the local devices; all local devices use local private data for model training; and uploading the local model parameters to the central server for aggregation to generate a new global model; pruning the newly generated global model by unstructured pruning technology to turn the global model into a teacher model friendly to students; transferring the final prediction and attention map knowledge in the pruned teacher model to the local model, and the local model is used as a student model at this time and is updated according to the transferred knowledge. However, this method does not explicitly mention a defense mechanism against malicious clients or model poisoning attacks. Especially when facing data poisoning or model poisoning attacks, there is a lack of additional detection and elimination mechanisms; at the same time, this method does not propose a dedicated solution to the special challenges of non-independent and identically distributed data, which may lead to a decline in the performance of the global model in the case of inconsistent data distribution.

[0004] Chinese Patent Document CN118246009A discloses a method for defending against poisoning attacks in federated learning. This method first initializes parameters. The server initializes the global model and the learning rate and broadcasts them to each client. Secondly, the client receives the global model from the server, calculates the local gradient and loss value using local data, and then uploads the local gradient and loss value to the server. Then, according to the gradient, the server uses the local outlier factor algorithm to eliminate malicious parameter clusters and retain normal parameter clusters. Finally, the server performs score calculation through the gradients and loss values uploaded by each client in the normal parameter cluster, calculates the global model of the (t + 1)-th round, and continuously iterates and updates. However, this method only relies on the local outlier factor to judge malicious clients. This defense strategy may not cover all attack patterns when facing diverse and complex poisoning attacks, especially limited in the effect on data poisoning attacks or misleading attacks; and this method only relies on gradients and loss values for defense, lacking effective knowledge transfer between the client and the global model.

[0005] To address the challenges of federated learning in the face of non-independent and identically distributed data and poisoning attacks, the present invention proposes a federated learning defense method, an electronic device, and a storage medium. Summary of the Invention

[0006] The present invention aims to overcome at least one defect of the above-mentioned prior art and provides a federated learning defense method. By constructing a historical global model storage pool to fuse global models of different rounds, using these historical models as teacher models, and performing local model training on the client through knowledge distillation technology; in addition, the server uses the K-means clustering algorithm to cluster the model updates uploaded by the client to detect and eliminate potential malicious clients; by combining these technologies, the system can not only enhance the resistance to poisoning attacks, but also improve the performance of the global model in a non-independent and identically distributed environment, ensuring the stability and accuracy of the model.

[0007] The present invention also discloses an electronic device for implementing the above method.

[0008] The present invention also discloses a machine-readable storage medium for implementing the above method.

[0009] To achieve the above object, the technical solution of the present invention is as follows:

[0010] A federated learning defense method, the method comprising:

[0011] S1. Initial work on the server side:

[0012] S11. Construct a historical global model storage pool on the server side, and randomly initialize a global model as the initialized global model and put it into the historical global model storage pool;

[0013] S12. Simultaneously distribute the initialized global model as both the teacher model and the global model to the clients randomly selected in the first round to participate in the training.

[0014] S2. Client training:

[0015] S21. After the client receives the global model and the teacher model sent by the server, the client conducts local training on the global model based on the local dataset to generate a client model and produce a local loss. At the same time, the client conducts knowledge distillation training on the teacher model based on the local dataset to generate a student model and produce a knowledge distillation loss.

[0016] S22. Jointly train the client model and the student model generated in step S21 to generate an updated local model, and send the updated local model back to the server to participate in the aggregation of the global model in the next round.

[0017] S3. Server aggregation:

[0018] S31. After the server receives the updated local model described in step S22, generate a new round of global model by aggregating the local models, and store it in the historical global model storage pool.

[0019] S32. Aggregate the global models in the historical global model storage pool to obtain a new round of teacher model.

[0020] S4. Simultaneously distribute the new round of global model generated in step S31 and the new round of teacher model generated in step S32 to the clients selected for training in the next round, and loop steps S2 and S3 until the specified number of rounds is reached or the global model converges to end the training.

[0021] Preferably, the dataset described in step S21 is an image dataset.

[0022] Preferably, the client conducts local training on the global model based on the local dataset to generate a client model and produce a local loss in step S21 specifically as follows: The client conducts training based on its local dataset and fine-tunes the global model. The local loss calculation formula (1) generated by the difference between the client model and the global model is as follows:

[0023] (1)

[0024] Where refers to the local loss of the round client The local loss of the is the local dataset of client k, refers to the global model sent by the server in the round, is the loss function, , respectively represent the input and output vectors, represents the model output calculated using the global model of represents the expected value of the random variable pair sampled from the distribution .

[0025] Preferably, the client in step S21 performs knowledge distillation training on the teacher model based on the local dataset to generate a student model and generates a knowledge distillation loss, specifically:

[0026] The client transfers the knowledge in the teacher model to the student model through knowledge distillation, which is achieved by minimizing the difference between the outputs of the student model and the teacher model. This process generates a knowledge distillation loss, and the student model is optimized by minimizing this loss:

[0027] (2)

[0028] where refers to the knowledge distillation loss of the round of the client , refers to the teacher model sent by the server in the round, Yes The distillation loss based on the Kullback-Leibler divergence measures the difference between the probability distributions of the outputs of the teacher model and the student model, 、 respectively represent the input and output vectors; represents the model output calculated using the teacher model of represents the expected value of the random variable pair sampled from the distribution .

[0029] Preferably, in step S22, the client model and the student model generated in step S21 are jointly trained to first generate a local model loss, and the updated local model is calculated through the local model loss.

[0030] More preferably, the formula for calculating the local model loss is as follows:

[0031] (3)

[0032] where refers to the round of the client The calculated local model loss, is a weighted coefficient that controls the balance between the local loss and the distillation loss, is the round of the client 's local loss, is the round of the client 's knowledge distillation loss;

[0033] The generation of the updated local model is specifically as follows:

[0034] (4)

[0035] where represents the learning rate, is the round of the client 's local model, that is, the updated local model, is the round of the client 's local model, represents the gradient of the local model loss with respect to the local model ;

[0036] Preferably, the generation of a new round of the global model by aggregating the local models in step S31 is specifically as follows:

[0037] After the server receives the updated local model uploaded by the client in step S22, it generates a global model by minimizing the clustering loss through the K-means aggregation method, and at the same time identifies and removes abnormal models:

[0038] (5)

[0039] where refers to the total loss function of K-means, represents the number of clusters for clustering. The number of attackers is less than the number of benign clients, so the local model in the largest cluster is selected for global model aggregation; represents the cluster 's local model , represents the t round of the client 's local model, represents the centroid of the cluster , represents the L2 norm, which is used to measure the distance between the local model and the clustering centroid . This formula makes similar models be grouped into the same cluster, reducing the probability of malicious clients participating in training;

[0040] The newly generated new round of global model :

[0041] ( 6)

[0042] wherein refers to the round of global model, represents the number of clients in the largest cluster participating in this round of training, refers to the round of client local model.

[0043] Preferably, the specific process of aggregating the global models in the historical global model storage pool in step S32 to obtain a new round of teacher model is as follows:

[0044] Use the FedAvg aggregation method to obtain the teacher model sent to the client in the round: :

[0045] (7)

[0046] wherein, refers to the teacher model sent by the server in the round, refers to the th global model in the historical global model storage pool.

[0047] In another aspect of the present invention, an electronic device is further provided, including:

[0048] At least one processor; and

[0049] A memory storing instructions, when the instructions are executed by the at least one processor, causing the at least one processor to execute the federated learning defense method as described above.

[0050] In another aspect of the present invention, a machine-readable storage medium is further provided, which stores executable instructions, and when the instructions are executed, the machine is caused to execute the federated learning defense method as described above.

[0051] Compared with the prior art, the present invention has the following beneficial effects:

[0052] (1) The present invention proposes a defense method against poisoning attacks applicable to non-independent and identically distributed data environments. By fusing historical global models to construct a teacher model, performing knowledge distillation and personalized training on the client side, and jointly using the K-means aggregation method on the server side to defend against model poisoning attacks, it can effectively resist model poisoning attacks in federated learning. It is mainly used to remove malicious clients in federated learning and improve the performance of the final global model. By simulating the attack and defense mechanisms in the training process, this method system can identify and eliminate malicious clients, not only improving the defense ability of the federated learning framework against model poisoning attacks, but also significantly enhancing the performance of the final global model, ensuring the overall robustness and security of the federated learning system.

[0053] (2) The establishment of the "historical global model storage pool" in the present invention ensures that the system can call historical models for fusion at any time, thereby constructing a more stable and robust teacher model, providing support for the subsequent knowledge distillation process, and at the same time avoiding the requirement of the server side relying on a small dataset in traditional knowledge distillation, further protecting the privacy of the client.

[0054] (3) In each round of training, the server side not only distributes the global model screened and generated by the K-means clustering aggregation algorithm (the global model randomly initialized in the first round of training), but also distributes the teacher model fused by the FedAvg aggregation algorithm. These two models play different roles: the global model is used to guide the local training of the client, while the teacher model transfers the knowledge it contains to the student model of the client through knowledge distillation. After receiving these two models, the client applies knowledge distillation technology to extract the knowledge in the teacher model and transfer it to the student model of the client, which effectively reduces the model bias caused by inconsistent data distributions. At the same time, the client also integrates the knowledge of the global model into its own model to further optimize the global model.

[0055] (4) After completing local training, the client uploads the updated local model to the server side to participate in the new round of model aggregation. The server side uses the K-means clustering aggregation algorithm to aggregate the models uploaded by the clients, and generates a new round of global model after screening and eliminating potential malicious models. By combining the use of the K-means clustering algorithm, the server side can effectively identify abnormal client models, which are often uploaded by malicious clients with the intention of poisoning. Through clustering analysis, the similarity and difference of each client model can be effectively evaluated, so as to delete those potential attack models and further ensure the robustness of the global model. Through such continuous iteration and optimization, the entire federated learning system resists poisoning attacks while gradually improving the overall performance of the global model, ensuring the security and reliability of the model.

[0056] (5) The present invention aims to resist model poisoning attacks in the federated learning process and improve the performance of the global model by combining the collaboration of the client and the server. Brief Description of the Drawings

[0057] Figure 1 is a schematic diagram of the defense method framework of the present invention;

[0058] Figure 2 is a schematic diagram of the local model training process of the client of the present invention;

[0059] Figure 3 is a comparison chart of the number of training rounds and test accuracy of the global model in the MNIST dataset of Embodiment 1 of the present invention under the gradient ascent attack mode, comparing three defense methods and the case without defense;

[0060] Figure 4 is a comparison chart of the number of training rounds and test accuracy of the global model in the MNIST dataset of Embodiment 1 of the present invention under the gradient misleading attack mode, comparing three defense methods and the case without defense;

[0061] Figure 5 is a comparison chart of the number of training rounds and test accuracy of the global model in the CIFAR-10 dataset of Embodiment 1 of the present invention under the gradient ascent attack mode, comparing three defense methods and the case without defense;

[0062] Figure 6 is a comparison chart of the number of training rounds and test accuracy of the global model in the CIFAR-10 dataset of Embodiment 1 of the present invention under the gradient misleading attack mode, comparing three defense methods and the case without defense. Detailed Embodiments

[0063] Embodiment 1

[0064] The present invention provides a federated learning defense method, which realizes effective defense against federated learning attacks, such as Figure 1 , Figure 2 shown, the method includes:

[0065] S1. Initial work of the server:

[0066] S11. Build a historical global model storage pool on the server side, and randomly initialize a global model as the initial global model and put it into the historical global model storage pool for subsequent use;

[0067] Specifically, the historical global model storage pool is used to store the global models generated after aggregation in each round of federated learning.

[0068] S12. Simultaneously distribute the initialized global model as both the teacher model and the global model to the clients randomly selected in the first round for training.

[0069] S2. Client training:

[0070] S21. After receiving the global model and the teacher model sent by the server, the client performs local training on the global model based on its local dataset to generate a client model and produce a local loss. At the same time, the client performs knowledge distillation training on the teacher model based on its local dataset to generate a student model and produce a knowledge distillation loss. The dataset is an image dataset.

[0071] Specifically, the client performing local training on the global model based on its local dataset to generate a client model and produce a local loss in step S21 is as follows:

[0072] First, the client performs training based on its local dataset to fine-tune the global model. The difference between the client model and the global model generates a local loss:

[0073] (1)

[0074] where refers to the local loss of the round client , is the local dataset of the client k , refers to the global model sent by the server in the round, is the loss function, , represent the input and output vectors respectively, represents the model output calculated using the global model weights , represents the expected value of the random variable pair sampled from the distribution . Here, the difference loss between the global model and the client model output is calculated.

[0075] Specifically, the client performing knowledge distillation training on the teacher model based on its local dataset to generate a student model and produce a knowledge distillation loss in step S21 is as follows:

[0076] The client transfers the knowledge in the teacher model to the student model through knowledge distillation, which is achieved by minimizing the difference between the outputs of the student model and the teacher model. This process generates a knowledge distillation loss, and the student model is optimized by minimizing this loss.

[0077] (2)

[0078] Wherein refers to the round of client knowledge distillation loss, refers to the teacher model sent by the server in the round, which is a distillation loss based on Kullback-Leibler divergence, measuring the difference between the probability distributions of the outputs of the teacher model and the student model. 、 respectively represent the input and output vectors; represents the model output calculated using the teacher model and represents the expected value of the random variable pair sampled from the distribution . The formula calculates the knowledge distillation loss between the student model and the output of the teacher model.

[0079] Through knowledge distillation, not only can the performance of the local model be improved, but also the learning bias caused by the non-i.i.d. environment can be effectively reduced. In addition, by learning from the teacher model, the client can better resist attacks from malicious participants or contaminated local models. Even if some clients are attacked or upload unreliable models, the stability of the teacher model and the knowledge integration based on multiple rounds of global models can effectively filter out the influence of malicious or abnormal models, thus ensuring the robustness and security of the model.

[0080] S22. Jointly train the client model and the student model generated in step S21 to generate an updated local model, and send the updated local model back to the server to participate in the aggregation of the new round of global model.

[0081] Specifically, in step S22, when jointly training the client model and the student model generated in step S21, first generate a local model loss, and calculate the updated local model through the local model loss.

[0082] Specifically, the formula for calculating the generated local model loss is as follows:

[0083] (3)

[0084] Wherein refers to the round of client calculated local model loss, is a weighting coefficient that controls the balance between the local loss and the distillation loss, is the Round client of the local loss, is the round client of the knowledge distillation loss;

[0085] Finally, the client the round local model uploaded to the server side is the updated local model:

[0086] (4)

[0087] where represents the learning rate, is the round client of the local model, that is, the updated local model, is the round client of the local model, represents the gradient of the local model loss with respect to the local model . This local model combines both the characteristics of local data and the shared knowledge learned from the global and teacher models through distillation.

[0088] S3. Server-side aggregation:

[0089] S31. After the server side receives the updated local model described in step S22, it generates a new round of global model by aggregating the local models and stores it in the historical global model storage pool;

[0090] Specifically, the generation of a new round of global model by aggregating the local models described in step S31 is as follows:

[0091] After the server side receives the updated local model uploaded by the client described in step S22, it generates a global model by minimizing the clustering loss through the K-means aggregation method, and at the same time identifies and eliminates abnormal models. The specific process is as follows:

[0092] (5)

[0093] where refers to the total loss function of K-means, represents the number of clusters for clustering. The number of attackers is less than the number of benign clients, so we select the local model in the largest cluster for global model aggregation. represents the cluster in the local model , Represents the local model of the client , represents the centroid of the cluster . Represents the L2 norm, which is used to measure the distance between the local model and the cluster centroid . This formula groups similar models into the same cluster, reducing the probability of malicious clients participating in training.

[0094] Next, aggregate the local models in all the largest clusters to obtain a new round of global model :

[0095] (6)

[0096] where refers to the -th round of global model, is the number of clients in the largest cluster, refers to the -th round of client 's local model.

[0097] S32. Aggregate the global models in the historical global model storage pool to obtain a new round of teacher model;

[0098] Specifically, use the FedAvg aggregation method to obtain the teacher model sent to the client in the -th round

[0099] (7)

[0100] where refers to the teacher model sent by the server in the -th round, refers to the -th global model in the historical global model storage pool. Using this method, the server does not need to have an additional dataset, which can further protect the privacy of the data.

[0101] S4. At the same time, send the new round of global model generated in step S31 and the new round of teacher model generated in step S32 to the clients selected for training in the next round, and loop steps S2 and S3 until the specified number of rounds is reached or the global model converges to end the training.

[0102] The present invention proposes an effective federated learning defense method by combining knowledge distillation of the historical global model and the K-means clustering algorithm. By constructing a storage pool of historical global models, the effective fusion of historical global models is realized on the server side to form a stable teacher model, overcoming the dependence on an additional dataset with the same distribution as the client side in traditional knowledge distillation methods and protecting client privacy. On the client side, by combining global model guidance and teacher model knowledge distillation, the local model not only incorporates local data characteristics but also absorbs the shared knowledge of the global model, effectively reducing the bias caused by inconsistent data distributions and enhancing the ability to resist malicious attacks. At the same time, the server side uses the K-means clustering aggregation algorithm to accurately identify and eliminate potential models to ensure the robustness of the global model.

[0103] Specifically, experiments were conducted on two public datasets, including MNIST and CIFAR-10. MNIST is a grayscale image dataset containing handwritten digits, with each image sized 28×28 pixels. The training set contains 60,000 images, and the test set contains 10,000 images, divided into 10 categories corresponding to the digits 0-9, and the image content is relatively simple. CIFAR-10 is a color natural image dataset containing 10 categories, with each image sized 32×32 pixels and containing 3 color channels (RGB). The training set contains 50,000 images, and the test set contains 10,000 images, also divided into 10 categories, but it is more complex than MNIST.

[0104] In each round of experiments, 50 clients were set up, using a fully connected neural network with one hidden layer. The size of the hidden layer of the neural network model for each client was 1024, and a total of 200 rounds of training were conducted. Each client trained locally for two rounds per round, and the learning rate was 0.001. Initially, the initial parameters of each client model were the same as those of the global model. The attacker ratio was 0.4, and the attacker started attacking from the first round. For the defense mechanism of the present invention, the weight of the local loss was set to 0.3, the corresponding distillation loss weight was 0.7, and the number of clusters was 2, meaning it was divided into two clusters. The number of attackers was less than the number of benign clients, so we selected the category with the larger number in the two categories for global model aggregation.

[0105] The experiments of the present invention compared two attack methods: Gradient Ascent Attack and Mislead Attack. Three defense methods: Median Aggregation Method, Trimmed Mean Aggregation Method, Our Defense Method of the present invention, and the Baseline without defense.

[0106] In the gradient ascent attack, malicious clients deliberately increase the loss of the global model by updating the gradient in the reverse direction, making the model inaccurate.

[0107] The gradient misleading attack is a combined attack strategy that combines label flipping and gradient ascent. Its purpose is to disrupt the normal training process of the model through the combined operation of the normal training gradient and the cumulative gradient of the local model.

[0108] The median aggregation method is a method that offsets the influence of abnormal or attack points by selecting the median value in the training data, enhancing the robustness of the model.

[0109] The modified mean aggregation method maintains the consistency of the model weights by using the trained average weights when updating the model parameters, reducing the risk of being disturbed by extreme data points.

[0110] Figure 3 and Figure 4 and are the test accuracies of the global model under the gradient ascent attack and the gradient misleading attack in the MNIST dataset, respectively, when comparing three defense methods, namely the defense method of the present invention, the median aggregation method, and the modified mean aggregation method, and the case of no defense. In the gradient ascent attack, the test accuracy of the global model of the no-defense method fluctuates with the number of rounds, indicating that the model is vulnerable to gradient perturbations. In the gradient misleading attack, the test accuracy drops slightly and stabilizes at about 96% after 200 rounds because the misleading attack directly interferes with the model's decision boundary. The median aggregation method has a relatively good resistance to the gradient ascent attack, and the test accuracy remains at 92.5% after 200 rounds, but its effect on the gradient misleading attack is limited because the median aggregation cannot filter malicious updates with the same direction. The modified mean aggregation method is relatively stable under the gradient misleading attack, and the difference in the test accuracy from the defense method of the present invention is not significant. However, in the gradient ascent attack, the effect does not meet the expectation, possibly because excessive pruning of extreme values results in the loss of effective updates. The defense method of the present invention performs best in the gradient ascent attack, and the test accuracy remains above 97.5% after 200 rounds; it has significant robustness to the gradient misleading attack, and the test accuracy reaches above 98% by effectively distinguishing malicious updates through historical model fusion.

[0111] Figure 5 and Figure 6Shows the test accuracy of the global model in the CIFAR-10 dataset under gradient ascent attack and gradient misdirection attack, comparing three defense methods and the case without defense. Due to the high dimensionality of the data in CIFAR-10, the impact of the attack is more significant and the defense is more difficult. The test accuracy of the method without defense under misdirection attack fluctuates around 60%, and further drops from 65% to 22% under gradient ascent attack, indicating that complex data is more vulnerable to attack. The median aggregation method has limited effect on gradient ascent attack and gradient misdirection attack. Due to the dispersion of gradient directions in high-dimensional data, the median cannot effectively denoise. Figure 6 In Figure 6 , the modified mean aggregation method has a relatively high initial test accuracy under gradient misdirection attack, but it continues to fluctuate as the number of rounds increases and still cannot converge after 200 rounds, possibly due to slow model convergence caused by excessive pruning. The defense method of the present invention combines knowledge distillation and dynamic aggregation, alleviates the attack diffusion under high-dimensional data, and the test accuracy is stable above 73% in gradient ascent attack and remains at 75% under gradient misdirection attack, significantly superior to other methods.

[0112] In summary, the present invention combines knowledge distillation of the historical global model and the K-means clustering algorithm, and proposes an effective federated learning defense method, which improves the performance of the model in the non-independent and identically distributed data environment while enhancing the defense ability against model poisoning attacks, ensuring the security and robustness of the federated learning system, and avoiding the requirement of relying on a small dataset on the server side in traditional knowledge distillation, further protecting the privacy of the client.

[0113] Example 2

[0114] This embodiment also provides an electronic device, including:

[0115] At least one processor; and

[0116] A memory, where the memory stores instructions, and when the instructions are executed by the at least one processor, the at least one processor executes the federated learning defense method as described above.

[0117] In this embodiment, the electronic device may include, but is not limited to: personal computers, server computers, workstations, desktop computers, laptop computers, notebook computers, mobile computing devices, smart phones, tablet computers, cellular phones, personal digital assistants (PDAs), handheld devices, messaging devices, wearable computing devices, consumer electronic devices, and so on.

[0118] Example 3

[0119] This embodiment also provides a machine-readable storage medium, which stores executable instructions, and when the instructions are executed, the machine executes the federated learning defense method as described above.

[0120] Specifically, a system or device equipped with a readable storage medium can be provided. On this readable storage medium, software program codes for implementing the functions of any one of the above embodiments are stored, and the computer or processor of the system or device is made to read and execute the instructions stored in the readable storage medium.

[0121] In this case, the program code read from the readable medium itself can implement the functions of any one of the above embodiments. Therefore, the machine-readable code and the readable storage medium storing the machine-readable code constitute a part of this specification.

[0122] Examples of the readable storage medium include floppy disks, hard disks, magneto-optical disks, optical disks (such as CD-ROM, CD-R, CD-RW, DVD-ROM, DVD-RAM, DVD-RW, DVD-RW), magnetic tapes, non-volatile memory cards, and ROMs. Optionally, the program code can be downloaded from a server computer or a cloud via a communication network.

[0123] Obviously, the above embodiments of the present invention are merely examples for clearly illustrating the technical solutions of the present invention, rather than limitations on the specific implementation manners of the present invention. Any modifications, equivalent replacements, improvements, etc. made within the spirit and principle of the claims of the present invention shall be included within the protection scope of the claims of the present invention.

Claims

1. A federated learning defense method, characterized in that, The method includes: S1. Server - side initial work: S11. Construct a historical global model storage pool on the server - side, and randomly initialize a global model as the initial global model and put it into the historical global model storage pool; S12. Send the initial global model as both the teacher model and the global model to the clients randomly selected in the first round to participate in training at the same time; S2. Client - side training: S21. After the client receives the global model and the teacher model sent by the server - side, the client conducts local training on the global model based on the local dataset to generate a client model and produce a local loss. At the same time, the client conducts knowledge distillation training on the teacher model based on the local dataset to generate a student model and produce a knowledge distillation loss; S22. Jointly train the client model and the student model generated in step S21 to generate an updated local model, and send the updated local model back to the server - side to participate in the aggregation of the global model in the next round; S3. Server - side aggregation: S31. After the server - side receives the updated local model described in step S22, identify and eliminate abnormal local models through the K - means aggregation method, select the local models in the largest cluster for global model aggregation to generate a new round of global model, and store it in the historical global model storage pool; S32. Aggregate the global models in the historical global model storage pool to obtain a new round of teacher model; S4. Send the new round of global model generated in step S31 and the new round of teacher model generated in step S32 to the clients selected for training in the next round at the same time, and loop steps S2 and S3 until the specified number of rounds is reached or the global model converges to end the training.

2. The federated learning defense method according to claim 1, wherein, Specifically, in step S21, when the client conducts local training on the global model based on the local dataset to generate a client model and produce a local loss: The client conducts training based on its local dataset, fine - tunes the global model, and the local loss calculation formula (1) generated by the difference between the client model and the global model is as follows: (1) Among them refers to the local loss of the k round client, is the k local dataset of the client, refers to the global model sent by the server in the round, is the loss function, and represent the input and output vectors respectively, represents the model output calculated using the global model represents the expected value of the random variable pair sampled from the distribution for 3. The federated learning defense method according to claim 1, characterized in that Specifically, in step S21, when the client conducts knowledge distillation training on the teacher model based on the local dataset to generate a student model and produce a knowledge distillation loss: The client transfers the knowledge in the teacher model to the student model in the way of knowledge distillation, and realizes it by minimizing the difference between the outputs of the student model and the teacher model. This process will produce a knowledge distillation loss, and optimize the student model by minimizing this loss; (2) Among them refers to the round of client knowledge distillation loss, refers to the teacher model sent by the server in the round, which is the distillation loss based on the Kullback-Leibler divergence, measuring the difference between the probability distributions of the outputs of the teacher model and the student model, 、 respectively represent the input and output vectors; represents the model output calculated using the teacher model and represents the expected value of the random variable pair sampled from the distribution for​ 4. The federated learning defense method according to claim 1, wherein In step S22, when jointly training the client model and the student model generated in step S21, first generate a local model loss, and calculate the updated local model through the local model loss.

5. The federated learning defense method according to claim 4, wherein The calculation formula of the local model loss is as follows: (3) Among them refers to the local model loss calculated by the client in the round, is a weighted coefficient that controls the balance between the local loss and the distillation loss, is the local loss of the client in the round, and is the knowledge distillation loss of the client in the round; Specifically, for generating the updated local model: (4) Among them represents the learning rate, is the local model of the round of clients, that is, the updated local model, is the local model of the round of clients, represents the gradient of the local model loss with respect to the local model gradient.

6. The federated learning defense method according to claim 1, wherein Specifically, in step S31, when identifying and eliminating abnormal local models through the K - means aggregation method and selecting the local models in the largest cluster for global model aggregation to generate a new round of global model: After the server receives the updated local model uploaded by the client in step S22, it minimizes the clustering loss through the K-means aggregation method while identifying and removing abnormal local models: (5) Among them refers to the total loss function of K-means represents the number of clusters in the clustering. The number of attackers is less than that of benign clients. Therefore, the local models in the largest cluster are selected for global model aggregation; represents a cluster the local model in , represents the t round of local model of the client ; represents the centroid of the cluster ; represents the L2 norm, which is used to measure the distance between the local model and the clustering centroid . This formula makes similar models be grouped into the same cluster, reducing the probability of malicious clients participating in training; The newly generated global model : (6) Among them refers to the round of the global model, indicating the number of clients in the largest cluster participating in this round of training, refers to the round of the client local model.

7. The federated learning defense method according to claim 1, wherein Specifically, the new round of teacher model is obtained by aggregating the global models in the historical global model storage pool in step S32: The teacher model sent to the client in the round using the FedAvg aggregation method : (7) Among them, refers to the teacher model sent by the server in the round, refers to the th global model in the historical global model storage pool.

8. An electronic device, characterized in that, The electronic device includes: At least one processor; and A memory storing instructions that, when executed by the at least one processor, cause the at least one processor to execute the federated learning defense method according to any one of claims 1-7.

9. A machine-readable storage medium storing executable instructions, characterized in that, When executed, the instructions cause the machine to execute the federated learning defense method according to any one of claims 1-7.

Citation Information

Patent Citations

  • Federal learning poisoning attack defense method

    CN118246009A

  • Personalized federal learning method based on knowledge distillation

    CN118761455A

  • Federal learning backdoor defense method based on attention distillation

    CN115630361A

  • Federal learning privacy protection model training method and system based on knowledge distillation

    CN116957064A