Heterogeneous federal learning framework and method based on multi-knowledge distillation fusion

By introducing a multi-knowledge distillation fusion framework in heterogeneous federated learning, and adopting temperature adaptive and batch sample correlation knowledge distillation methods, the shortcomings of heterogeneity and data diversity in the prior art are solved, and more efficient model training and lower communication overhead are achieved.

CN119940476AActive Publication Date: 2025-05-06NANJING DAKANG AUTOMATION TECHNOLOGY CO LTD

Patent Information

Application Number
CN202411975676.2
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2024-12-30
Publication Date
2025-05-06
Estimated Expiration
2044-12-30

AI Technical Summary

Technical Problem

When using knowledge distillation technology, existing heterogeneous federated learning methods lack sufficient consideration of client heterogeneity and data diversity, resulting in poor effectiveness of the model in practical applications.

Method used

A heterogeneous federal learning framework based on multi-knowledge distillation fusion is proposed. By introducing temperature adaptive knowledge distillation and batch sample category correlation knowledge distillation methods, global class particle size knowledge is generated and weighted aggregation is performed, and local model parameters are updated in combination with cross entropy loss.

Benefits of technology

It effectively improves the adaptability and performance of the model in a heterogeneous environment, achieves higher accuracy and lower communication overhead, and alleviates the problems of large differences in equipment performance, high communication overhead and uneven data distribution.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN119940476A_ABST
    Figure CN119940476A_ABST
Patent Text Reader

Abstract

The invention relates to the technical field of machine learning, in particular to a heterogeneous federal learning framework and method based on multi-knowledge distillation fusion, and the method comprises the steps that a client side uses local data to calculate local knowledge; the server receives the class granularity knowledge uploaded by the client, and performs weighted aggregation on all the client class granularity knowledge through a class granularity knowledge aggregation module to generate global class granularity knowledge; the server issues global knowledge to each client through a knowledge distribution module as teacher knowledge for subsequent training of the client model; and the client receives global knowledge, temperature self-adaptive knowledge distillation is performed through the knowledge distillation module, and the effect of teacher and student knowledge transmission is maximized. According to the method, two different knowledge distillation methods, namely temperature self-adaptive knowledge distillation and batch sample category relevance knowledge distillation, are provided and fused, so that many challenges faced by a traditional federal learning framework in processing heterogeneous equipment and heterogeneous data environments are effectively relieved.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention relates to the field of federated learning and knowledge distillation technology, and in particular to a heterogeneous federated learning framework and method based on multi-knowledge distillation fusion. Background Art

[0002] Federated learning is a distributed machine learning method that aims to train models with multiple participants while protecting data privacy. Traditional centralized machine learning requires data to be aggregated to a central server for training, which poses challenges in terms of data privacy, security, and data transmission costs. Federated learning trains models locally on each participant and only shares model parameters or gradient information, avoiding direct transmission of raw data, thereby effectively protecting data privacy. However, federated learning faces the problem of device and data heterogeneity in practical applications. Different clients may have different computing power, network bandwidth, and data distribution, which makes it difficult for a unified model architecture to adapt to the needs of all clients. In addition, the non-independent and identically distributed characteristics of data will also affect the convergence speed and accuracy of the model. In order to solve the above problems, knowledge distillation technology is introduced into federated learning. Knowledge distillation is a method of model compression and acceleration. By transferring the knowledge of a complex model (teacher model) to a simple model (student model), the student model can obtain performance similar to that of the teacher model while maintaining a small scale. In federated learning, knowledge distillation can be used to aggregate knowledge from different clients on the server side and transfer it to each client, thereby achieving model personalization and adaptability. However, existing heterogeneous federated learning methods often lack sufficient consideration of client heterogeneity and data diversity when using knowledge distillation technology, resulting in poor performance of the model in practical applications. Therefore, a heterogeneous federated learning method is urgently needed to improve the adaptability and performance of the model in heterogeneous environments.

[0003] CN116227624A discloses a federated knowledge distillation method and system for heterogeneous models, aiming to solve the problems of non-independent and identically distributed data and model heterogeneity. By introducing domain classifier self-supervised learning to extract domain data sets from open data sets, combined with knowledge distillation of global models and local models, the model is updated iteratively in rounds to achieve performance optimization of global models and local models in heterogeneous environments. This method weakens the dependence on open data sets, and improves the robustness and generalization ability of the model by utilizing intermediate layer features and domain data. However, this method has high design and training complexity for domain classifiers and global models, and requires large computing and communication resources.

[0004] CN118153666A discloses a personalized federated knowledge distillation model construction method (Fed-PKD). Aiming at the heterogeneity of clients and data in the Internet of Things, a federated learning framework combining personalized model construction, a two-paradigm weight aggregation algorithm, and a federated knowledge distillation strategy is designed to improve the generalization ability of the model, accelerate the model convergence speed, and reduce communication overhead. However, this method requires a public data set, and the quality of the public data set selection directly affects the effect of knowledge distillation. If the client data is relatively private or lacks high-quality public data samples, knowledge transfer may fail.

[0005] CN116629376A discloses a federated learning aggregation method and system (FedDTG) based on data-free distillation. The core of the method is to use distributed generative adversarial networks (GANs) and knowledge distillation technology to solve the problems of not supporting model heterogeneity, privacy leakage, and public data set dependency in traditional federated learning methods. However, this method relies on generator training, and the quality of the generator is the key to the entire method. However, when the data is small or the distribution is extremely uneven, the training of the generator may be difficult to achieve the ideal effect, which directly affects the quality of distillation. Secondly, the computational and communication overheads increase. The client needs to perform three-party adversarial training and upload the generator and discriminator parameters multiple times, which may bring higher computational costs and communication overheads, especially on resource-constrained devices. Summary of the invention

[0006] The technical problem to be solved by the present invention is to overcome the shortcomings of the above-mentioned prior art and provide a heterogeneous federated learning framework and method based on multi-knowledge distillation fusion.

[0007] The technical solution adopted to solve the above technical problems is: a heterogeneous federated learning framework based on multi-knowledge distillation fusion, including: an edge server and at least one client device;

[0008] Each client k∈{1, 2, 3, ...K} has its local dataset, where the dataset is expressed as follows:

[0009]

[0010] And there is |D k | samples, each of which has D data dimensions and belongs to one of C different classes. Due to different user behaviors, the local training and test data sets between clients are not independent and identically distributed. The personalized model parameters of client k are recorded as where d k represents the number of parameters in the client k model. There is system heterogeneity between client devices, and the scale of models deployed on the client is different. For example, d l ≠d m , Each client k has a local optimization goal, which depends on its corresponding local data distribution, and the overall goal is to minimize the expected goal of all clients, which is expressed as:

[0011]

[0012] The edge server is configured with: a class granularity knowledge aggregation module, which is used to receive class granularity knowledge from multiple clients and generate global class granularity knowledge; a knowledge distribution module, which is used to send the global class granularity knowledge to each client; the client is configured with: a local knowledge calculation module, which is used to calculate local class granularity knowledge based on local data; a knowledge distillation module, which is used to perform knowledge distillation based on the received global class granularity knowledge and local class granularity knowledge.

[0013] The technical solution adopted to solve the above technical problems is: a heterogeneous federated learning method based on multi-knowledge distillation fusion, which is applicable to the heterogeneous federated learning framework based on multi-knowledge distillation fusion, including:

[0014] Step 1: The client uses local data to calculate local knowledge. For example, the first client input is Represents the sample data of the u-th Batchsize of client 1, through the predicted label and the true label y u The error between them is used to calculate the cross entropy loss L ce At the same time, each client calculates the local class granularity knowledge based on local data through the local knowledge calculation module according to the formula Calculate and aggregate into a kind of granular knowledge;

[0015] Step 2: The server receives the class granularity knowledge uploaded by the client, and generates global class granularity knowledge by weighted aggregation of all client class granularity knowledge through the class granularity knowledge aggregation module;

[0016] Step 3: The server distributes global knowledge to Distribute to each client as teacher knowledge for subsequent training of the client model;

[0017] Step 4: The client receives global knowledge and performs temperature-adaptive knowledge distillation through the knowledge distillation module to maximize the effect of knowledge transfer between teachers and students;

[0018] Step 5: The client receives global knowledge and performs batch sample correlation knowledge distillation through the knowledge distillation module. By introducing batch-level sample correlation distillation loss, the model's excessive bias towards specific samples or categories is suppressed, and the model's understanding of the overall characteristics of the data is strengthened;

[0019] Step 6: The client updates the local model and temperature prediction model parameters based on the knowledge distillation loss and cross entropy loss.

[0020] Preferably, in step 2, a knowledge aggregation scheme that takes into account differences in knowledge clarity is adopted, which specifically includes the following steps:

[0021] Step 2.1: Calculate the clarity of the teacher and student knowledge respectively, using the logarithmic sum function to quantify the smoothness of the output, denoted as clarity. Assume Where C represents the probability of the sample belonging to the category, and clarity is defined as follows:

[0022]

[0023] Step 2.2: Different aggregation weights are assigned according to the clarity of the knowledge differentiation. Specifically, the definition of the knowledge weight aggregation for class j is as follows:

[0024]

[0025] Among them, Sharpness j represents the sum of the clarity of all knowledge belonging to category j, represents the global knowledge of class j, Represents the set of clients that have samples of type j.

[0026] Preferably, the temperature adaptive knowledge distillation in step 4 specifically includes the following steps:

[0027] Step 4.1: For each client, a separate lightweight model is used to learn a dynamic temperature prediction module θ k ,The training starts by first inversely optimizing the module model parameters of the ,temperature prediction module with the aim of maximizing the distillation loss between the ,student and the teacher;

[0028] Step 4.2: Pass the student logits and teacher logits as input to the temperature prediction module to predict the distillation temperature T suitable for the current sample. The prediction process can be expressed as:

[0029]

[0030] in, For the relationship mapping of the temperature prediction module, in order to ensure the predicted temperature value T pred Within a reasonable range, the predicted value is mapped to the preset temperature range [T start ,T end ], the specific formula is as follows:

[0031] T pred =T start +T end δ(T pred )

[0032] Where T start and T end is the set temperature range, and δ(·) is the activation function, which aims to map the model prediction value to between 0 and 1.

[0033] Step 4.3: Standardize the teacher knowledge and student knowledge to ensure that the student model can learn knowledge more effectively rather than mechanically imitate, thereby reducing the distillation loss in the case of accurate prediction and making the model more balanced in the performance of all categories during the learning process. According to the formula Get the knowledge after the teachers and students have standardized the process. is the mean value of students’ or teachers’ knowledge, and σ(Z) is the corresponding standard deviation;

[0034] Step 4.4: According to the different clarity of the standardized teacher and student knowledge, obtain the teacher-student differential distillation temperature. Specifically, the calculation method of the clarity sharp is as follows:

[0035]

[0036] Among them, Z is the student or teacher knowledge. Compared with student knowledge, teacher knowledge with relatively large clarity should have a higher temperature, while knowledge with less clarity should have a lower temperature. The distillation temperature between teachers and students is dynamically adjusted for each sample to better reflect the characteristics of the sample. This adaptive adjustment mechanism makes up for the defect of standardization lacking sample-level characteristic considerations, making knowledge transfer more targeted. The distillation temperature of teachers and students is adjusted in the following ways:

[0037]

[0038]

[0039] Where T s and T t Differentiated temperature for knowledge distillation and T s ≠T t ,This step further forms personalized knowledge transfer at the sample level, by adaptively generating appropriate temperature values ​​for each sample and category, and fine-tuning the temperature according to sample characteristics during the knowledge transfer process, maximizing the absorption of teacher knowledge by the student model;

[0040] Step 4.5: Calculate each element c∈[C] in the knowledge vector at temperature T according to the teacher-student differentiation temperature sor T t The transformation mapping under the action of the student model output knowledge as an example, each element c∈[C] in the knowledge vector is at temperature T s The transformation mapping under action is:

[0041]

[0042] Among them, C is the category that the sample may belong to;

[0043] Step 4.6: Based on the teacher-student knowledge, the temperature T is the temperature at which the distillation temperature is differentiated. s and T t Transformation mapping under action and Calculate its temperature adaptive knowledge distillation loss, the calculation formula is as follows:

[0044]

[0045] Preferably, in step 5, a knowledge distillation of the correlation of samples within a batch specifically includes the following steps:

[0046] Step 5.1: Uniformly quantify the performance of each sample in a batch in a category, so that the model considers the influence of other samples while considering the prediction of a single sample, so as to prevent the model from being overly biased towards certain specific samples or categories. Formally, the normalized logits values ​​of the teacher and student models for sample i in category c∈[C] are expressed as:

[0047]

[0048] Where B is the batch size, T is the set distillation temperature, and are the logits values ​​of teacher knowledge and student knowledge for sample i on category c respectively;

[0049] Step 5.2: Taking the normalized knowledge in the category direction as input, the calculation method of the knowledge distillation of the sample association within the batch is as follows:

[0050]

[0051] Batch class normalization introduces an implicit regularization effect, which prevents the model from overfitting the prediction of a single sample during training. In the loss function of batch class normalization, the output of the student model is is the logits normalization function of all samples in the batch on category c, so the loss function is about the student model logits The gradient of is:

[0052]

[0053] By using the chain rule, we can get the formula By taking partial derivatives, we can get the gradient formula:

[0054]

[0055] The gradient of sample i on category c depends not only on its own probability distribution, but also on the probability distribution of other samples in the batch. represents Kronecker delta, which is 1 when i = b and 0 otherwise. This gradient formula shows that the intra-batch category normalization introduces mutual constraints between samples in the batch, thereby limiting the gradient update amplitude of a single sample.

[0056] Preferably, step 6 specifically includes the following steps:

[0057] Step 6.1: According to the temperature-adaptive knowledge distillation method, obtain the temperature-adaptive knowledge distillation loss L TAKD ;

[0058] Step 6.2: Obtain the temperature-adaptive knowledge distillation loss L according to the knowledge distillation method of the sample association within the batch CRKD ;

[0059] Step 6.3: Combine these two loss functions and the cross entropy loss to form the final distillation loss function:

[0060] Loss = λ·L ce +α·L TAKD +β·L CRKD

[0061] Among them, λ, α, and β are the weight parameters of the three loss components respectively. The overall fusion strategy improves the robustness and generalization ability of the model by balancing the contribution of each loss function;

[0062] Step 6.4: The client calculates the total loss Loss relative to the local model parameters W through the back propagation algorithm k and the temperature prediction module parameter θ k Gradient And use the stochastic gradient descent optimization algorithm to update the model parameters. The formula is as follows:

[0063]

[0064] Among them, η is the learning rate of the local model, η T is the learning rate of the temperature prediction module parameters, and are the gradients of the total loss function with respect to the local model parameters and the temperature prediction module parameters, respectively.

[0065] The beneficial effects of the present invention are as follows: The present invention proposes and integrates two different knowledge distillation methods, namely temperature-adaptive knowledge distillation and batch sample category correlation knowledge distillation, to achieve higher accuracy with minimal communication overhead, effectively alleviating the many challenges faced by traditional federated learning frameworks when dealing with heterogeneous devices and heterogeneous data environments, especially in the case of large differences in device performance, high communication overhead and uneven data distribution, the accuracy and efficiency of client model training are significantly limited. BRIEF DESCRIPTION OF THE DRAWINGS

[0066] Figure 1 This is a heterogeneous federated learning framework diagram of multi-knowledge distillation fusion in the present invention.

[0067] Figure 2 It is a schematic diagram of the temperature-adaptive knowledge distillation method under the heterogeneous federated learning framework of multi-knowledge distillation fusion of the present invention.

[0068] Figure 3 It is a schematic diagram of the batch sample category correlation knowledge distillation method under the heterogeneous federated learning framework of multi-knowledge distillation fusion of the present invention.

[0069] Figure 4 This is a comparison chart of the effects of the heterogeneous federated learning method of multi-knowledge distillation fusion of the present invention and other federated learning methods. DETAILED DESCRIPTION

[0070] Embodiment 1, as Figure 1 As shown, the present invention proposes a heterogeneous federated learning framework based on multi-knowledge distillation fusion, including: an edge server and at least one client device:

[0071] Each client k∈{1, 2, 3, …K} has its local dataset, where the dataset is expressed as follows:

[0072]

[0073] And there is |D k | samples, each of which has D data dimensions and belongs to one of C different classes. Due to different user behaviors, the local training and test data sets between clients are not independent and identically distributed. The personalized model parameters of client k are recorded as where d k represents the number of parameters in the client k model. There is system heterogeneity between client devices, and the scale of models deployed on the client is different. For example, d l ≠d m , Each client k has a local optimization goal, which depends on its corresponding local data distribution, and the overall goal is to minimize the expected goal of all clients, which is expressed as:

[0074]

[0075] The edge server is configured with: a class granularity knowledge aggregation module, which is used to receive class granularity knowledge from multiple clients and generate global class granularity knowledge; a knowledge distribution module, which is used to send the global class granularity knowledge to each client; the client is configured with: a local knowledge calculation module, which is used to calculate local class granularity knowledge based on local data; a knowledge distillation module, which is used to perform knowledge distillation based on the received global class granularity knowledge and local class granularity knowledge.

[0076] Embodiment 2, as Figure 2-Figure 3 As shown, the present invention proposes a heterogeneous federated learning method based on multi-knowledge distillation fusion, which is applicable to a heterogeneous federated learning framework based on multi-knowledge distillation fusion, including:

[0077] Step 1: The client uses local data to calculate local knowledge. For example, the first client input is Represents the sample data of the u-th Batchsize of client 1, through the predicted label and the true label y u The error between them is used to calculate the cross entropy loss L ce At the same time, each client calculates the local class granularity knowledge based on local data through the local knowledge calculation module according to the formula Calculate and aggregate into a kind of granular knowledge;

[0078] Step 2: The server receives the class granularity knowledge uploaded by the client, and generates global class granularity knowledge by weighted aggregation of all client class granularity knowledge through the class granularity knowledge aggregation module;

[0079] Step 3: The server distributes global knowledge to Distributed to each client as teacher knowledge for subsequent client model training;

[0080] Step 4: The client receives global knowledge and performs temperature-adaptive knowledge distillation through the knowledge distillation module to maximize the effect of knowledge transfer between teachers and students;

[0081] Step 5: The client receives global knowledge and performs batch sample correlation knowledge distillation through the knowledge distillation module. By introducing batch-level sample correlation distillation loss, the model's excessive bias towards specific samples or categories is suppressed, and the model's understanding of the overall characteristics of the data is strengthened;

[0082] Step 6: The client updates the local model and temperature prediction model parameters based on the knowledge distillation loss and cross entropy loss.

[0083] In this embodiment, in step 2, a knowledge aggregation scheme that takes into account differences in knowledge clarity is adopted, which specifically includes the following steps:

[0084] Step 2.1: Calculate the clarity of the teacher and student knowledge separately, using the logarithmic sum function to quantify the smoothness of the output, denoted as clarity. Assume Where C represents the probability of the sample belonging to the category, and clarity is defined as follows:

[0085]

[0086] Step 2.2: According to the clarity of knowledge differentiation, different aggregation weights are assigned. Specifically, the definition of knowledge weight aggregation for class j is as follows:

[0087]

[0088] Among them, Sharpness j represents the sum of the clarity of all knowledge belonging to category j, represents the global knowledge of class j, Represents the set of clients that have samples of type j.

[0089] It should be noted that, in an optional embodiment, the temperature adaptive knowledge distillation in step 4 specifically includes the following steps:

[0090] Step 4.1: For each client, a separate lightweight model is used to learn a dynamic temperature prediction module θ k ,The training starts by first reversely optimizing the ,module model parameters of the temperature prediction module aiming to maximize the ,distillation loss between the student and the teacher;

[0091] Step 4.2: Pass the student logits and teacher logits as input to the temperature prediction module to predict the distillation temperature T suitable for the current sample. The prediction process can be expressed as:

[0092]

[0093] in, For the relationship mapping of the temperature prediction module, in order to ensure the predicted temperature value T pred Within a reasonable range, the predicted value is mapped to the preset temperature range [T start ,T end ], the specific formula is as follows:

[0094] T pred =T start +T end δ(T pred )

[0095] Where T start and T end is the set temperature range, and δ(·) is the activation function, which aims to map the model prediction value to between 0 and 1.

[0096] Step 4.3: Standardize the teacher knowledge and student knowledge to ensure that the student model can learn knowledge more effectively rather than mechanically imitate, thereby reducing the distillation loss in the case of accurate prediction and making the model more balanced in the performance of all categories during the learning process. According to the formula Get the knowledge after the teachers and students have standardized the process. is the mean value of students’ or teachers’ knowledge, and σ(Z) is the corresponding standard deviation;

[0097] Step 4.4: According to the different clarity of the standardized teacher and student knowledge, obtain the differentiated distillation temperature of the teacher and the student. Specifically, the calculation method of the clarity sharp is as follows:

[0098]

[0099] Among them, Z is the student or teacher knowledge. Compared with student knowledge, teacher knowledge with relatively large clarity should have a higher temperature, while knowledge with less clarity should have a lower temperature. The distillation temperature between teachers and students is dynamically adjusted for each sample to better reflect the characteristics of the sample. This adaptive adjustment mechanism makes up for the defect of standardization lacking sample-level characteristic considerations, making knowledge transfer more targeted. The distillation temperature of teachers and students is adjusted in the following ways:

[0100]

[0101] Where T s and T t Differentiated temperature for knowledge distillation and T s ≠T t ,This step further forms personalized knowledge transfer at the sample level, by adaptively generating appropriate temperature values ​​for each sample and category, and fine-tuning the temperature according to sample characteristics during the knowledge transfer process, maximizing the absorption of teacher knowledge by the student model;

[0102] Step 4.5: Calculate the temperature T of each element c∈[C] in the knowledge vector according to the teacher-student differentiation temperature s or T t The transformation mapping under the action of the student model output knowledge as an example, each element c∈[C] in the knowledge vector is at temperature Ts The transformation mapping under action is:

[0103]

[0104] Among them, C is the category that the sample may belong to;

[0105] Step 4.6: Based on the knowledge of teachers and students, the temperature T is the temperature of the distillation at the differential distillation temperature. s and T t Transformation mapping under action and Calculate its temperature adaptive knowledge distillation loss, the calculation formula is as follows:

[0106]

[0107] It should be noted that, in an optional embodiment, in step 5, a knowledge distillation of the correlation of samples within a batch specifically includes the following steps:

[0108] Step 5.1: Uniformly quantify the performance of each sample in a batch in a category, so that the model considers the influence of other samples while considering the prediction of a single sample, so as to prevent the model from being overly biased towards certain specific samples or categories. Formally, the normalized logits values ​​of the teacher and student models for sample i in category c∈[C] are expressed as:

[0109]

[0110] Where B is the batch size, T is the set distillation temperature, and are the logits values ​​of teacher knowledge and student knowledge for sample i on category c respectively;

[0111] Step 5.2: Take the normalized knowledge in the category direction as input, and the calculation method of the knowledge distillation of the correlation between samples in the batch is as follows:

[0112]

[0113] Batch class normalization introduces an implicit regularization effect, which prevents the model from overfitting the prediction of a single sample during training. In the loss function of batch class normalization, the output of the student model is is the logits normalization function of all samples in the batch on category c, so the loss function is about the student model logits The gradient of is:

[0114]

[0115] By using the chain rule, we can By taking partial derivatives, we can get the gradient formula:

[0116]

[0117] The gradient of sample i on category c depends not only on its own probability distribution, but also on the probability distribution of other samples in the batch. represents Kronecker delta, which is 1 when i = b and 0 otherwise. This gradient formula shows that the intra-batch category normalization introduces mutual constraints between samples in the batch, thereby limiting the gradient update amplitude of a single sample.

[0118] It should be noted that, in an optional embodiment, step 6 specifically includes the following steps:

[0119] Step 6.1: According to the temperature-adaptive knowledge distillation method, obtain the temperature-adaptive knowledge distillation loss L TAKD ;

[0120] Step 6.2: Obtain the temperature-adaptive knowledge distillation loss L based on the knowledge distillation method of sample correlation within the batch CRKD ;

[0121] Step 6.3: Combine these two loss functions and the cross entropy loss to form the final distillation loss function:

[0122] Loss = λ·L ce +α·L TAKD +β·L CRKD

[0123] Among them, λ, α, and β are the weight parameters of the three loss components respectively. The overall fusion strategy improves the robustness and generalization ability of the model by balancing the contribution of each loss function;

[0124] Step 6.4: The client calculates the total loss Loss relative to the local model parameters W through the back propagation algorithm k and the temperature prediction module parameter θ k Gradient And use the stochastic gradient descent optimization algorithm to update the model parameters. The formula is as follows:

[0125]

[0126] Among them, η is the learning rate of the local model, η T is the learning rate of the temperature prediction module parameters, and are the gradients of the total loss function with respect to the local model parameters and the temperature prediction module parameters, respectively.

[0127] In order to verify the effectiveness of the proposed method, this paper conducted a comparative experiment on the proposed method and other federated learning methods:

[0128] 1) FD: A federated distillation method based on a class-granular knowledge interaction system;

[0129] 2) FedCache: A federated learning algorithm based on sample-granular knowledge caching;

[0130] 3) FedMkd: A heterogeneous federated learning method for multi-knowledge distillation fusion disclosed in the present invention;

[0131] Figure 4 The experimental results comparing the above federated learning algorithms on the MNIST dataset are shown. Some of the same conditions are set in the experiment for each method, such as the type of heterogeneous models, the number of clients, and the degree of data heterogeneity, to ensure the fairness of the experiment. From the results, it can be seen that the convergence curve of the method of the present invention is much steeper than that of FedCache in the case of heterogeneous client models, similar to FD, but FedMkd obtains a higher MAUA of 89.02%, and the communication overhead is also reduced by an order of magnitude compared to FedCache.

[0132] The present invention proposes and integrates two different knowledge distillation methods, namely temperature-adaptive knowledge distillation and batch sample category correlation knowledge distillation, to achieve higher accuracy with minimal communication overhead, effectively alleviating the many challenges faced by traditional federated learning frameworks when dealing with heterogeneous devices and heterogeneous data environments, especially in the case of large differences in device performance, high communication overhead and uneven data distribution, which significantly limits the accuracy and efficiency of client model training.

[0133] It should be noted that the embodiments of the present invention are described in detail above in conjunction with the accompanying drawings, but the present invention is not limited thereto, and various changes can be made within the knowledge scope of technicians in the relevant technical field without departing from the purpose of the present invention.

Claims

1. A heterogeneous federated learning framework based on multi-knowledge distillation fusion, characterized by: include: an edge server and at least one client device; Each client k∈{1, 2, 3, ...K} has its local dataset, where the dataset is expressed as follows: And there is |D k | samples, each of which has D data dimensions and belongs to one of C different classes. Due to different user behaviors, the local training and test data sets between clients are not independent and identically distributed. The personalized model parameters of client k are recorded as where d k represents the number of parameters in the client k model. There is system heterogeneity between client devices, and the scale of models deployed on the client is different. For example, d l ≠d m , Each client k has a local optimization goal, which depends on its corresponding local data distribution, and the overall goal is to minimize the expected goal of all clients, which is expressed as: The edge server is configured with: a class granularity knowledge aggregation module, which is used to receive class granularity knowledge from multiple clients and generate global class granularity knowledge; a knowledge distribution module, which is used to send the global class granularity knowledge to each client; the client is configured with: a local knowledge calculation module, which is used to calculate local class granularity knowledge based on local data; The knowledge distillation module is used to perform knowledge distillation based on the received global class granularity knowledge and local class granularity knowledge.

2. A heterogeneous federated learning method based on multi-knowledge distillation fusion, which is applicable to a heterogeneous federated learning framework based on multi-knowledge distillation fusion as described in claim 1, characterized in that: include: Step 1: The client uses local data to calculate local knowledge. For example, the first client input is Represents the sample data of the u-th Batchsize of client 1, through the predicted label and the true label y u The error between them is used to calculate the cross entropy loss L ce At the same time, each client calculates the local class granularity knowledge based on local data through the local knowledge calculation module according to the formula Calculate and aggregate into a kind of granular knowledge; Step 2: The server receives the class granularity knowledge uploaded by the client, and generates global class granularity knowledge by weighted aggregation of all client class granularity knowledge through the class granularity knowledge aggregation module; Step 3: The server distributes global knowledge to Distribute to each client as teacher knowledge for subsequent training of the client model; Step 4: The client receives global knowledge and performs temperature-adaptive knowledge distillation through the knowledge distillation module to maximize the effect of knowledge transfer between teachers and students; Step 5: The client receives global knowledge and performs batch sample correlation knowledge distillation through the knowledge distillation module. By introducing batch-level sample correlation distillation loss, the model's excessive bias towards specific samples or categories is suppressed, and the model's understanding of the overall characteristics of the data is strengthened; Step 6: The client updates the local model and temperature prediction model parameters based on the knowledge distillation loss and cross entropy loss.

3. A heterogeneous federated learning method based on multi-knowledge distillation fusion according to claim 2, characterized in that: In step 2, a knowledge aggregation scheme that takes into account differences in knowledge clarity is adopted, which specifically includes the following steps: Step 2.1: Calculate the clarity of the teacher and student knowledge respectively, using the logarithmic sum function to quantify the smoothness of the output, denoted as clarity. Assume Where C represents the probability of the sample belonging to the category, and clarity is defined as follows: Step 2.2: Different aggregation weights are assigned according to the clarity of the knowledge differentiation. Specifically, the definition of the knowledge weight aggregation for class j is as follows: Among them, Sharpness j represents the sum of the clarity of all knowledge belonging to category j, represents the global knowledge of class j, Represents the set of clients that have samples of type j.

4. A heterogeneous federated learning method based on multi-knowledge distillation fusion according to claim 3, characterized in that: The temperature adaptive knowledge distillation in step 4 specifically includes the following steps: Step 4.1: For each client, a separate lightweight model is used to learn a dynamic temperature prediction module θ k ,The training starts by first reversely optimizing the ,module model parameters of the temperature prediction module with the aim of maximizing ,the distillation loss between the student and the teacher; Step 4.2: Pass the student logits and teacher logits as input to the temperature prediction module to predict the distillation temperature T suitable for the current sample. The prediction process can be expressed as: in, For the relationship mapping of the temperature prediction module, in order to ensure the predicted temperature value T pred Within a reasonable range, the predicted value is mapped to the preset temperature range [T start , T end ], the specific formula is as follows: T pred =T start +T end δ(T pred ) Where T start and T end is the set temperature range, and δ(·) is the activation function, which aims to map the model prediction value to between 0 and 1. Step 4.3: Standardize the teacher knowledge and student knowledge to ensure that the student model can learn knowledge more effectively rather than mechanically imitate, thereby reducing the distillation loss in the case of accurate prediction and making the model more balanced in the performance of all categories during the learning process. According to the formula Get the knowledge after the teachers and students have standardized the process. is the mean value of students’ or teachers’ knowledge, and σ(Z) is the corresponding standard deviation; Step 4.4: According to the different clarity of the teacher and student knowledge after the standardization, the teacher and student differentiated distillation temperature is determined. Specifically, the calculation method of the sharpness is as follows: Among them, Z is the student or teacher knowledge. Compared with student knowledge, teacher knowledge with relatively large clarity should have a higher temperature, while knowledge with less clarity should have a lower temperature. The distillation temperature between teachers and students is dynamically adjusted for each sample to better reflect the characteristics of the sample. This adaptive adjustment mechanism makes up for the defect of standardization lacking sample-level characteristic considerations, making knowledge transfer more targeted. The distillation temperature of teachers and students is adjusted in the following ways: Where T s and T t Differentiated temperature for knowledge distillation and T s ≠T t ,This step further forms personalized knowledge transfer at the sample level, by adaptively generating appropriate temperature values ​​for each sample and category, and fine-tuning the temperature according to sample characteristics during the knowledge transfer process, maximizing the absorption of teacher knowledge by the student model; Step 4.5: Calculate each element c∈[C] in the knowledge vector at temperature T according to the teacher-student differentiation temperature s or T t The transformation mapping under the action of the student model output knowledge as an example, each element c∈[C] in the knowledge vector is at temperature T s The transformation mapping under action is: Among them, C is the category that the sample may belong to; Step 4.6: Based on the teacher-student knowledge, the temperature T is the temperature at which the distillation temperature is differentiated. s and T t Transformation mapping under action and Calculate its temperature adaptive knowledge distillation loss, the calculation formula is as follows:

5. A heterogeneous federated learning method based on multi-knowledge distillation fusion according to claim 4, characterized in that: In step 5, a knowledge distillation of the correlation of samples within a batch specifically includes the following steps: Step 5.1: Uniformly quantify the performance of each sample in a batch in a category, so that the model considers the influence of other samples while considering the prediction of a single sample, so as to prevent the model from being overly biased towards certain specific samples or categories. Formally, the normalized logits values ​​of the teacher and student models for sample i in category c∈[C] are expressed as: Where B is the batch size, T is the set distillation temperature, and are the logits values ​​of teacher knowledge and student knowledge for sample i on category c respectively; Step 5.2: Taking the normalized knowledge in the category direction as input, the calculation method of the knowledge distillation of the sample association within the batch is as follows: Batch class normalization introduces an implicit regularization effect, which prevents the model from overfitting the prediction of a single sample during training. In the loss function of batch class normalization, the output of the student model is is the logits normalization function of all samples in the batch on category c, so the loss function is about the student model The gradient of is: By using the chain rule, we can get the formula By taking partial derivatives, we can get the gradient formula: The gradient of sample i on category c depends not only on its own probability distribution, but also on the probability distribution of other samples in the batch. represents Kronecker delta, which is 1 when i = b and 0 otherwise. This gradient formula shows that the intra-batch category normalization introduces mutual constraints between samples in the batch, thereby limiting the gradient update amplitude of a single sample.

6. A heterogeneous federated learning method based on multi-knowledge distillation fusion according to claim 5, characterized in that: The step 6 specifically includes the following steps: Step 6.1: According to the temperature-adaptive knowledge distillation method, obtain the temperature-adaptive knowledge distillation loss L TAKD ; Step 6.2: Obtain the temperature-adaptive knowledge distillation loss L according to the knowledge distillation method of the sample association within the batch CRKD ; Step 6.3: Combine these two loss functions and the cross entropy loss to form the final distillation loss function: Loss=λ·L ce +α·L TAKD +β·L CRKD Among them, λ, α, and β are the weight parameters of the three loss components respectively. The overall fusion strategy improves the robustness and generalization ability of the model by balancing the contribution of each loss function; Step 6.4: The client calculates the total loss Loss relative to the local model parameters W through the back propagation algorithm k and the temperature prediction module parameter θ k Gradient And use the stochastic gradient descent optimization algorithm to update the model parameters. The formula is as follows: Among them, η is the learning rate of the local model, η T is the learning rate of the temperature prediction module parameters, and are the gradients of the total loss function with respect to the local model parameters and the temperature prediction module parameters, respectively.

Citation Information

Patent Citations

  • Federal learning model aggregation method based on dynamic adaptive knowledge distillation

    CN116681144A

  • Image classification method based on federal knowledge distillation and ensemble learning

    CN117523291A

  • Federal self-supervised contrast learning image classification system and method based on knowledge distillation

    CN117893807A

  • Distillation personalized federal learning method based on discriminator knowledge

    CN118095412A

  • Methods, devices and media for improving knowledge distillation using intermediate representations

    US20220335303A1

Cited By

  • Data transaction method based on multi-constraint unsupervised federated knowledge distillation

    CN120671771A

  • A Data Transaction Approach Based on Multi-Constraint Unsupervised Federated Knowledge Distillation

    CN120671771B

  • Equipment data-free federation incremental learning method and system under resource limitation

    CN121766393A

  • A device federated incremental learning method and system under resource constraints

    CN121766393B

  • Photovoltaic power prediction-oriented multi-level security data resource pool construction method

    CN122113142A