Training Method of Domain Adaptation Type Neural Network
The method addresses challenges in unsupervised domain adaptation by employing a voting scheme for pseudo-label improvement, self-learning decay rates for adaptive knowledge distillation, and data distillation for simplified and effective domain adaptation.
Patent Information
- Application Number
- JP2021136658
- Authority / Receiving Office
- JP · JP
- Patent Type
- Patents
- Current Assignee / Owner
- Priority Date
- 2020-09-02
- Filing Date
- 2021-08-24
- Publication Date
- 2025-06-18
- Estimated Expiration
- 2041-08-24
AI Technical Summary
Existing methods for unsupervised domain adaptation face challenges such as reliance on accurate pseudo-labels, fixed decay rates in knowledge distillation, and the need for multiple training stages, which can lead to suboptimal performance and increased complexity.
The proposed method involves a domain adaptation neural network with a voting scheme to improve pseudo-label accuracy, a self-learning decay rate for adaptive ensemble learning, and a data distillation approach using a domain discriminator to weight source data based on similarity to target data.
This approach enhances the accuracy of pseudo-labels, improves the performance of knowledge distillation by dynamically adjusting the ensemble speed, and simplifies the training process by integrating data distillation, resulting in improved model performance and adaptability across domains.
Smart Images

Figure 0007694254000048 
Figure 0007694254000049 
Figure 0007694254000050
Abstract
Description
Technical Field
[0001] The present invention generally relates to domain adaptation, and more specifically, to a neural network for unsupervised domain adaptation and a method for training the same.
Background Art
[0002] Unsupervised domain adaptation means transferring a model trained using labeled source data to an unlabeled target domain of data while maintaining the performance of the model in the target domain as much as possible. There is a deviation in the data set between the source domain and the target domain, and since there is insufficient labeled data in the target domain, the performance of a model trained using labeled source data may decrease in the target domain. The training process of unsupervised domain adaptation can effectively reduce the difference between domains and improve the robustness of the model by using both the labeled data in the source domain and the unlabeled data in the target domain.
[0003] Currently, the mainstream methods of unsupervised domain adaptation include learning methods for domain-invariant features such as adversarial training. A typical adversarial training method is a domain adversarial neural network. In this neural network, a domain discriminator is added after the feature extraction network to determine whether the features are from the source domain or the target domain, and a gradient reversal layer is added between the feature extraction network and the domain discriminator. When minimizing the loss function of the domain discriminator, the feature extraction network can learn domain-invariant features through the gradient reversal layer.
[0004] Furthermore, knowledge distillation has recently been introduced into teacherless domain adaptation, and many new methods have been proposed as follows. For example, use a self-ensembling teacher model to let the student model learn unlabeled data in the target domain. Utilize the self-ensembling teacher model to obtain more accurate pseudo-labels of the target data. Distill data similar to the target data from the source data to fine-tune the pre-trained model. Align the features of the source domain and the target domain at the semantic level (class level), that is, make the average features (class centers) of the same class in the source domain and the target domain approach each other.
[0005] The following briefly introduces these conventional methods.
[0006] Figure 1 shows the structure of a typical domain adversarial neural network. As shown in Figure 1, the domain adversarial neural network includes a feature extractor F, a classifier C s , and a domain discriminator D. The domain discriminator D is connected to the feature extractor F through a gradient reversal layer, and the gradient reversal layer multiplies the gradient by a specific negative number and sends it back to the feature extractor F. I s represents the labeled source data, and I t represents the unlabeled target data, and both are input into the feature extractor F. The features extracted by the feature extractor F for the source data are input into the classifier C s to predict the class of the source data. Also, the features extracted by the feature extractor F for both the source data and the target data are input into the domain discriminator D, and the domain discriminator D identifies whether the currently processed data is from the source domain or the target domain based on the input features. In the training of the domain adversarial neural network, the classification cross-entropy loss function L c of the source domain and the binary cross-entropy loss function L advUsing this, the loss functions L c and L adv are minimized by training according to the standard backpropagation algorithm, enabling the feature extractor F to learn domain-invariant features.
[0007] Figure 2 shows the structure of the self-ensembling teacher model, where the teacher network is constructed using the exponential moving average of the parameters of the student network. In Figure 2, x Si represents the labeled source data, x Ti represents the unlabeled target data, y Si represents the true label of the source data, z Ti represents the predicted probability of the target data by the student network, (Outer 1) TIFF0007694254000001.tif16170 represents the predicted probability of the target data by the teacher network.
[0008] The premise of this scheme is to assume that the prediction accuracy of the teacher network is higher than that of the student network. Since the student network can learn the hidden knowledge of the target data from the prediction probability of the teacher network, this scheme is knowledge distillation. For the source data x Si , the cross-entropy loss function based on the predicted probability z Ti of the student network and the true label y Si is adopted. For the target data x Ti , the mean squared error between the predicted probability (Outer 2) TIFF0007694254000002.tif17170 of the teacher network and the predicted probability z Ti of the student network is used as the loss function. Then, weighted addition is performed on the above two loss functions to obtain the final loss function.
[0009] Furthermore, regarding the alignment of features at the semantic level, the following loss function has been proposed. [Number] [Number]
[0010] Here, X s,k represents all data samples belonging to the k-th class in the source domain X s (determined based on the true label), and X t,k represents all data samples labeled as the k-th class in the target domain X t (determined based on the pseudo label). λ s,k represents the class center of the k-th class in the source domain, that is, the average value of the features F of all source data belonging to the k-th class. Similarly, λ t,k represents the class center of the k-th class in the target domain, that is, the average value of the features F of all target data labeled as the k-th class. The pseudo label of the target data is obtained by using a classifier to predict the class of the target data. The semantic alignment loss function L a (X s , X t ) represents the distance between the class centers of the same class in the source domain and the target domain.
[0011] Although the above method has achieved good results, there are still some problems that need to be improved. First, in semantic alignment, the accuracy of the pseudo-labels of the target data has a great impact on the class centers in the target domain. For some data near the classification boundary, if the pseudo-labels are incorrect, a large deviation will occur in the calculation results of the class centers. Also, in contrastive learning, incorrect pseudo-labels undermine the constraints on the clustering of data samples within a class and the separation of data samples between classes. Furthermore, in the mean teacher model of self-ensembling, a fixed decay rate is often used for exponential moving average. However, since the performance of the current model is variable, the ensemble speed cannot be adjusted according to the performance of the current model with a fixed decay rate. Additionally, when fine-tuning using distilled data, this method requires two stages, an additional switching operation in the middle, and the training cannot be completed in one go. Summary of the Invention Problems to be Solved by the Invention
[0012] The present invention provides a method and apparatus for training a domain-adaptive neural network. Means for Solving the Problems
[0013] In one aspect of the present invention, there is provided a method for training a domain adaptation neural network, which is executed by a computer. The domain adaptation neural network includes a first feature extraction unit, a first classification unit, and an identification unit. The computer includes a memory storing instructions and a processor. When the instructions are executed by the processor, the processor is caused to execute the method. The method includes: a step of the first feature extraction unit extracting first features for source data in a labeled source data set; a step of the first classification unit predicting, based on the first features, the probability that the source data belongs to each of a plurality of classes; a step of the first feature extraction unit extracting second features for target data in an unlabeled target data set; a step of the first classification unit predicting, based on the second features, the probability that the target data belongs to each of the classes, and determining the class corresponding to the maximum probability as the first label of the target data; a step of calculating the distance between the class center for each class in the source data set and the features of the target data, and determining the class corresponding to the closest class center as the second label of the target data; a step of selecting, from the target data set, target data for which the determined first label and the second label are the same, where the first label or the second label is used as the pseudo label of the selected target data; a step of calculating the class center for each class in the target data set based on the selected target data; a step of constructing a first loss function based on the distance between the class center of the source data set and the calculated class center of the target data set; a step of constructing a second loss function based on the selected target data and its pseudo label; a step of constructing a third loss function for the source data in the source data set and the selected target data; and a step of training the domain adaptation neural network based on the first loss function, the second loss function, and the third loss function.
[0014] In another aspect of the present invention, there is provided an apparatus for training a domain adaptation neural network, wherein the domain adaptation neural network extracts a first feature for source data in a labeled source data set and extracts a second feature for target data in an unlabeled target data set; a first feature extraction unit; based on the first feature, predicts the probability that the source data belongs to each of a plurality of classes, and based on the second feature, predicts the probability that the target data belongs to each of the classes, and determines, as a first label of the target data, a class corresponding to the maximum probability; a first classification unit; and an identification unit that determines the probability that currently input data is source data based on the first feature and the second feature, the apparatus including a memory storing a program and one or more processors, wherein the processor, by executing the program, calculates the distance between the class center for each class of the source data set and the features of the target data, and determines, as a second label of the target data, a class corresponding to the class center with the closest distance; selects target data for which the determined first label and the second label are the same from the target data set, wherein the first label or the second label is used as a pseudo-label of the selected target data; calculates the class center for each class of the target data set based on the selected target data; constructs a first loss function based on the distance between the class center of the source data set and the calculated class center of the target data set; constructs a second loss function based on the selected target data and its pseudo-label; constructs a third loss function for the source data in the source data set and the selected target data; and trains the domain adaptation neural network based on the first loss function, the second loss function, and the third loss function.
[0015] In another aspect of the present invention, there is provided a storage medium storing a program for training a domain adaptation neural network, wherein the domain adaptation neural network includes a first feature extraction unit, a first classification unit, and an identification unit. When the program is executed by a computer, the computer is caused to perform the steps of: the first feature extraction unit extracting first features from source data in a labeled source data set; the first classification unit predicting, based on the first features, the probability that the source data belongs to each of a plurality of classes; the first feature extraction unit extracting second features from target data in an unlabeled target data set; the first classification unit predicting, based on the second features, the probability that the target data belongs to each of the classes, and determining, as a first label of the target data, the class corresponding to the maximum probability; calculating the distance between the class center of each class in the source data set and the features of the target data, and determining, as a second label of the target data, the class corresponding to the class center with the closest distance; selecting, from the target data set, target data for which the determined first label and second label are the same, wherein the first label or the second label is used as a pseudo-label of the selected target data; calculating, based on the selected target data, the class center of each class in the target data set; constructing a first loss function based on the distance between the class center of the source data set and the calculated class center of the target data set; constructing a second loss function based on the selected target data and its pseudo-label; constructing a third loss function for the source data and the selected target data in the source data set; and training the domain adaptation neural network based on the first loss function, the second loss function, and the third loss function. A storage medium for causing the above method to be executed is provided.
Brief Description of the Drawings
[0016]
Figure 1
Figure 2
Figure 3
Figure 4
Figure 5
Figure 6
Figure 7
Figure 8
Figure 9
Figure 10
Embodiments for Carrying Out the Invention
[0017] Figure 3 schematically shows the structure of a neural network for teacherless domain adaptation according to the present invention. As shown in Figure 3, the neural network includes the domain adversarial neural network described with reference to Figure 1, and the neural network includes a first feature extractor 310, a first classifier 320, a domain discriminator 330, and a gradient reversal layer (not shown). Further, the neural network further includes a second feature extractor 310_T and a second classifier 320_T. Note that, as an existing technique, the first feature extractor 310, the second feature extractor 310_T, the first classifier 320, the second classifier 320_T, and the domain discriminator 330 in Figure 3 may all be implemented by a convolutional neural network. In this specification, the structure of the convolutional neural network for realizing these units will not be described in detail.
[0018] The first feature extractor 310 and the first classifier 320 constitute a student network, and the second feature extractor 310_T and the second classifier 320_T constitute a teacher network. The parameters of the second (teacher) feature extractor 310_T are the exponential moving average of the parameters of the first (student) feature extractor 310, and the parameters of the second (teacher) classifier 320_T are the exponential moving average of the parameters of the first (student) classifier 320.
[0019] Source data X s and target data X t are input to the first feature extractor 310 and the second feature extractor 310_T respectively. The first feature extractor 310 inputs the features extracted for the source data X s and target data X t to the first classifier 320, and the second feature extractor 310_T inputs the features extracted for the source data X s and target data X t to the second classifier 320_T.
[0020] In the training of the domain adaptation type neural network shown in Figure 3, the present invention proposes a plurality of loss functions, and the details thereof will be described below.
[0021] In one aspect of the present invention, a voting scheme is proposed to improve the accuracy of pseudo-labels of target data. The voting scheme means voting on the predicted label of the target data using at least two prediction methods. For example, for target data (Outer 3) TIFF0007694254000005.tif19170, the classifier is used to predict its class label and the prediction result (Outer 4) TIFF0007694254000006.tif18170 is obtained. Also, the class center nearest neighbor algorithm is used to predict its label, and as shown in the following mathematical formulas (2) and (3), the prediction result l d is obtained. [Number] [Number]
[0022] Here, λ s,k represents the class center of the k-th class in the source domain, that is, the average value of the characteristics of all source data belonging to the k-th class, K represents the number of all classes in the source domain, and l d represents the class corresponding to the class center closest to the target data (Outer 5) TIFF0007694254000009.tif17170 among all K class centers in the source domain.
[0023] When the predicted label l c and the predicted label l d match, the target data (Outer 6) TIFF0007694254000010.tif19170 is selected, and the predicted label l c or l d is assigned to the target data (External 7) Use TIFF0007694254000011.tif17170 as the pseudo-label. Prediction label l c and prediction label l d If they do not match, the target data (External 8) Ignore TIFF0007694254000012.tif17170. All selected target data (External 9) TIFF0007694254000013.tif18170 is the optimal target data set (External 10) TIFF0007694254000014.tif19170 constitutes. Compared with the case of performing only classifier prediction or only class-centered nearest neighbor prediction, the data set selected by such filtering (External 11) The probability of the pseudo-label of each target data in TIFF0007694254000015.tif18170 becomes higher. Therefore, the voting scheme according to the present invention can effectively select target data with more accurate prediction results.
[0024] Note that the above-mentioned classifier prediction and class-centered nearest neighbor prediction are merely examples of at least two different prediction methods, and the present invention is not limited thereto, and those skilled in the art can easily conceive of other appropriate prediction methods.
[0025] Next, based on the optimal target data set (External 12) Construct a semantic alignment loss function L for training the neural network shown in FIG. 3 based on TIFF0007694254000016.tif17170 a (not shown in FIG. 3), and the semantic alignment loss function L a is also referred to as the first loss function. Specifically, taking the k-th class among the K predetermined classes as an example, first, according to Equation (2), the class center λ of the k-th class in the source domains,k Calculate it and obtain the optimal target dataset according to the following formula (4) (Outer 13) The class center λ of the k-th class in TIFF0007694254000017.tif16170 t,k Calculate it and obtain the class center λ according to formula (5) s,k Between the class center λ t,k And the class center λ (Outer 14) Calculate TIFF0007694254000018.tif17170. In this way, for all k classes, calculate the distance between the class center of the source domain and the class center of the target domain as the semantic alignment loss function respectively. In training, the distance (Outer 15) TIFF0007694254000019.tif19170 is targeted to be minimized.
Number
Number
[0026] Optimal target dataset (Outer 16) Since the pseudo-labels of the target data in TIFF0007694254000022.tif16170 have higher accuracy, the dataset (Outer 17) The class center λ of the target domain calculated using TIFF0007694254000023.tif17170 t,k Is more accurate and can improve the semantics alignment loss function.
[0027] Furthermore, the optimal target dataset (Outer 18) Using the target data and its pseudo-labels in TIFF0007694254000024.tif at 19170, the cross-entropy loss function (in Figure 3 (Outer 19) for training the first classifier 320 shown in Figure 3 (in TIFF0007694254000025.tif at 18170) may be constructed. The cross-entropy loss function is also referred to as the second loss function and is specifically as shown in the following mathematical formula (6).
Equation
[0028] Here,[[]] (Outer 20) TIFF0007694254000027.tif at 17170 is the optimal target data set (Outer 21) When predicting the label for the target data in TIFF0007694254000028.tif at 16170 (Outer 22) TIFF0007694254000029.tif at 15170, the prediction result is the probability of its pseudo-label.
[0029] In the prior art, usually only the source data with true labels is used to train the first classifier 320. However, in the present invention, since the accuracy of the pseudo-labels of the target data in the optimal target data set (Outer 23) TIFF0007694254000030.tif at 17170 is relatively high, the present invention further uses the optimal target data set (Outer 24) TIFF0007694254000031.tif at 16170 to train the first classifier 320, thereby improving the recognition ability of the network model for the target data.
[0030] Furthermore, the optimal target data set (Outer 25) By using the target data at TIFF0007694254000032.tif16170 together with the source data for contrastive learning, the following effects can be achieved. Constraining to converge the features within a class while separating the features between different classes so that the distance between the features of different classes becomes large. Here, for example, a contrastive learning loss function L con (not shown in FIG. 3) may be constructed, and the contrastive learning loss function L con is also referred to as the third loss function.
Number
[0031] Here, x i or x j represents the source data set and the optimal target data set (Outer 26) TIFF0007694254000034.tif18170 represents a data sample, and f(x i ) and f(x j ) represent the features of the data sample. δ ij is an indicator variable. When x i and x j are data of the same class, δ ij is 1. When x i and x j are data of different classes, δ ij is 0. d(f(x i ), f(x j )) represents the distance between the feature of data x i and data x j . m is a constant, for example, m = 3.
[0032] As described above, the current knowledge distillation method used for unsupervised domain adaptation constructs the teacher network using exponential moving average, but since its decay rate is usually set to a fixed value, it is difficult to obtain a teacher network with excellent performance. Specifically, exponential moving average means slowly updating the parameters of the teacher network based on a specific decay rate, as shown in the following mathematical formula (8).
Equation
[0033] Here, S represents the current parameters of the student network, and T t represents the current parameters (updated parameters) of the teacher network, and T t-1 represents the previous parameters (unupdated parameters) of the teacher network, and the decay rate decay is usually fixedly set to 0.99.
[0034] In another aspect of the present invention, the present invention proposes a self-learning decay rate to improve the performance of the teacher model. "Self-learning" means that the decay rate is a learnable parameter or the output of a learnt network. In the present invention, a differentiable variable may be used as the decay rate, or the output of one fully connected layer may be used as the decay rate. In the latter case, for example, the fully connected layer may be set at the same level as the output layer of the second classifier 320_T so that the fully connected layer is connected to the layer immediately before the output layer in parallel with the output layer. Since the decay rate set in this way is not a fixed value and can adjust the ensembling speed according to the change in the performance of the model, the performance of knowledge distillation can be improved.
[0035] In another aspect of the present invention, data distillation based on a domain discriminator is proposed. Specifically, when training a classifier based on a cross-entropy loss function using source data, higher weights are given to source data that is similar to the target data among the source data. As a result, source data with a high degree of similarity to the target data can play a greater role in training, so that the classifier obtained by training can achieve better performance in the target domain.
[0036] The output of the domain discriminator can be used to determine whether the source data is similar to the target data. Since the domain discriminator can predict the probability that the current data is source data, the smaller this probability, the higher the degree of similarity between the current data and the target data. In other words, there is an inverse proportional relationship between the probability output by the domain discriminator and the degree of similarity. Therefore, the output of the domain discriminator can be used to weight the source data.
[0037] According to this principle, as shown in the following mathematical formula (9) or (10), a data distillation loss function L dd (not shown in FIG. 3) may be constructed, and the data distillation loss function L dd is also referred to as the fourth loss function.
Equation
Equation
[0038] Here, p s represents the probability that the prediction result is the true label when predicting the label for the source data. p d represents the probability that the source data determined by the domain discriminator is from the source domain, and 1 - p d or 1 / p d represents the weight assigned to the source data.
[0039] The probability p determined by the domain discriminator d is relatively small (which means that the similarity between the current source data and the target data is relatively high), 1 - p d or 1 / p d Since the value of becomes large, the weight given to the current source data becomes large. Therefore, the current source data (similar to the target data) can play a greater role in training.
[0040] Also, in another aspect of the present invention, the present invention further improves the structure of the self-ensembling teacher model shown in FIG. 2. FIG. 4 shows the improved network structure.
[0041] As shown in FIG. 4, the source data x Si and the target data x Ti are input not only to the student network but also to the teacher network. In contrast, in FIG. 2, only the target data x Ti is input to the teacher network. Therefore, the present invention performs distillation learning not only for the target domain but also for the source domain.
[0042] In FIG. 4, y Si represents the true label of the source data x Si , z Ti represents the probability predicted by the student network for the target data x Ti (that is, the probability that the target data x Ti belongs to each class), (Outer 27) TIFF0007694254000038.tif16170 represents the probability predicted by the teacher network for the target data x Ti , Z Si represents the probability predicted by the student network for the source data x Si (that is, the probability that the source data x Si belongs to each class), (Outer 28) TIFF0007694254000039.tif17170 represents the probability predicted by the teacher network for the source data x Si Further, the student network in FIG. 4 may include the first feature extractor 310 and the first classifier 320 shown in FIG. 3. The teacher network in FIG. 4 may include the second feature extractor 310_T and the second classifier 320_T shown in FIG. 3. Each of the above prediction probabilities may be generated by the first classifier 320 or the second classifier 320_T.
[0043] As shown in the following formula (11), based on the above prediction probabilities, a knowledge distillation loss function L kd (L in FIG. 3 kd-s and L kd-t included) may be constructed. The knowledge distillation loss function L kd is also referred to as the fifth loss function. [Number]
[0044] Here,[[]]END]] (Outer 29) TIFF0007694254000041.tif27170 represents the mean squared error of the probabilities predicted by the first classifier 320 and the second classifier 320_T respectively for the source data x Si and (Outer 30) TIFF0007694254000042.tif26170 represents the mean squared error of the probabilities predicted by the first classifier 320 and the second classifier 320_T respectively for the target data x Ti n represents the number of source data, and m represents the number of target data.
[0045] As shown in formula (12), based on the first loss function to the fifth loss function described above, a final loss function L for training the neural network shown in FIG. 3 may be constructed.
Number
[0046] Here, L c-s represents the classification cross - entropy loss function for source data and is the same as the loss function L c shown in FIG. 1. L adv represents the binary cross - entropy loss function of the domain discriminator and is the same as the loss function L adv shown in FIG. 1. The loss functions L c-s and L adv are known loss functions in the prior art, so their detailed descriptions are omitted in this specification.
[0047] Also, λ1 and λ2 in Equation (12) are weights for weighting the fourth loss function L kd and the fifth loss function L dd respectively, and may control the degree of action of the fourth and fifth loss functions in the training process. Specifically, the weight λ1 may be determined according to Equation (13).
Number
[0048] Here, p = step / total step , that is, the quotient obtained by dividing the current number of iteration steps by the total number of training steps, so p can represent the progress of training. α and n represent hyperparameters and may be set, for example, as α = 200 and n = 10. FIG. 5 shows a graph of the curve that changes as the number of training steps of the weight λ1 increases (assuming that the total number of training steps is 5000).
[0049] The weight λ2 may be determined according to Equation (14).
Number
[0050] Here, p has the same meaning as p in Equation (13). α and n represent hyperparameters, and for example, they may be set to α = 5 and n = 10. FIG. 6 is a diagram showing a curve that changes as the number of training steps of the weight λ2 increases (assuming that the total number of training steps is 5000).
[0051] As shown in FIGS. 5 and 6, at the beginning stage of training, since neither the prediction of the classifier nor the prediction of the domain discriminator is accurate, preferably, the values of λ1 and λ2 are set small. As the training progresses, since the predictions of the classifier and the domain discriminator of the teacher network gradually become accurate, the values of λ1 and λ2 may be gradually increased. Thereby, the knowledge distillation loss function L kd and the data distillation loss function L dd can play a greater role.
[0052] FIG. 7 is a flowchart showing a method for generating an optimal target data set according to the present invention. This method may be executed by the optimal target data set generation unit 960 in FIG. 9.
[0053] As shown in FIG. 7, in step S710, the first feature extractor 310 extracts features for the source data, and the first classifier 320 predicts the probability that the source data belongs to each of a plurality of predetermined classes based on the extracted features. The class corresponding to the maximum probability is determined as the label of the source data.
[0054] In step S720, the first feature extractor 310 extracts features for the target data, and the first classifier 320 predicts the probability that the target data belongs to each class based on the extracted features. The class corresponding to the maximum probability is determined as the first label of the target data.
[0055] In step S730, according to Equations (2) and (3), the second label of the target data is determined using the class center nearest neighbor algorithm.
[0056] In step S740, target data for which the determined first label and second label are the same is selected, and the first label or the second label is used as the pseudo label of the selected target data. Then, all the selected target data may constitute an optimal target data set.
[0057] FIG. 8 is a flowchart showing a method for training a domain adaptation neural network according to the present invention, and FIG. 9 is a block diagram showing a modular configuration of a training apparatus for a domain adaptation neural network according to the present invention.
[0058] As shown in FIG. 8, in step S810, according to equations (2), (4), and (5), based on the distance between the class center of the source data set and the class center of the optimal target data set, a first loss function L a (semantic alignment loss function) is constructed. This step may be executed by the first loss function generation unit 910 in FIG. 9.
[0059] In step S820, according to equation (6), based on the target data and its pseudo label in the optimal target data set, a second loss function (outer 31) TIFF0007694254000046.tif15170 (cross-entropy loss function) is constructed. This step may be executed by the second loss function generation unit 920 in FIG. 9.
[0060] In step S830, according to equation (7), for the source data in the source data set and the target data in the optimal target data set, a third loss function L con (contrastive learning loss function) is constructed. This step may be executed by the third loss function generation unit 930 in FIG. 9.
[0061] As can be seen from FIG. 9, the optimal target dataset generated by the method shown in FIG. 7 is used to construct the first to third loss functions.
[0062] Next, in step S840, according to formula (9) or (10), based on the probability output by the domain discriminator, a fourth loss function L dd (data distillation loss function) is constructed. This step may be executed by the fourth loss function generation unit 940 in FIG. 9.
[0063] In step S850, the second (teacher) feature extractor 310_T extracts the features of the source data and the target data, and the second (teacher) classifier 320_T predicts the labels of the source data and the target data. Next, in step S860, according to formula (11), based on the prediction results of the first classifier 320 and the prediction results of the second classifier 320_T, a fifth loss function L kd (knowledge distillation loss function) is constructed. Step S860 may be executed by the fifth loss function generation unit 950 in FIG. 9.
[0064] Next, in step S870, according to formula (12), based on the weighted combination of the first to fifth loss functions, the neural network is trained. This step may be executed by the training unit 970 in FIG. 9.
[0065] Note that the training method of the present invention does not necessarily need to be executed in the order shown in FIG. 8. For example, the order of generating the first to fifth loss functions may be different from that shown, or they may be generated simultaneously.
[0066] The inventor of the present invention conducted tests based on MNIST, USPS, and SVHN (all of which are known character datasets). Here, it includes domain adaptation in three directions: MNIST→USPS, USPS→MNIST, and SVHN→MNIST. Table 1 below shows a comparison of the performance between the solution means of the present invention and prior arts (such as ADDA and DANN). The values in Table 1 represent the classification accuracy rate, and the higher the accuracy rate, the better the performance of the solution means. As can be seen from Table 1, the performance of the solution means of the present invention is equivalent to or better than that of the prior arts.
Table 1
[0067] In particular, "source only" in Table 1 represents a method of training using only source data without using target data, that is, the simplest method, which is used as a comparison criterion. DANN (Domain-Adversarial Training of Neural Networks) represents the domain adversarial neural network shown in FIG. 1, and ADDA (Adversarial Discriminative Domain Adaptation) represents adversarial discriminative domain adaptation. CAT+RevGrad is described in the technical literature "Cluster Alignment with a Teacher for Unsupervised Domain Adaptation[C]", Deng Z et al., Proceedings of IEEE International Conference on Computer Vision, 2019:9944-9953.
[0068] The domain adaptation method according to the present invention can be applied to a wide range of domains. The following exemplifies and explains typical application scenarios.
[0069] [Application Scenario 1] Semantic segmentation Semantic segmentation means marking the parts representing different objects in an image with different colors. In the application scenarios of semantic segmentation, since the cost of manually labeling real-world images is very high, real-world images are hardly labeled. In this case, as an alternative method, training is performed using images of scenes in a simulation environment (such as a 3D game). Since automatic labeling of objects in the simulation environment can be easily realized by programming, labeled data can be easily obtained. In this way, a model is trained using the labeled data generated in the simulation environment, and the trained model is used to process images of the actual environment. However, since the simulation environment does not exactly match the actual environment, the performance of the model trained using the data of the simulation environment will drop significantly when processing images of the actual environment.
[0070] In this case, by using the domain adaptation method according to the present invention, training can be performed based on the labeled data of the simulation environment and the unlabeled data of the actual environment, so that the performance of the model when processing images of the actual environment can be improved.
[0071] [Application Scenario 2] Recognition of Handwritten Characters Handwritten characters generally include handwritten numbers, characters (such as Chinese, Japanese), etc. In the recognition of handwritten characters, commonly used labeled character sets include MNIST, USPS, SVHN, etc., and usually, a model is trained using these labeled character data. However, when applying the trained model to the recognition of actual (unlabeled) handwritten characters, the accuracy rate may decrease.
[0072] In this case, by using the domain adaptation method according to the present invention, training can be performed based on the labeled source data and the unlabeled target data, so that the performance of the model when processing the target data can be improved.
[0073] [Application Scenario 3] Classification and Prediction of Time-Series Data The prediction of time-series data includes, for example, the prediction of air pollution indices, the prediction of the length of stay (LOS) of ICU patients, the prediction of the stock market, etc. Taking the time-series data of particulate matter (PM2.5) indices as an example, a prediction model may be trained using a labeled training sample set. After the training is completed, the trained model may be applied to actual predictions. For example, based on the data of the previous 24 hours (unlabeled data) immediately before the current time, the range of PM2.5 indices three days later may be predicted.
[0074] In this scenario, by using the domain adaptation method according to the present invention, a model can be trained based on labeled data and unlabeled data, so that the prediction accuracy of the model can be improved.
[0075] [Application Scenario 4] Classification and Prediction of Table-Type Data Table-type data may include financial data such as online loan data. In this example, a prediction model may be constructed to predict whether there is a possibility of repayment delay in the way of loan combination, and the model may be trained using the method according to the present invention.
[0076] [Application Scenario 5] Image Recognition In the application scenario of image recognition or image classification, similar to the scenario of semantic segmentation, there is also a problem that the cost of labeling real-world image datasets is high. Therefore, in order to obtain a model with performance meeting the requirements, the domain adaptation method according to the present invention may be used to select a labeled dataset (such as ImageNet) as the source dataset, and training may be performed based on the source dataset and an unlabeled target dataset.
[0077] The embodiments of the present invention have been described with reference to specific examples. The methods according to the above examples may be implemented by software, hardware, or a combination of software and hardware. The programs included in the software may be pre-stored in a storage medium installed inside or outside the device. As an example, during execution, these programs are written into a random access memory (RAM) and executed by a processor (e.g., a CPU) to implement each process described herein.
[0078] FIG. 10 is a block diagram showing an exemplary configuration of the hardware of a computer capable of implementing the present invention, and the hardware of the computer is an example of a training device for a domain adaptation type neural network according to the present invention. Also, the domain adaptation type neural network according to the present invention may also be implemented based on the computer hardware.
[0079] As shown in FIG. 10, in the computer 1000, a central processing unit (CPU) 1001, a read-only memory (ROM) 1002, and a random access memory (RAM) 1003 are interconnected by a bus 1004.
[0080] An input / output interface 1005 is further connected to the bus 1004. Connected to the input / output interface 1005 are an input unit 1006 composed of a keyboard, a mouse, a microphone, etc., an output unit 1007 composed of a display, a speaker, etc., a storage unit 1008 composed of a hard disk, a non-volatile memory, etc., a communication unit 1009 composed of a network interface card (local area network (LAN) card, modem, etc.), and a driver 1010 for driving a removable medium 1011. The removable medium 1011 is, for example, a magnetic disk, an optical disk, a magneto-optical disk, or a semiconductor memory.
[0081] In the computer having the above configuration, the CPU 1001 loads the program stored in the storage unit 1008 into the RAM 1003 via the input / output interface 1005 and the bus 1004, and executes the above method by executing the program.
[0082] The program executed by the computer (CPU 1001) may be recorded on a movable medium 1011 which is a package medium. The package medium is formed by, for example, a magnetic disk (including a floppy disk), an optical disk (including a compact disk read only memory (CD-ROM), a digital versatile disk (DVD), etc.), a magneto-optical disk, or a semiconductor memory. Also, the program executed by the computer (CPU 1001) may be provided via a wired or wireless transmission medium of a local area network, the Internet, or digital satellite broadcasting.
[0083] When the movable medium 1011 is installed in the driver 1010, the program can be installed in the storage unit 1008 via the input / output interface 1005. Also, the program is received by the communication unit 1009 via a wired or wireless transmission medium and installed in the storage unit 1008. Alternatively, the program may be installed in the ROM 1002 or the storage unit 1008 in advance.
[0084] The program executed by the computer may be a program that executes processing according to the order described in this specification, or a program that executes processing in parallel, or a program that executes processing as needed (for example, at the time of a call).
[0085] The devices or units described in this specification are logical and not limited to physical devices or entities. For example, the functions of each unit described in this specification may be implemented by multiple physical entities, or the functions of multiple units described in this specification may be implemented by a single physical entity. Also, features, components, elements, steps, etc. described in one embodiment are not limited to that embodiment, and for example, may be applied to other embodiments, or may be used instead of or in combination with specific features, components, elements, steps, etc. of other embodiments.
[0086] The scope of the present invention is not limited to the specific embodiments described herein. As can be understood by those skilled in the art, various modifications or changes may be made to the embodiments herein without departing from the principles and gist of the present invention according to design requirements and other factors. The scope of the present invention is limited by the appended claims and their equivalents.
[0087] Also, with respect to the embodiments including each of the above-described embodiments, the following supplementary notes are disclosed, but are not limited thereto. (Supplementary Note 1) A method for training a domain adaptation type neural network executed by a computer, wherein the domain adaptation type neural network includes a first feature extraction unit, a first classification unit, and an identification unit, the computer includes a memory storing instructions and a processor, and when the instructions are executed by the processor, the processor is caused to execute the method, and the method includes: a step in which the first feature extraction unit extracts a first feature for source data in a labeled source data set; a step in which the first classification unit predicts the probability that the source data belongs to each of a plurality of classes based on the first feature; a step in which the first feature extraction unit extracts a second feature for target data in an unlabeled target data set; The step in which the first classification unit predicts the probability that the target data belongs to each class based on the second feature, and determines the class corresponding to the maximum probability as the first label of the target data; The step of calculating the distance between the class center for each class of the source data set and the features of the target data, and determining the class corresponding to the closest class center as the second label of the target data; The step of selecting target data from the target data set for which the determined first label and second label are the same, wherein the first label or the second label is used as the pseudo-label of the selected target data; The step of calculating the class center for each class of the target data set based on the selected target data; The step of constructing a first loss function based on the distance between the class center of the source data set and the class center of the calculated target data set; The step of constructing a second loss function based on the selected target data and its pseudo-label; The step of constructing a third loss function for the source data and the selected target data in the source data set; The method includes the step of training the domain adaptation type neural network based on the first loss function, the second loss function, and the third loss function. (Appendix 2) The step in which the identification unit determines the probability that the currently input data is source data based on the first feature and the second feature; The step of constructing a fourth loss function based on the probability determined by the identification unit; The method according to Appendix 1, including the step of training the domain adaptation type neural network based on the fourth loss function. (Appendix 3) The method according to Supplementary Note 2, wherein the fourth loss function is constructed based on the reciprocal of the probability determined by the identification unit and one of the differences obtained by subtracting the probability determined by the identification unit from 1. (Supplementary Note 4) The domain adaptation type neural network further includes a second feature extraction unit and a second classification unit, The method includes a step of the second feature extraction unit extracting a third feature from the source data; a step of the second classification unit predicting the probability that the source data belongs to each class based on the third feature; a step of the second feature extraction unit extracting a fourth feature from the target data; a step of the second classification unit predicting the probability that the target data belongs to each class based on the fourth feature; a step of constructing a fifth loss function based on the probability predicted by the first classification unit and the probability predicted by the second classification unit; a step of training the domain adaptation type neural network based on the fifth loss function, the method according to Supplementary Note 2. (Supplementary Note 5) The method according to Supplementary Note 4, wherein the fifth loss function is constructed based on the mean squared error of the probabilities predicted by each of the first classification unit and the second classification unit for the source data, and the mean squared error of the probabilities predicted by each of the first classification unit and the second classification unit for the target data. (Supplementary Note 6) The parameters of the second feature extraction unit are the exponential moving average of the parameters of the first feature extraction unit, and the parameters of the second classification unit are the exponential moving average of the parameters of the first classification unit. When obtaining the decay rate used in the exponential moving average, a differentiable variable is used as the decay rate, or a fully connected layer is used to generate the decay rate. The method according to Supplementary Note 4, wherein the fully connected layer is set to be connected to the layer immediately before the output layer in parallel with the output layer of the second classification unit. (Supplementary Note 7) Training the domain adaptation type neural network based on a weighted combination of the first loss function, the second loss function, the third loss function, the fourth loss function, and the fifth loss function, The method according to Supplementary Note 4, wherein the weights of the fourth loss function and the fifth loss function are gradually increased as training is executed. (Supplementary Note 8) The method according to Supplementary Note 1, wherein the second loss function is a cross-entropy loss function for training the first classification unit. (Supplementary Note 9) The discrimination unit is connected to the first feature extraction unit via a gradient reversal unit, The method according to Supplementary Note 1, wherein the discrimination unit and the first feature extraction unit operate adversarially to each other. (Supplementary Note 10) The domain adaptation type neural network is used to perform image recognition, and the source data and the target data are image data, or The domain adaptation type neural network is used to process financial data, and the source data and the target data are table type data, or The method according to Supplementary Note 1, wherein the domain adaptation type neural network is used to process environmental weather data or medical data, and the source data and the target data are time series data. (Supplementary Note 11) An apparatus for training a domain adaptation type neural network, The domain adaptation type neural network is A first feature extraction unit that extracts first features for source data in a labeled source data set and extracts second features for target data in an unlabeled target data set, Based on the first feature, predict the probability that the source data belongs to each of a plurality of classes, and based on the second feature, predict the probability that the target data belongs to each of the classes, and determine the class corresponding to the maximum probability as the first label of the target data, a first classification unit; An identification unit that determines the probability that currently input data is source data based on the first feature and the second feature; The apparatus A memory in which a program is stored; One or more processors; By executing the program, the processor Calculating the distance between the class center for each class of the source data set and the features of the target data, and determining the class corresponding to the class center with the closest distance as the second label of the target data; Selecting target data for which the determined first label and second label are the same from the target data set, where the first label or the second label is used as the pseudo-label of the selected target data; Calculating the class center for each class of the target data set based on the selected target data; Constructing a first loss function based on the distance between the class center of the source data set and the class center of the calculated target data set; Constructing a second loss function based on the selected target data and its pseudo-label; Constructing a third loss function for the source data and the selected target data in the source data set; Training the domain adaptation type neural network based on the first loss function, the second loss function, and the third loss function. (Appendix 12) An apparatus for training a domain adaptation neural network, wherein the domain adaptation neural network comprises a first feature extraction unit that extracts a first feature for source data in a labeled source dataset and extracts a second feature for target data in an unlabeled target dataset; a first classification unit that predicts, based on the first feature, the probability that the source data belongs to each of a plurality of classes, predicts, based on the second feature, the probability that the target data belongs to each of the classes, and determines, as a first label of the target data, the class corresponding to the maximum probability; and an identification unit that determines, based on the first feature and the second feature, the probability that currently input data is source data; the apparatus further comprises an optimal target dataset generation unit that calculates the distance between the class center for each class in the source dataset and the features of the target data, determines, as a second label of the target data, the class corresponding to the closest class center, selects, from the target dataset, target data for which the determined first label and the second label are the same, and forms an optimal target dataset, where the first label or the second label is used as the pseudo-label of the selected target data; a first loss function generation unit that calculates, based on the target data in the optimal target dataset, the class center for each class in the target dataset, and constructs a first loss function based on the distance between the class center of the source dataset and the calculated class center of the target dataset; a second loss function generation unit that constructs a second loss function based on the target data and its pseudo-label in the optimal target dataset; A third loss function generation unit that constructs a third loss function for the source data in the source data set and the target data in the optimal target data set, An apparatus that executes a training unit that trains the domain adaptation type neural network based on the first loss function, the second loss function, and the third loss function. (Appendix 13) A storage medium storing a program for training a domain adaptation type neural network, wherein the domain adaptation type neural network includes a first feature extraction unit, a first classification unit, and an identification unit, and when the program is executed by a computer, the computer The step of the first feature extraction unit extracting first features for the source data in the labeled source data set; The step of the first classification unit predicting the probability that the source data belongs to each of a plurality of classes based on the first features; The step of the first feature extraction unit extracting second features for the target data in the unlabeled target data set; The step of the first classification unit predicting the probability that the target data belongs to each of the classes based on the second features and determining the class corresponding to the maximum probability as the first label of the target data; Calculating the distance between the class center for each class of the source data set and the features of the target data, and determining the class corresponding to the class center with the closest distance as the second label of the target data; Selecting target data for which the determined first label and second label are the same from the target data set, wherein the first label or the second label is used as the pseudo label of the selected target data; Calculating the class center for each class of the target data set based on the selected target data; Constructing a first loss function based on the distance between the class center of the source data set and the class center of the calculated target data set; Constructing a second loss function based on the selected target data and its pseudo label; Constructing a third loss function for the source data and the selected target data in the source data set; Training the domain adaptation type neural network based on the first loss function, the second loss function, and the third loss function. A storage medium for executing a method including these steps.
Claims
1. A method for training a domain adaptation neural network executed by a computer, wherein the domain adaptation neural network includes a first feature extraction unit, a first classification unit, and an identification unit, and the computer includes a memory storing instructions and a processor, and when the instructions are executed by the processor, the processor is caused to execute the method, and the method includes: a step in which the first feature extraction unit extracts first features for source data in a labeled source data set; a step in which the first classification unit predicts, based on the first features, the probability that the source data belongs to each of a plurality of classes; a step in which the first feature extraction unit extracts second features for target data in an unlabeled target data set; a step in which the first classification unit predicts, based on the second features, the probability that the target data belongs to each of the classes, and determines, as a first label of the target data, the class corresponding to the maximum probability; a step of calculating the distance between the class center for each class of the source data set and the features of the target data, and determining, as a second label of the target data, the class corresponding to the class center with the closest distance; a step of selecting, from the target data set, target data for which the determined first label and second label are the same, wherein the first label or the second label is used as a pseudo-label of the selected target data; a step of calculating, based on the selected target data, the class center for each class of the target data set; a step of constructing a first loss function based on the distance between the class center of the source data set and the class center of the calculated target data set; a step of constructing a second loss function based on the selected target data and its pseudo-label; For the source data and the selected target data in the source data set, constructing a third loss function; training the domain adaptation type neural network based on the first loss function, the second loss function, and the third loss function, including; The first loss function is a semantic alignment loss function for minimizing the distance between the class center of the source data set and the class center of the calculated target data set; The second loss function is a cross-entropy loss function for training the first classification unit; The third loss function is a contrastive learning loss function for making the distance between features within the same class small and the distance between features of different classes large. Claim 2 The identification unit determines the probability that the currently input data is source data based on the first feature and the second feature; constructing a fourth loss function based on the probability determined by the identification unit; training the domain adaptation type neural network based on the fourth loss function, including the method according to claim 1. Claim 3 Constructing the fourth loss function based on one of the reciprocal of the probability determined by the identification unit and the difference obtained by subtracting the probability determined by the identification unit from 1, according to the method of claim 2. Claim 4 The domain adaptation type neural network further includes a second feature extraction unit and a second classification unit; The method includes; the second feature extraction unit extracts a third feature for the source data; the second classification unit predicts the probability that the source data belongs to each class based on the third feature; The step in which the second feature extraction unit extracts a fourth feature from the target data; The step in which the second classification unit predicts the probability that the target data belongs to each class based on the fourth feature; The step of constructing a fifth loss function based on the probability predicted by the first classification unit and the probability predicted by the second classification unit; The method according to claim 2, comprising the step of training the domain adaptation type neural network based on the fifth loss function.
5. The fifth loss function is constructed based on the mean squared error of the probabilities predicted by each of the first classification unit and the second classification unit for the source data, and the mean squared error of the probabilities predicted by each of the first classification unit and the second classification unit for the target data. The method according to claim 4.
6. The parameters of the second feature extraction unit are the exponential moving average of the parameters of the first feature extraction unit, and the parameters of the second classification unit are the exponential moving average of the parameters of the first classification unit. When obtaining the decay rate used in the exponential moving average, a differentiable variable is used as the decay rate, or a fully connected layer is used to generate the decay rate. The method according to claim 4, wherein the fully connected layer is set to be connected to the layer immediately before the output layer in parallel with the output layer of the second classification unit.
7. Training the domain adaptation type neural network based on a weighted combination of the first loss function, the second loss function, the third loss function, the fourth loss function, and the fifth loss function. The method according to claim 4, wherein the weights of the fourth loss function and the fifth loss function are gradually increased as the training is executed.
8. An apparatus for training a domain adaptation type neural network, The domain adaptation type neural network a first feature extraction unit that extracts a first feature for source data in a labeled source data set and extracts a second feature for target data in an unlabeled target data set; based on the first feature, predicting the probability that the source data belongs to each of a plurality of classes, and based on the second feature, predicting the probability that the target data belongs to each of the classes, and determining the class corresponding to the maximum probability as the first label of the target data; a first classification unit; an identification unit that determines the probability that currently input data is source data based on the first feature and the second feature; The apparatus includes a memory in which a program is stored; and one or more processors. By executing the program, the processor calculating the distance between the class center for each class of the source data set and the feature of the target data, and determining the class corresponding to the class center with the closest distance as the second label of the target data; selecting target data for which the determined first label and the second label are the same from the target data set, where the first label or the second label is used as the pseudo-label of the selected target data; calculating the class center for each class of the target data set based on the selected target data; constructing a first loss function based on the distance between the class center of the source data set and the calculated class center of the target data set; constructing a second loss function based on the selected target data and its pseudo-label; Constructing a third loss function for the source data and the selected target data in the source data set; Training the domain adaptation neural network based on the first loss function, the second loss function, and the third loss function; The first loss function is a semantic alignment loss function for minimizing the distance between the class centers of the source data set and the calculated class centers of the target data set; The second loss function is a cross-entropy loss function for training the first classification unit; The third loss function is a contrastive learning loss function for making the distance between features within the same class small and the distance between features of different classes large. Claim 9 A storage medium storing a program for training a domain adaptation neural network, the domain adaptation neural network including a first feature extraction unit, a first classification unit, and an identification unit, and when the program is executed by a computer, causing the computer to execute the method according to any one of claims 1 to 7.
Citation Information
Patent Citations
Knowledge transfer method, information processing apparatus, and storage medium
JP2019215861A