A Federated Domain Adaptation Method for Data Heterogeneity

The global model is optimized through knowledge distillation and weight calculation, and combined with the comparison training loss function to guide the local model update, solving the problem of global model performance degradation caused by data heterogeneity, and achieving higher precision federated domain adaptation.

CN114881134BActive Publication Date: 2025-07-29SHANGHAI UNIV OF ENG SCI
View PDF 3 Cites 0 Cited by

Patent Information

Application Number
CN202210450589.X
Authority / Receiving Office
CN · China
Patent Type
Patents(China)
Current Assignee / Owner
Filing Date
2022-04-26
Publication Date
2025-07-29
Estimated Expiration
2042-04-26

AI Technical Summary

Technical Problem

In the case of data heterogeneity, the performance of the global model declines. The existing knowledge distillation method fails to make full use of integrated knowledge to guide local source domain model learning, resulting in the model deviating from global optimality.

Method used

Knowledge distillation is used to extract high-quality knowledge, combine weight calculation and comparison training methods, optimize the update of global models and local models, acquire high-quality knowledge through knowledge distillation and assign corresponding weights to each local source domain model, and guide the update of local models with contrast training loss function.

Benefits of technology

It improves the performance of the global model in the case of data heterogeneity, reduces the requirements for source domain data quality, and enhances the accuracy and consistency of the model.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN114881134B_ABST
    Figure CN114881134B_ABST
Patent Text Reader

Abstract

The present invention belongs to the technical field of machine learning and discloses a federated domain adaptation method applied to data heterogeneity. Each federated source domain node performs a current round of training based on a relevant local data set to obtain a corresponding local source domain model. Knowledge distillation is performed on all local source domain models to obtain high-quality knowledge and the number of local source domain models that support the high-quality knowledge. The knowledge is uploaded together with all local source domain models to a central server. The weight corresponding to each local source domain model is then calculated, and a global model is established through an aggregation operation. Finally, the central server sends the global model to each source domain node. Combined with a comparative training method, the next round of training is performed to obtain a local source domain model corresponding to each source domain node. The process is iterated in sequence until the global model converges, thereby completing the adaptive learning of the federated domain.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention belongs to the technical field of machine learning, and particularly relates to a federated domain adaptation method applied to data heterogeneity. Background Art

[0002] In unsupervised deep learning, in order to avoid the costly annotation process, we usually use other similar datasets (i.e., source domains) to train models that can be applied to new datasets (i.e., target domains), which is the domain adaptation problem. Traditional domain adaptation methods do not consider the problem that source domain data is unavailable. In order to train models that can be applied to the target domain without obtaining source domain data, researchers have proposed federated domain adaptation, which deploys domain adaptation on the framework of federated learning. Federated learning is a distributed machine learning setting that can reduce the privacy leakage risk and data transmission cost of traditional centralized machine learning methods, and ensure the privacy security of source domain data.

[0003] The implementation process of traditional federated learning is mainly divided into four parts: ① Each federated learning source domain node trains a local source domain model based on local source domain data, ② uploads the local source domain model parameters and weights to the central server, ③ the central server aggregates the local source domain model parameters of each node to obtain a global model, and ④ finally each federated learning source domain node downloads the global model for the next round of local source domain model update.

[0004] However, during the collection process of each source domain data, due to uncertain factors such as environment, equipment, and the personal style of photographers, there will be heterogeneity among the data of each source domain. The data heterogeneity among each source domain will cause the local source domain models to tend to local optimal models and deviate from the global optimal model, thereby degrading the performance of the global model learned by federated domain adaptation. In the prior art, knowledge distillation is often used to alleviate the impact of data heterogeneity on the global model. Knowledge distillation enriches the global model by using the integrated knowledge from local source domain models, but it does not fully utilize the integrated knowledge to guide the learning of local source domain models, which in turn reduces the performance of the global model. Summary of the Invention

[0005] The present invention provides a federated domain adaptation method for data heterogeneity, which consists of three sub-methods in series and is deployed in the framework of federated learning. First, the knowledge distillation integration strategy can extract high-quality knowledge from source domain data. Then, in the aggregation method of the federated learning model, the high-quality knowledge extracted by the previous method is used to assign corresponding weights to each local source domain model to optimize the global model, improving the performance of the global model in the case of data heterogeneity. Finally, the model comparison loss is calculated by combining the local source domain model and the original global model to guide the update training of the local source domain model, reducing the requirements for the quality of the source domain data for federated domain adaptation. The present invention optimizes both the update stage of the local source domain model and the aggregation stage of the central server, enabling the present invention to better handle the data heterogeneity problem in federated domain adaptation.

[0006] The present invention can be realized through the following technical solutions:

[0007] A federated domain adaptation method for data heterogeneity, where each federated source domain node performs current-round training based on the relevant local dataset to obtain the corresponding local source domain model. Knowledge distillation is performed on all local source domain models to obtain high-quality knowledge and the number of local source domain models supporting the high-quality knowledge, and they are jointly uploaded to the central server together with all local source domain models. Then, the corresponding weight of each local source domain model is calculated, and a global model is established through an aggregation operation. Finally, the central server distributes the global model to each source domain node, and then, in combination with the contrast training method, the next-round training is performed to obtain the corresponding local source domain model for each source domain node, and the iteration is carried out in turn until the global model converges, thus completing the adaptive learning of the federated domain.

[0008] Furthermore, the method for performing knowledge distillation on all local source domain models includes the following steps:

[0009] Step I: Denote the local source domain model as Input the sample data of the target domain D t into the K local source domain models successively for training and learning, and calculate the set of confidence prediction values for all classes

[0010] Step II: Set a confidence threshold, and filter out the local source domain models whose confidence prediction values for any class do not exceed the confidence threshold;

[0011] Step III: For the remaining local source domain models, sum the confidence prediction values of the same class, set the class with the largest sum as the resonance class, and then filter out the local source domain models corresponding to the confidence prediction values of the resonance class that are less than the confidence threshold;

[0012] ​Step Ⅳ: For the locally sourced models that are retained again, average and integrate the corresponding confidence prediction values to obtain high-quality knowledge p i , and record the number of locally sourced models that support p i

[0013] Furthermore, the method for establishing the global model includes the following steps:

[0014] Step i: Use the following equation to calculate the weight of each locally sourced model

[0015]

[0016]

[0017]

[0018] where represents the contribution degree of the k-th locally sourced model , α k represents the weight of the k-th sourced model

[0019] Step ii: Multiply the original parameters of each locally sourced model by the corresponding weight to obtain new parameters, and then perform an aggregation operation in combination with the corresponding locally sourced models to establish a global model.

[0020] Furthermore, the comparison training method includes the following steps:

[0021] S1: Each source domain node downloads the global model from the central server and saves a corresponding copy of the global model locally, which are the original global model and the original global model copy;

[0022] S2: For each source domain node, based on the corresponding local dataset, train and learn the original global network for the current epoch, calculate the cross-entropy loss function l cro corresponding to the current epoch and the initial global model of the current epoch. At the same time, use the source domain model of the previous round, the initial global model of the current epoch, and the original global model copy to extract features from the local dataset, and the obtained features are denoted as z prev , z, z glob . Use the following equation to calculate the comparison loss function l con

[0023]

[0024] ​​​​Furthermore, the total loss function $l = l$ cro $+ \mu l$ con is obtained, where $\mu$ is a hyperparameter used to control the weight of the model contrast loss, and $t$ represents the temperature parameter. Then, the total loss function is applied to the training learning of the current epoch to obtain the final global model of the current epoch;

[0025] S3. Repeat step S2, use the final global model of the current epoch to replace the original global model for the training learning of the next epoch until convergence to obtain the final global model, which is used as the local source domain model for the next round. At the same time, each source domain node releases the copy of the original global model and the local source domain model of the previous round.

[0026] The beneficial technical effects of the present invention are as follows:

[0027] Compared with the prior art, the adaptive method of the present invention enables the high-quality knowledge extracted through knowledge distillation to not only act on the optimization of the global model but also guide the training of the local source domain model according to the idea of contrast learning. It reduces the impact of source domain data heterogeneity on the accuracy of the final global model during both the aggregation stage of federated learning and the update training stage of the local source domain model, improves the accuracy of the model finally applied to the target domain, and reduces the requirements for the quality of source domain data for federated domain adaptation. BRIEF DESCRIPTION OF THE DRAWINGS

[0028] Figure 1 is the overall flow schematic diagram of the present invention;

[0029] Figure 2 is the process schematic diagram of knowledge distillation of the present invention;

[0030] Figure 3 is the process schematic diagram of the federated learning model aggregation method of the present invention;

[0031] Figure 4 is the process schematic diagram of the contrast training method of the present invention. DETAILED DESCRIPTION OF THE EMBODIMENTS

[0032] The following describes the specific embodiments of the present invention in detail with reference to the accompanying drawings and preferred embodiments.

[0033] As Figure 1As shown in the figure, the present invention provides a federated domain adaptation method applied to data heterogeneity. Each federated source domain node performs current round of training based on the relevant local dataset to obtain the corresponding local source domain model. Knowledge distillation is performed on all the local source domain models to obtain high-quality knowledge and the number of local source domain models supporting the high-quality knowledge, and they are uploaded to the central server together with all the local source domain models. Then, the weight corresponding to each local source domain model is calculated, and a global model is established through an aggregation operation. Finally, the central server distributes the global model to each source domain node, and then, in combination with the contrast training method, the next round of training is performed to obtain the corresponding local source domain model for each source domain node, and the iteration is carried out in turn until the global model converges, thereby completing the adaptive learning of the federated domain. Specifically as follows:

[0034] Step 1. Knowledge distillation

[0035] To prepare for the aggregation operation of the central server and solve the problem that traditional knowledge distillation integration strategies such as maximum integration and average integration cannot obtain high-quality knowledge, the knowledge distillation method of the present invention can distill higher-quality knowledge, laying a better foundation for subsequent operations of federated domain adaptation. As Figure 2 shown, the specific strategy implementation is as follows:

[0036] D1. Properly save the local source domain model of each source domain node locally;

[0037] D2. Successively input the sample data of the target domain D t into K local source domain models for training and learning, and calculate to obtain the set of confidence prediction values for all classes Then, a relatively high confidence threshold can be set to filter out the unconfident models in the local source domain models. The unconfident model is defined as the local source domain model whose confidence prediction value for any class does not exceed the confidence threshold;

[0038] D3. For the remaining local source domain models, sum the confidence prediction values of the same class, and set the class with the largest sum as the resonance class. Then, filter out the local source domain models corresponding to the confidence prediction value of the resonance class being less than the confidence threshold;

[0039] D4. At this point, a set of local source domain models that all support the resonance class can be obtained, and the confidence prediction values of each class of these local source domain models are simply averaged to obtain high-quality knowledge p i Meanwhile, record the number of local source domain models i supporting the high-quality knowledge p

[0040] ​To mitigate the impact of data heterogeneity on the global model during the model aggregation phase, the federated learning model aggregation method of the present invention utilizes high-quality knowledge p i and the number of supported local source domain models to calculate the weights of each local source domain model, as Figure 3 shown, including the following two processes:

[0041] Sⅰ. Using the following equation, calculate the weight of each local source domain model where,

[0042]

[0043]

[0044]

[0045] where represents the contribution degree of the k-th local source domain model , α k represents the weight of the k-th source domain model , K represents the total number of local source domain models,

[0046] Sⅱ. Multiply the original parameters of each local source domain model by the corresponding weight to obtain new parameters, and then jointly perform an aggregation operation with the corresponding local source domain model to establish a global model.

[0047] To address the drawback that the high-quality knowledge extracted by knowledge distillation can only be used to optimize the global model and cannot guide the training of local source domain models, the present invention provides a brand-new training method, enabling the local source domain models to be closer to the global model and preventing deviation from the globally optimal model, as Figure 4 shown, specifically as follows:

[0048] S1. Each source domain node downloads the global model from the central server and saves the corresponding global model copy locally, namely the original global model and the original global model copy;

[0049] S2. Whether it is a local source domain model or a global model, the network architecture can be divided into a feature extraction part and a classifier part. The feature extraction part is used to obtain the feature representation of the same dimension for each image, and the classifier is used to generate confidence prediction values for each class. In the embodiment of the present invention, the feature extraction part uses the classic ResNet network, and the classifier part uses two linear layers and a fully connected layer to achieve the purpose.

[0050] To implement contrastive training, we split the loss function of the local source domain model into two parts. One part is the most typical cross-entropy loss function lcro , which is obtained by supervised learning of each source domain node based on the relevant local dataset; the other part is the contrast loss function, which is obtained by contrastive learning of each source domain node based on the relevant local dataset, that is, contrastive learning needs to be performed based on three models: the local source domain model of the previous round, the initial global model of the current epoch, and the original global model copy. The specific process is as follows:

[0051] For each source domain node, the original global network is trained and learned for the current epoch based on the corresponding local dataset, and the cross-entropy loss function l corresponding to the current epoch is calculated cro and the initial global model of the current epoch. At the same time, the local source domain model of the previous round, the initial global model of the current epoch, and the original global model copy are used to extract features from the local dataset, and the obtained features are denoted as z prev , z, z glob respectively. The contrast loss function l is calculated using the following equation con ,

[0052]

[0053] Furthermore, the total loss function l = l cro + μl con is obtained, where μ is a hyperparameter used to control the weight of the model contrast loss, and t represents the temperature parameter. Then, the total loss function is applied to the training and learning of the current epoch to obtain the final global model of the current epoch;

[0054] S3. Repeat step S2, using the final global model of the current epoch to replace the original global model for the training and learning of the next epoch until the final global model is obtained through convergence. This is used as the local source domain model for the next round, and at the same time, each source domain node releases the original global model copy and the local source domain model of the previous round, and no longer saves them locally.

[0055] The method for federated domain adaptation applicable to data heterogeneity provided by the present invention is composed of the above three sub-methods in series. To make the overall method framework structure clearer, a schematic diagram of the overall process of the method for federated domain adaptation based on data heterogeneity is provided, as shown in Figure 1 . Only in the first round, each federated learning source domain node obtains the local source domain model through supervised learning based on local relevant data. During the training process, only the traditional cross-entropy loss function is involved, and the model contrast loss function is not involved. From the second round until the global model converges, the local source domain model is trained and updated according to the method for training the local source domain model provided by the present invention.

[0056] As described above, the above embodiments are only used to illustrate the technical solutions of the present invention and are not intended to limit them; although the present invention has been described in detail with reference to the foregoing embodiments, those of ordinary skill in the art should understand that they can still modify the technical solutions described in the foregoing embodiments, or perform equivalent replacements for some of the technical features; and these modifications or replacements do not cause the essence of the corresponding technical solutions to deviate from the spirit and scope of the technical solutions of the embodiments of the present invention.

Claims

1. A federated domain adaptation method applied to data heterogeneity, characterized in that: Each federal source domain node trains the current round based on the relevant local dataset to obtain the corresponding local source domain model, conducts knowledge distillation on all local source domain models, obtains high-quality knowledge and the number of local source domain models supporting the high-quality knowledge, uploads them to the central server together with all local source domain models, then calculates the weights corresponding to each local source domain model, establishes a global model through an aggregation operation, and finally the central server distributes the global model to each source domain node. Then, in combination with the contrast training method, the next round of training is carried out to obtain the local source domain model corresponding to each source domain node, and the iteration is carried out in turn until the global model converges, thus completing the adaptive learning of the federal domain; Among them, the method for establishing a global model includes the following steps: Step ⅰ. Calculate the weights of each local source domain model using the following equation and Among them, represents the contribution degree of the k-th local source domain model , and α k represents the weight of the k-th local source domain model . K represents the total number of local source domain models, p i represents high-quality knowledge, and n pi represents the number of local source domain models that support the high-quality knowledge p i . Step ii: Multiply the original parameters of each local source domain model by the corresponding weight to obtain new parameters, and then jointly perform an aggregation operation in combination with the corresponding local source domain model to establish a global model.

2. The federated domain adaptation method applied to data heterogeneity according to claim 1, wherein The method for conducting knowledge distillation on all local source domain models includes the following steps: Step Ⅰ. Denote the local source domain model as Input the sample data of the target domain D t successively into K local source domain models for training and learning, and calculate the set of confidence prediction values for all classes ​ Step II: Set a confidence threshold, and filter out the local source domain models whose confidence prediction values for any class do not exceed the confidence threshold; Step III: For the remaining local source domain models, sum the confidence prediction values of the same class, set the class with the largest sum as the resonance class, and then filter out the local source domain models corresponding to the confidence prediction value of the resonance class being less than the confidence threshold; Step Ⅳ. For the locally sourced domain models that are retained again, average and integrate the confidence prediction values corresponding to all classes to obtain high-quality knowledge p i , and at the same time record the number n i of the locally sourced domain models that support p pi .

3. The federated domain adaptation method applied to data heterogeneity according to claim 1, characterized in that The contrast training method includes the following steps: S1: Each source domain node downloads the global model from the central server and saves the corresponding global model copy locally, namely the original global model and the original global model copy; S2. For each source domain node, based on the corresponding local dataset, train and learn the original global network for the current epoch, and calculate the cross-entropy loss function \(l\) corresponding to the current epoch. cro And the initial global model for the current epoch. At the same time, use the local source domain model of the previous round, the initial global model of the current epoch, and the original global model copy to extract features from the local dataset. The obtained features are denoted as \(z\) prev , \(z\), \(z\) glob respectively. Use the following equation to calculate the contrastive loss function \(l\) con , Furthermore, the total loss function l = l cro + μl con , where μ is a hyperparameter used to control the weight of the model contrast loss, t represents the temperature parameter. Then, the total loss function is applied to the training learning of the current epoch to obtain the final global model of the current epoch; S3: Repeat step S2, use the final global model of the current epoch to replace the original global model for the training and learning of the next epoch until the final global model is obtained through convergence, and use this as the local source domain model for the next round. At the same time, each source domain node releases the original global model copy and the local source domain model of the previous round.

Citation Information

Patent Citations

  • Federal learning-based distributed language relationship identification method, system and device

    CN112101578A

  • Safety production early warning system based on multi-source heterogeneous data federal learning

    CN113160021A

  • Federal learning-based space-time prediction algorithm on industrial internet-of-things edge device

    CN114265913A