Adaptive Bilateral Distillation Personalized Federated Learning Method Based on Diffusion Model
By using the adaptive bilateral distillation method based on diffusion model and the technology of generating pseudo-data based on the diffusion model in personalized federated learning, the problem of insufficient support for personalized needs and forgotten knowledge in the existing technology is solved, and the balance between personalized performance and global generalization ability and the improvement of model performance is achieved.
Patent Information
- Application Number
- CN202411931128.X
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2024-12-26
- Publication Date
- 2025-06-24
- Estimated Expiration
- 2044-12-26
AI Technical Summary
While ensuring the generalization performance of the model, existing personalized federated learning methods are difficult to effectively support client personalized needs, and the global model may experience the problem of ‘catastrophic forgetting’ during the local update process.
Adaptive bilateral distillation personalized federal learning method based on diffusion model is adopted, and high-quality pseudo-data is generated through mutual distillation between the global model and the local model, combined with the conditional diffusion model, and the aggregated global model is fine-tuned to ensure efficient knowledge transfer and adaptability of the personalized model.
The balance between personalized performance and global generalization capabilities is achieved, and it is suitable for multi-client collaboration scenarios in non-independent and homogeneous data environments, avoiding the "catastrophic forgetting" of the global model and significantly improving the performance of the model.
Smart Images

Figure CN119358708B_ABST
Abstract
Description
Technical Field
[0001] The present invention belongs to the technical field of federated learning, and more specifically, relates to an adaptive bilateral distillation personalized federated learning method based on a diffusion model. Background Art
[0002] Federated learning is a distributed machine learning method, whose initial purpose is to solve the "data silo" problem. However, with the continuous iteration of technology and the sharp increase in the number of users, a single generalized global model has become difficult to meet the users' demand for personalization. To solve this problem, personalized federated learning has emerged. Personalized federated learning optimizes the performance of local models for the specific data distribution and task requirements of each user while sharing global knowledge. Compared with traditional federated learning, it can not only alleviate the data heterogeneity problem, but also significantly improve the performance of local models on specific clients with minimal sacrifice of global generalization ability.
[0003] As an efficient knowledge transfer means, knowledge distillation has been widely applied to personalized federated learning. However, due to privacy restrictions, the server cannot directly access the local data of clients, which makes generating pseudo-data an effective solution. The currently popular pseudo-data generation technology is the generative adversarial network GAN. However, GAN requires adversarial training between the generator Generator and the discriminator Discriminator, and the "mode collapse" phenomenon may occur during this process, that is, the generated samples lack diversity and only cover a small number of similar patterns. To overcome this problem, conditional diffusion models Conditional Diffusion Models have been proposed. This method generates pseudo-data by gradually adding noise to the original data and learning the denoising process, fundamentally solving the mode collapse problem of GAN, and significantly improving the quality and diversity of the generated samples.
[0004] Most existing personalized federated learning methods tend to focus on only one aspect of model personalization and generalization ability, resulting in an imbalance between the two, and may even forget previously learned knowledge. For example, FedDistill achieves information exchange by transferring logits between clients and servers. Although this improves the generalization performance of the global model, it ignores the support for client personalization needs, and the transfer of logits may bring privacy leakage risks. FedAMP achieves adaptive model collaboration aggregation by calculating model similarity between clients. However, its focus is still on the optimization of the global model and lacks effective support for client personalization. FML takes into account the learning of personalized knowledge and global knowledge through mutual distillation of local and global models, but the global model may forget some global knowledge when training with local data, resulting in a decrease in generalization performance. FedProto assumes that there are shared class prototypes between clients, but when the client class distribution is completely different or non-overlapping, the global prototype may not accurately express the global features, resulting in limited generalization ability.
[0005] In order to solve the above problems, some improvement schemes have been proposed in recent years. For example, FedKD achieves a balance between personalization and generalization performance through adaptive mutual distillation of local models and global models. However, it may still forget some global knowledge, thus affecting the overall effect. Therefore, there is an urgent need for a method to achieve more efficient personalization while ensuring the generalization performance of the model as much as possible and avoiding forgetting knowledge.
[0006] Chinese patent document CN117035057A discloses a personalized federated learning method based on model and data distillation, the steps of which are as follows: construct a local model on the client, including a shared encoder and a private decoder. The client trains the local model based on the private data set and uploads the model parameters of the shared encoder to the server. The client calculates the output logits of the public data set based on the local model and uploads the logits to the server. The server performs weighted aggregation on the logits and the shared encoder model parameters of multiple clients to obtain global logits and multiple global encoder model parameters. Each client downloads multiple global encoder models, updates multiple local shared encoder model parameters in the client, downloads global logits, and participates in the training of the decoder in the form of knowledge distillation. However, this method has the following shortcomings: 1. This method relies on the global encoder and global logits for knowledge distillation, but the global model has limitations in capturing the personalized features of the client, especially when the data distribution is not independent and identically distributed, it is difficult to meet the personalized needs. 2. This method does not fully consider the "catastrophic forgetting" problem that may occur in the local update process of the global model, that is, the loss of previous knowledge.
[0007] In view of this, the present invention designs an adaptive bilateral distillation personalized federated learning method based on a diffusion model. Summary of the Invention
[0008] The present invention aims to overcome at least one defect of the above-mentioned prior art, and provides an adaptive bilateral distillation personalized federated learning method based on a diffusion model. By designing a guidance mechanism, mutual distillation is carried out between the global model and the local model, so as to achieve efficient knowledge transfer and enhance the adaptability of the personalized model to the specific data distribution of the client. At the same time, a conditional diffusion model is introduced to generate high-quality pseudo-data, and these pseudo-data are used to fine-tune the aggregated global model. This process not only effectively makes up for the global information that may be lost in the local-global mutual distillation process, but also further optimizes the performance of the global model. By combining mutual distillation and conditional diffusion fine-tuning techniques, the present invention achieves a balance between personalized performance and global generalization ability while protecting data privacy, and is applicable to multi-client collaboration scenarios in a non-independent and identically distributed non-IID data environment.
[0009] The detailed technical solution of the present invention is as follows:
[0010] An adaptive bilateral distillation personalized federated learning method based on a diffusion model, the method comprising:
[0011] S1. The server initializes the global model , and broadcasts and sends the initial global model to each participating client;
[0012] S2. The client receives the global model sent by the server, and the client uses local data to train the received global model and the local local model to obtain local loss and global loss. Then, mutual distillation is carried out on the global model and the local local model using local data, and the local loss and global loss are used to guide the distillation process to obtain a local global model; at the same time, the client uses local data and class vectors to train a local conditional diffusion model to obtain a local local generator;
[0013] S3. The client sends the adjusted local global model and the local local generator to the server;
[0014] S4. The server aggregates the received local global model and the initial global model according to the local data volume of the client to obtain an aggregated global model; then the KL divergence is used to calculate the similarity between the local global model and the aggregated global model;
[0015] S5. The similarity is used to perform weighted aggregation on the received local local generator to obtain a global generator that can generate global pseudo-data;
[0016] S6. The server uses the global pseudo-data generated by the global generator to perform knowledge distillation on the aggregated global model and the historical global model, further optimizing the global model;
[0017] S7. The server re-broadcasts and sends the optimized global model to each participating client, repeating the above steps S1 - S6 until the preset number of rounds is reached and then ending, obtaining the finally fine-tuned global model.
[0018] According to the preference of the present invention, in step S2, the client uses local data to train the received global model and the local local model, and obtains the local loss and the global loss as follows:
[0019] Send the local data into the local local model and the global model respectively for training, and use the cross-entropy loss function to calculate their respective losses during the process:
[0020] (1)
[0021] (2)
[0022] In formulas (1) and (2), represents the local loss, represents the global loss, is the result predicted by the local model, is the result predicted by the global model, is the true label, represents the total number of categories in the client dataset.
[0023] According to the preference of the present invention, in step S2, the global model and the local local model are mutually distilled using local data, and the local loss and the global loss are used to guide the distillation process to obtain the local global model as follows:
[0024] Control the intensity of mutual distillation according to the prediction accuracy of the local local model and the global model on the dataset. If the accuracy is higher, the loss is smaller, and at this time the distillation intensity is smaller. If the accuracy is low, the loss is larger, and at this time the distillation intensity is larger, and more knowledge needs to be learned;
[0025] During the mutual distillation process, the KL divergence is used to control the difference between the prediction distributions of the local local model and the global model, and the purpose of knowledge transfer is achieved by minimizing the difference between the two. This process is shown as follows:
[0026] (3)
[0027] (4)
[0028] In formulas (3) and (4), represents the difference between the local model and the global model, represents the difference between the global model and the local model, is the result predicted by the local model, is the result predicted by the global model, represents the local loss, represents the global loss, and the denominator realizes the purpose of controlling the knowledge transfer intensity through the sum of two cross-entropy losses.
[0029] Preferably according to the present invention, in step S2, the client uses local data and class vectors to train a local conditional diffusion model to obtain a local local generator as follows:
[0030] First, use the conditional diffusion model to perform forward diffusion: use the local data as the target data, and gradually add noise to the original data to make the original data gradually approach pure noise. The formula is as follows:
[0031] (5)
[0032] In formula (5), represents the conditional probability distribution from the original data to in the diffusion model, represents the original data, represents the data after adding noise for time step t to the original data represents the Gaussian distribution, is used to control the size of the noise, represents the cumulative noise attenuation coefficient, and I represents the identity matrix; through this process, the original data gradually becomes pure noise;
[0033] Then, perform reverse diffusion. In the reverse diffusion process, a class embedding vector c is introduced to control the generation of data that conforms to the distribution of the local dataset; subsequently, use the variance and mean of the added noise learned in the forward diffusion process to gradually denoise to achieve the purpose of restoring the original data. The formula is expressed as follows:
[0034] (6)
[0035] (7)
[0036] In formulas (6) and (7), represents the reverse generation of the time t-1 state sample from The conditional probability distribution, where c is the class embedding vector, is the mean of the noise, is the variance of the noise, t represents the time step, represents sampling from the conditional probability distribution to obtain a sample at time state t - 1; Through the reverse diffusion process, the data is gradually restored to the target data conforming to the target distribution;
[0037] The optimization objective of the reverse optimization process is:
[0038] (8)
[0039] In formula (8), represents the mean squared error loss function, is the standard Gaussian noise, is the noise predicted by the model, represents the expected value with respect to the time step t, the original data and the noise ; The conditional diffusion model makes the reverse generation process more accurate by minimizing the distance between the noise predicted by the model and the actually added noise ; Through this optimization process, the conditional diffusion model can gradually learn the reverse diffusion process;
[0040] So far, the local generator training process ends, and a local generator for generating data conforming to the client data distribution is obtained .
[0041] According to the preference of the present invention, step S4 is specifically as follows:
[0042] The server performs global aggregation on the global model according to its local data volume, and the formula is as follows:
[0043] (9)
[0044] In formula (9), is the global model after the round of aggregation, N represents the total number of clients participating in the aggregation, represents the local data volume of client j, |D| represents the total data volume of all clients, represents the global model of client j in the
[0045] After the aggregation is completed, the server needs to calculate the similarity between each local global model and the aggregated global model. Since the model parameters here are all high-dimensional vectors, the cosine similarity needs to be used, and its calculation process is as follows:
[0046] (10)
[0047] In formula (10), is to calculate the initial similarity between the local global model and the aggregated global model, represents the global model parameter vector of client k in the i-th round, represents the global model parameter vector after aggregation in the i-th round, represents the L2 norm, and means to and normalize to a unit vector; then is standardized as shown in formula (11):
[0048] (11)
[0049] In formula (11), is the round client similarity between the local global model and the aggregated global model, satisfies .
[0050] Preferably according to the present invention, step S5 is specifically as follows:
[0051] According to the calculated similarity , weighted aggregation is performed on each local generator, and its calculation formula is as shown in (12):
[0052] (12)
[0053] In formula (12), represents the global generator after aggregation in the i-th round, represents the round client similarity between the local global model and the aggregated global model, represents the local generator of client k in the i-th round.
[0054] Preferably according to the present invention, since the aggregated global model may forget the knowledge learned before after local training, therefore, in order to cope with this situation, it is considered to use the pseudo data generated by the global generator and the rounds of saved global models to perform knowledge distillation fine-tuning on the aggregated global model, which fundamentally avoids the problem of "catastrophic forgetting" of the global model.
[0055] First, use the global generator to generate global pseudo data , then select the previous The global models of the rounds are weighted and aggregated to obtain the aggregated global model parameters of the previous rounds , and the specific formula is shown in (13):
[0056] (13)
[0057] In formula (13), represents the average value of the global model parameters of the previous m rounds, represents the global model parameter vector of the i - j round, m represents selecting the global model parameters of the previous m rounds. To avoid unnecessary computational costs, here , where represents the number of communication rounds between the client and the server;
[0058] Subsequently, based on the obtained and perform knowledge distillation on : First, calculate the probability distributions of the student and teacher models used in the distillation process:
[0059] (14)
[0060] (15)
[0061] In formulas (14) and (15), is the probability distribution of the teacher model after smoothing, is the probability distribution of the student model after smoothing, represents the output logits of the teacher model on the global pseudo - data, represents the output logits of the student model on the global pseudo - data, represents the distillation temperature;
[0062] Next, calculate the KL - divergence between the probability distribution of the student model and the probability distribution of the teacher model :
[0063] (16)
[0064] In formula (16), is used to offset the gradient scaling caused by probability distribution smoothing;
[0065] By minimizing the KL - divergence, update to achieve the fine - tuning effect, and its formula is shown in (17):
[0066] (17)
[0067] In Equation (17), represents the learning rate, which is used to control the step size of each update. represents the loss function gradient;
[0068] Finally, the globally fine-tuned model parameters are broadcast and distributed to each client, and used to update the local globally model parameters. The above process is continuously iterated until the global model converges.
[0069] Compared with the prior art, the beneficial effects of the present invention are as follows:
[0070] (1) Through the adaptive mutual distillation of the local model and the global model, the present invention makes full use of the local loss and the global loss, and guides the global model to achieve a dynamic balance between generalization performance and personalized requirements.
[0071] (2) The present invention uses the pseudo-data generated by the global generator to fine-tune the aggregated global model, effectively compensating for the global knowledge that may be lost during the distillation process, and avoiding the global model forgetting historical knowledge due to the introduction of new knowledge.
[0072] (3) Since there are differences in the data volume, data distribution, etc. of each client, the present invention avoids the influence brought by data heterogeneity by learning personalized models for each client. For the common non-independent and identically distributed data in federated learning, the introduction of the mutual distillation mechanism and the conditional diffusion model of the present invention can better adapt to and process the heterogeneity of the data of each client.
[0073] (4) The present invention uses global pseudo-data for global knowledge distillation, avoiding the risk of privacy leakage caused by directly using private data. BRIEF DESCRIPTION OF THE DRAWINGS
[0074] Figure 1 is a flowchart of the personalized federated learning method of the present invention.
[0075] Figure 2 is a schematic diagram of the framework principle of the personalized federated learning method of the present invention.
[0076] Figure 3 is a comparison chart of the test accuracies of the personalized federated learning algorithm of the present invention and other federated learning algorithms under CIFAR10 data in Embodiment 1 of the present invention.
[0077] Figure 4 is a comparison chart of the test accuracies of the personalized federated learning algorithm of the present invention and other federated learning algorithms under CIFAR100 data in Embodiment 1 of the present invention. DETAILED DESCRIPTION OF THE INVENTION
[0078] The present disclosure will be further described below with reference to the accompanying drawings and embodiments.
[0079] Example 1
[0080] Refer to Figure 1 and Figure 2 , this example provides an adaptive bilateral distillation personalized federated learning method based on a diffusion model, and the method includes:
[0081] S1. The server initializes the global model , and broadcasts and sends the initial global model to each participating client;
[0082] S2. The client receives the global model sent by the server. The client uses local data to train the received global model and the local local model, obtains the local loss and the global loss. Then, the client uses local data to mutually distill the global model and the local local model, and uses the local loss and the global loss to guide the distillation process to obtain the local global model; at the same time, the client uses local data and the category vector to train the local conditional diffusion model to obtain the local local generator;
[0083] S3. The client sends the adjusted local global model and the local local generator to the server;
[0084] S4. The server aggregates the received local global model and the initial global model according to the local data volume of the client to obtain the aggregated global model; then uses the KL divergence to calculate the similarity between the local global model and the aggregated global model;
[0085] S5. Use the similarity to perform weighted aggregation on the received local local generator to obtain a global generator that can generate global pseudo data;
[0086] S6. The server uses the global pseudo data generated by the global generator to perform knowledge distillation on the aggregated global model and the historical global model to further optimize the global model;
[0087] S7. The server re-broadcasts and sends the optimized global model to each participating client, and repeats the above steps S1-S6 until the preset number of rounds is reached and then ends, to obtain the finally fine-tuned global model.
[0088] In step S2, the client uses local data to train the received global model and the local local model, and obtains the local loss and the global loss as follows:
[0089] Send the local data into the local local model and the global model respectively for training, and use the cross-entropy loss function to calculate their respective losses during the process:
[0090] (1)
[0091] (2)
[0092] In formulas (1) and (2), represents the local loss, represents the global loss, is the result predicted by the local model, is the result predicted by the global model, is the true label.
[0093] The method of using local data to mutually distill the global model and the local local model, and using the local loss and the global loss to guide the distillation process to obtain the local global model is as follows:
[0094] Control the intensity of mutual distillation according to the prediction accuracy of the local local model and the global model on the dataset. If the accuracy is higher, the loss is smaller, and the distillation intensity is smaller at this time. If the accuracy is low, the loss is larger, and the distillation intensity is larger at this time, and more knowledge needs to be learned;
[0095] In the process of mutual distillation, KL divergence is used to control the difference between the prediction distributions of the local local model and the global model, and the purpose of knowledge transfer is achieved by minimizing the difference between the two. This process is as follows:
[0096] (3)
[0097] (4)
[0098] In formulas (3) and (4), represents the difference between the local model and the global model, represents the difference between the global model and the local model, is the result predicted by the local model, is the result predicted by the global model, represents the local loss, represents the global loss. The denominator achieves the purpose of controlling the knowledge transfer intensity through the sum of two cross-entropy losses.
[0099] The client uses local data and class vectors to train a local conditional diffusion model to obtain the local local generator as follows:
[0100] First, use the conditional diffusion model for forward diffusion: Use the local data as the target data, and gradually add noise to the original data to make the original data gradually approach pure noise. The formula is as follows:
[0101] (5)
[0102] In formula (5), represents the conditional probability distribution from the original data to in the diffusion model, represents the original data, represents the data after adding noise for the time step to the original data , represents the Gaussian distribution, which is used to control the magnitude of the noise, represents the cumulative noise attenuation coefficient, and I represents the identity matrix; through this process, the original data gradually becomes pure noise;
[0103] Then, reverse diffusion is performed. During the reverse diffusion process, a class embedding vector c is introduced to control the generation of data that conforms to the distribution of the local dataset; subsequently, the variance and mean of the added noise learned during the forward diffusion process are used to gradually denoise to achieve the purpose of restoring the original data. The formula is as follows:
[0104] (6)
[0105] (7)
[0106] In formulas (6) and (7), represents the conditional probability distribution of reverse generating the state sample at time t - 1 from , c is the class embedding vector, is the mean of the noise, is the variance of the noise, t represents the time step, represents sampling from the conditional probability distribution to obtain the sample at time state t - 1 ; through the reverse diffusion process, the data is gradually restored to the target data that conforms to the target distribution; ;
[0107] The optimization objective of the reverse optimization process is:
[0108] (8)
[0109] In formula (8), represents the mean squared error loss function, is the standard Gaussian noise, is the noise predicted by the model, represents for the time step t, the original data and the noise Expected value; The conditional diffusion model minimizes the noise predicted by the model and the actually added noise to make the reverse generation process more accurate; Through this optimization process, the conditional diffusion model can gradually learn the reverse diffusion process;
[0110] At this point, the local generator training process ends, and a local generator for generating data conforming to the client data distribution is obtained .
[0111] Step S4 is specifically as follows:
[0112] The server performs global aggregation on the global model according to its local data volume, and the formula is as follows:
[0113] (9)
[0114] In formula (9), is the global model after the round of aggregation, N represents the total number of clients participating in the aggregation, represents the local data volume of client j, |D| represents the total data volume of all clients, represents the global model of client j in the i-th round;
[0115] After the aggregation is completed, the server needs to calculate the similarity between each local global model and the aggregated global model. Since the model parameters here are all high-dimensional vectors, the cosine similarity needs to be used, and its calculation process is as follows:
[0116] (10)
[0117] In formula (10), is the preliminary similarity between the local global model and the aggregated global model, represents the global model parameter vector of client k in the i-th round, represents the global model parameter vector after the i-th round of aggregation, represents the L2 norm, and mean to and normalize to unit vectors; Subsequently, is standardized as shown in formula (11):
[0118] (11)
[0119] In formula (11), is the round of client The similarity between the local global model and the aggregated global model satisfies .
[0120] Step S5 is specifically as follows:
[0121] According to the calculated similarity , perform weighted aggregation on each local generator, and its calculation formula is as shown in (12):
[0122] (12)
[0123] In formula (12), represents the global generator after the i-th round of aggregation, represents the round of the client The similarity between the local global model and the aggregated global model represents the local generator of the k-th client in the i-th round.
[0124] According to the preference of the present invention, since the aggregated global model may forget the knowledge learned before after local training, therefore, in order to cope with this situation, consider using the pseudo-data generated by the global generator and the rounds of the saved global models to perform knowledge distillation fine-tuning on the aggregated global model, which fundamentally avoids the problem of "catastrophic forgetting" of the global model.
[0125] First, use the global generator to generate global pseudo-data , then select the global models of the previous rounds for weighted aggregation to obtain the parameters of the global model of the previous rounds after aggregation , and the specific formula is as follows:
[0126] (13)
[0127] In formula (13), represents the average value of the global model parameters of the previous m rounds, represents the global model parameter vector of the i-j-th round, m represents the global model parameters of the previous m rounds selected, and in order to avoid unnecessary computational costs, here , where represents the number of communication rounds between the client and the server;
[0128] Subsequently, according to the obtained and perform knowledge distillation on : First, calculate the probability distributions of the student and teacher models used in the distillation process:
[0129] (14)
[0130] (15)
[0131] In equations (14) and (15), is the probability distribution of the teacher model after smoothing, is the probability distribution of the student model after smoothing, represents the output logits of the teacher model on the global pseudo-data, represents the output logits of the student model on the global pseudo-data, represents the distillation temperature;
[0132] Next, calculate the KL divergence between the probability distribution of the student model and the probability distribution of the teacher model :
[0133] (16)
[0134] In equation (16), is used to offset the gradient scaling caused by the smoothing of the probability distribution;
[0135] By minimizing the KL divergence, update to achieve the effect of fine-tuning. The formula is as follows:
[0136] (17)
[0137] In equation (17), represents the learning rate, which is used to control the step size of each update, represents the loss function of the gradient;
[0138] Finally, broadcast and distribute the fine-tuned global model parameters to each client, and use them to update the local global model parameters. Continuously iterate the above process until the global model converges.
[0139] The optimization objective of the present invention is as follows:
[0140] The present invention selects clients to participate in the aggregation. Each client owns its local dataset as , where | | represents the number of datasets it owns. The global dataset represents the set of all client datasets. The main local objective of the invention is to simultaneously obtain a local model and a global model , whose optimization goal is to minimize the total loss value on the global dataset and , and their calculation formulas are as follows respectively:
[0141] (18)
[0142] Among them,
[0143] (19)
[0144] Among them,
[0145] represents the local model loss of the -th client, is the cross-entropy loss, which is used to measure the difference between the predicted value and the true value. represents the global model loss of the -th client. Here, and are the KL divergences after weighting the local model prediction distribution and the global model prediction distribution.
[0146] The adaptive bilateral distillation personalized federated learning method based on the diffusion model provided by the present invention can be applied to the scenario of personalized blood glucose prediction and management for diabetic patients. High-quality simulated data is generated through the diffusion model to solve the problem of insufficient individual data, and the global knowledge and local characteristics are combined through the bilateral distillation mechanism to generate a high-precision personalized prediction model for each patient. In practice, the intelligent devices of patients collect data such as blood glucose, exercise, and diet in real time, and protect data privacy through the federated learning framework to avoid uploading sensitive information. At the same time, the model can dynamically adapt according to the living habits and disease conditions of patients, and provide real-time blood glucose trend prediction and personalized management suggestions for patients, such as diet, exercise, and drug adjustment plans.
[0147] Experimental examples,
[0148] Next, experimental verification is carried out. First, a basic introduction to the experiment is given:
[0149] (1) Datasets: CIFAR10, CIFAR100;
[0150] (2) Model: CNN;
[0151] (3) Data partitioning method: practical non-IID Dirichlet partitioning method;
[0152] (4) Baseline algorithms: FedAvg, FedProx, FML, DaFKD, FedKD, FedTweet;
[0153] To ensure fairness, all baseline algorithms and this algorithm use the same network architecture, devices, and hyperparameter settings. Next, the Figure 3 and Figure 4 experimental results will be analyzed in detail.
[0154] For Figure 3 , this part compares the performance of the proposed personalized federated learning algorithm with other baseline algorithms on the CIFAR10 dataset under the experimental condition of heterogeneous coefficient α = 0.1. It can be seen from the figure that the accuracies of the personalized algorithms (such as FML, DaFKD, FedKD, FedFD, FedTweet) are significantly higher than those of the traditional federated learning algorithms FedAvg and FedProx. This shows that the effectiveness of the personalized federated learning scheme has been verified in the data heterogeneous scenario. At the same time, among various personalized federated learning algorithms, the algorithm FedFDM of the present invention has a better accuracy than other algorithms after convergence, demonstrating its significant performance advantages.
[0155] For Figure 4 , the dataset is changed from CIFAR10 to CIFAR100 while keeping other experimental configurations unchanged. It can be seen from the experimental results that FedFDM still shows significantly better performance than other algorithms in the case of data heterogeneity, further verifying its generalization ability and applicability.
[0156] The above analysis shows that the algorithm FedFDM of the present invention has significant advantages in the personalized federated learning scenario with data heterogeneity, and its key modules play a crucial role in improving the overall performance.
Claims
1. Adaptive bilateral distillation personalized federated learning method based on diffusion model, characterized by: The method comprises: S1. Server initializes the global model , broadcast the initial global model to each participating client; S2, the client receives the global model from the server, and trains the received global model and the local local model using local data to obtain local loss and global loss, then uses local data to distill the global model and the local local model, and uses local loss and global loss to guide the distillation process to obtain a local global model; at the same time, the client uses local data and category vectors to train a local conditional diffusion model to obtain a local local generator; wherein the local data is image data; S3, the client sends the adjusted local global model and the local local generator to the server; S4, the server aggregates the received local global models according to the amount of local data of the client to obtain an aggregated global model; then uses cosine similarity to calculate the similarity between the local global model and the aggregated global model; wherein the local data is image data; The specific steps of S4 are as follows: The server performs global aggregation on the global model based on its local data volume. The formula is as follows: (1) In formula (1), It is The global model after round aggregation, N represents the total number of clients participating in the aggregation, represents the amount of local data of client j, Represents the total amount of data from all clients. represents the global model of client j in round i; After the aggregation is completed, the server uses cosine similarity to calculate the similarity between each local global model and the aggregated global model. The calculation process is as follows: (2) In formula (2), It is to calculate the preliminary similarity between the local global model and the aggregated global model. represents the global model parameter vector of client k in round i, represents the global model parameter vector after the i-th round of aggregation, and Indicates that and Normalized to unit vector; Then will Perform standardization, as shown in formula (3): (3) In formula (3), It is Round Client The similarity between the local global model and the aggregated global model, satisfy ; S5. Use similarity to perform weighted aggregation on the received local local generators to obtain a global generator that can generate global pseudo data. The specific steps are as follows: Based on the calculated similarity , weighted aggregation is performed on each local generator, and the calculation formula is shown in (4): (4) In formula (4), represents the global generator after the i-th round of aggregation, Representative Round Client The similarity between the local global model and the aggregated global model, represents the local generator of client k in round i; S6. The server uses the global pseudo data generated by the global generator to perform knowledge distillation on the aggregated global model and the historical global model to further optimize the global model. First, use the global generator Generate global pseudo data , then select Previous The global model of the round is weighted aggregated to obtain the aggregated front Global model parameters of the wheel , the specific formula is shown in (5): (5) In formula (5), represents the average value of the global model parameters in the first m rounds, represents the global model parameter vector of the ijth round, m represents the global model parameters of the first m rounds, and ,in, Represents the number of communication rounds between the client and the server; Then, according to the obtained and right Perform knowledge distillation: First, calculate the probability distribution of the student and teacher models used in the distillation process: (6) (7) In formulas (6) and (7), is the probability distribution of the teacher model after smoothing, is the probability distribution of the student model after smoothing, represents the output logits of the teacher model on the global pseudo data, represents the output logits of the student model on the global pseudo data, represents the distillation temperature; Next, calculate the KL divergence between the probability distribution of the student model and the probability distribution of the teacher model : (8) In formula (8), Used to offset the gradient scaling caused by smoothing of probability distribution; By minimizing the KL divergence, we update , thus achieving the effect of fine-tuning, the formula is shown in (9): (9) In formula (9), Represents the learning rate, which is used to control the step size of each update. Represents the loss function The gradient of S7. The server rebroadcasts the optimized global model to each participating client, uses it to update the local global model parameters, and repeats the above steps S1-S6 until the preset rounds are reached to obtain the final fine-tuned global model, which is a CNN model for processing image data.
2. The adaptive bilateral distillation personalized federated learning method based on the diffusion model according to claim 1, characterized in that: In step S2, the client trains the received global model and the local local model using local data, and obtains the local loss and the global loss as follows: Local data They are sent to the local model and the global model for training respectively, and the cross entropy loss function is used to calculate their respective losses in the process: (10) (11) In formulas (10) and (11), Indicates local loss, represents the global loss, is the result predicted by the local model, is the result predicted by the global model, Is a real label, Represents the total number of categories in the client dataset.
3. The adaptive bilateral distillation personalized federated learning method based on the diffusion model according to claim 1, characterized in that: In step S2, the global model and the local local model are mutually distilled using local data, and the local loss and the global loss are used to guide the distillation process, and the local global model is obtained as follows: The strength of mutual distillation is controlled according to the accuracy of the local model and the global model in predicting the data set. If the accuracy is higher, the loss is smaller, and the distillation strength is smaller. If the accuracy is lower, the loss is greater, and the distillation strength is greater. In the process of mutual distillation, KL divergence is used to control the difference between the local model prediction distribution and the global model prediction distribution, and the purpose of knowledge transfer is achieved by minimizing the difference between the two. The process is as follows: (12) (13) In formulas (12) and (13), represents the difference between the local model and the global model, represents the difference between the global model and the local model, is the result predicted by the local model, is the result predicted by the global model, Indicates local loss, represents the global loss.
4. The adaptive bilateral distillation personalized federated learning method based on the diffusion model according to claim 1, characterized in that: In step S2, the client uses local data and category vectors to train a local conditional diffusion model to obtain a local local generator as follows: First, use the conditional diffusion model to perform forward diffusion: convert the local data As the target data, by gradually adding noise to the original data, the original data gradually approaches pure noise. The formula is as follows: (14) In formula (14), In the diffusion model, the original data arrive The conditional probability distribution of Represents the original data, Represents the original data The data after the noise addition time step t, represents a Gaussian distribution, To control the noise level, represents the cumulative noise attenuation coefficient, I represents the unit matrix; Then, reverse diffusion is performed. In the reverse diffusion process, the category embedding vector c is introduced to control the generation of data that conforms to the distribution of the local data set; then, the variance and mean of the added noise learned in the forward diffusion process are used to gradually denoise to restore the original data. The formula is as follows: (15) (16) In equations (15) and (16), Indicates from Reverse generation of state samples at time t-1 The conditional probability distribution of , c is the category embedding vector, is the mean of the noise, is the variance of the noise, t represents the time step, Represents the conditional probability distribution The sample at time state t-1 is obtained by sampling ; Through the reverse diffusion process, the data is gradually restored to the target data that conforms to the target distribution; The optimization goal of the reverse optimization process is: (17) In formula (17), represents the mean square error loss function, is standard Gaussian noise, is the noise predicted by the model, Represents the time step t, the original data and noise The expected value of ; through this optimization process, the conditional diffusion model can gradually learn the reverse diffusion process; At this point, the local generator training process is completed, and the local generator used to generate data that matches the client data distribution is obtained. .
Citation Information
Patent Citations
Personalized federal learning method based on model and data distillation
CN117035057A
Federal learning algorithm based on diffusion model and weight adaptive knowledge distillation
CN116665000A