Self-supervised cross-domain few-sample classification method based on iterative pruning

Through iterative pruning and self-supervised learning methods, the problem of the difference between pre-training data and downstream task data in transfer learning is solved, and better performance in cross-domain small sample classification tasks is achieved.

CN120105147APending Publication Date: 2025-06-06NANJING UNIV OF INFORMATION SCI & TECH
View PDF 0 Cites 0 Cited by

Patent Information

Application Number
CN202510109257.9
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2025-01-23
Publication Date
2025-06-06

AI Technical Summary

Technical Problem

The prior art ignores the differences between pre-trained data and downstream task data in transfer learning tasks, resulting in poor performance of the model in target domain tasks.

Method used

The self-supervised cross-domain small sample classification method based on iterative pruning is used to pre-train the large amount of labeled source domain data, and then retrain the small amount of labeled target domain data, and optimize the model through pruning operations until a classification model with excellent performance on the target domain task is obtained.

Benefits of technology

The full utilization of the tagless target domain data is achieved, and the pruning operation cuts redundant parameters, realizes model balance between the source domain and the target domain, and improves the performance of cross-domain small sample classification.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN120105147A_ABST
    Figure CN120105147A_ABST
Patent Text Reader

Abstract

The invention discloses a self-supervised cross-domain few-sample classification method based on iterative pruning. The method comprises the following steps: step 1, training a classification model through labeled source domain data of a large data volume; 2, the classification model is retrained through small-data-volume unlabeled target domain data; step 3, performing pruning operation on the retrained classification model in the step 2 according to the parameter assignment of the classification model; 4, repeating the steps 2 and 3 until a classification model with excellent target domain task performance is obtained; and 5, finely adjusting the classifier in the classification model obtained in the step 4, and carrying out cross-domain few-sample classification through the fine-adjusted classification model. According to the method, model pruning and a comparative learning method are effectively combined, full utilization of label-free target domain data can be achieved, and meanwhile model balance between a source domain and a target domain can be achieved while redundant parameters are cut.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention relates to the technical field of cross-domain few-sample classification, and in particular to a self-supervised cross-domain few-sample classification method based on iterative pruning. Background Art

[0002] Deep learning methods have shown great superiority in various fields. One of the reasons for the success of deep learning algorithms is the ability to obtain large-scale annotated data. However, in some scenarios, it may be difficult or expensive to collect and annotate a large number of data samples, and directly training the model with raw data samples cannot achieve significant performance. Recently, a new research topic called cross-domain few-shot learning has been proposed to solve the above problems. Cross-domain few-shot learning first extracts features from the source domain containing rich samples, and then uses the trained feature extractor to solve tasks in unseen domains containing only a few annotated samples. Cross-domain few-shot learning has recently attracted great attention due to its outstanding performance in solving problems such as domain shift, heterogeneous data, and missing data.

[0003] Traditional cross-domain few-shot classification methods are committed to extracting informative features in the pre-training stage. For example, the Meta-FT method first uses a graph neural model to train a high-performance feature extractor, and then introduces mechanisms such as meta-learning to post-process the extracted features. The ATA method uses adversarial learning ideas to enhance learning tasks and increase the "difficulty" of task learning, effectively improving the performance of pre-trained models. The StyleAdv method uses adversarial ideas to perform adversarial learning on the high-order semantic information of the data, effectively improving the generalization ability of the model, and enabling the model to achieve excellent performance in the target domain. The NASE method uses the mechanism of image generation to extract features from the source domain data. First, the encoder is used to encode the data, and then the decoder is used to restore the features. This method can fully extract the features of the source domain data.

[0004] Although the above methods can achieve relatively excellent target domain task performance, they ignore the existence of some unlabeled target domain data during model training. In order to make full use of target domain samples, some works use self-supervised learning technology to extract data features. These methods can achieve better performance than the above supervised learning-based methods. Representative works include STARTUP, CLD-FD, and DynDistill. In addition, due to the rise of contrastive learning methods, Oh et al. used contrastive learning methods to solve the problem of cross-domain few-sample learning. Oh et al. also systematically analyzed the advantages and disadvantages of different contrastive learning methods in cross-domain few-sample tasks. Although the above methods can make full use of unlabeled target domain data, they only focus on optimizing the pre-trained model by updating the model parameters in the process of extracting data features, and solve downstream tasks by retraining. The structure of the model is fixed during the parameter update process of the above methods, so the mismatch between the source domain pre-trained model and the target domain task is ignored. Summary of the invention

[0005] In view of the deficiencies in the prior art, the present invention provides a self-supervised cross-domain few-sample classification method based on iterative pruning to solve the technical problem in the prior art that the difference between pre-training data and downstream task data in transfer learning tasks is too large.

[0006] The present invention provides a self-supervised cross-domain few-sample classification method based on iterative pruning, comprising the following steps:

[0007] Step 1: Train the classification model using large amounts of labeled source domain data;

[0008] Step 2: Retrain the classification model using a small amount of unlabeled target domain data;

[0009] Step 3: Prune the classification model retrained in step 2 according to the parameter assignment of the classification model;

[0010] Step 4: Repeat steps 2 and 3 until a classification model with excellent performance for the target domain task is obtained;

[0011] Step 5: Fine-tune the classifier in the classification model obtained in step 4, and perform cross-domain few-sample classification using the fine-tuned classification model.

[0012] Furthermore, in step 1, the model is trained by means of supervised loss guidance, and the specific process is as follows:

[0013]

[0014] Where W represents the model parameter; W s Represents data from the source domain Model parameters obtained by feature extraction on It is a loss of supervision.

[0015] Furthermore, the supervised loss is a cross entropy loss.

[0016] Furthermore, in step 2, the model is trained by constructing a contrastive learning loss function, and the specific process is as follows:

[0017]

[0018] in,

[0019] In the formula, aug i (x) is the i-th data augmentation output of input x; f is the initialization parameter W s The feature extractor of is; h is the projection head; It is a pre-trained model for unlabeled target domain dataset; l sim is the self-supervised loss function.

[0020] Furthermore, in step 3, the specific method of the pruning operation is:

[0021] The model is pruned according to the absolute value of the classification model. The specific pruning method is:

[0022]

[0023] Among them, the hard threshold cutoff for:

[0024]

[0025] Where n is the total number of parameters in each layer of the model.

[0026] Furthermore, in step 5, only the last layer classifier in the classification model obtained in step 4 is retrained to achieve fine-tuning.

[0027] Beneficial effects of the present invention:

[0028] The present invention effectively combines model pruning with contrastive learning methods, which can fully utilize unlabeled target domain data, and can achieve model balance between source domain and target domain while trimming redundant parameters. Compared with other methods, the present invention can still achieve better performance without introducing additional modules. The present invention can be deployed in different cross-domain few-sample classification tasks based on self-supervised learning in a plug-and-play manner. BRIEF DESCRIPTION OF THE DRAWINGS

[0029] The features and advantages of the present invention will be more clearly understood by referring to the accompanying drawings, which are schematic and should not be construed as limiting the present invention in any way. In the accompanying drawings:

[0030] Figure 1 is a flow chart of a specific embodiment of the present invention;

[0031] Figure 2 It is a process schematic diagram of a specific embodiment of the present invention;

[0032] Figure 3 is a visualization result of the existing pre-training in a specific embodiment of the present invention;

[0033] Figure 4 is a visualization result of the existing BYOL method in a specific embodiment of the present invention;

[0034] Figure 5 It is the visualization result of the method of the present invention in a specific embodiment of the present invention. DETAILED DESCRIPTION

[0035] In order to make the purpose, technical solution and advantages of the embodiments of the present invention clearer, the technical solution in the embodiments of the present invention will be clearly and completely described below in conjunction with the drawings in the embodiments of the present invention. Obviously, the described embodiments are part of the embodiments of the present invention, not all of the embodiments. Based on the embodiments of the present invention, all other embodiments obtained by those skilled in the art without creative work are within the scope of protection of the present invention.

[0036] The present invention is further illustrated below in conjunction with specific embodiments. Those skilled in the art should understand that these embodiments are only used to illustrate the present invention and are not used to limit the scope of the present invention, and modifications to various equivalent forms of the present invention fall within the scope defined by the appended claims of this application.

[0037] like Figure 1 , 2 The present invention provides a self-supervised cross-domain few-sample classification method based on iterative pruning, comprising the following steps:

[0038] Step 1: Train the classification model using large amounts of labeled source domain data;

[0039] First, a large amount of labeled source domain data is used to pre-train the classification model to obtain a pre-trained model with good performance. The classification model can use ResNet10. The process can be expressed as:

[0040]

[0041] Where W represents the model parameter; W s Represents data from the source domain Model parameters obtained by feature extraction on It is the supervised loss, and the supervised loss is preferably the cross entropy loss;

[0042] Step 2: Retrain the classification model using a small amount of unlabeled target domain data. Specifically, the model is trained by constructing a contrastive learning loss function. The specific process is as follows:

[0043]

[0044] in,

[0045] In the formula, aug i (x) is the i-th data augmentation output of input x; f is the initialization parameter W s The feature extractor of is; h is the projection head; It is a pre-trained model for unlabeled target domain dataset; l sim is the self-supervised loss function, that is, the contrastive learning loss function, which can make z 1 With z 2 as close as possible, while making z i With negative samples z - Stay away as much as possible;

[0046] Step 3: Prune the classification model retrained in step 2 according to the parameter assignment of the classification model;

[0047] After the comparative training in step 2, a classification model with slightly better performance for downstream targets and tasks is obtained. In order to improve the performance of the classification model for downstream target domain tasks, the classification model at this time is pruned:

[0048] First, all parameters are sorted according to their absolute values, and then pruning is performed on parameters with smaller amplitudes, and the top-k parameters are retained, that is,

[0049] For hard threshold cutoff The definition is as follows:

[0050]

[0051] Set k = n × (1-p), where n is the total number of parameters in each layer of the model. During the pruning process, only the deep parameters of the model are pruned, and the shallow parameters of the model are ignored.

[0052] Step 4: Repeat steps 2 and 3 until a classification model with excellent performance for the target domain task is obtained;

[0053] The pruning process will have a certain impact on the performance of the classification model. In order to eliminate this impact, the pruned classification model is updated again, that is, steps 2 and 3 are repeated; during the repetition process, the pruned elements also need to participate in the model training, and the initial value is 0 to participate in the model training. Repeat the above operation several times to obtain a classification model with stable performance.

[0054] Step 5: Fine-tune the classifier in the classification model obtained in step 4, and perform cross-domain few-sample classification using the fine-tuned classification model.

[0055] Finally, after the pre-training and retraining steps, the model is fine-tuned using labeled target domain data: the remaining operation values ​​are fixed, and only the parameters of the last layer classifier are retrained and updated. During the retraining and updating process, the loss function represents the supervision loss. The expression form of the loss function is the same as the loss function in step 1. The differences are mainly reflected in the following two aspects: First, the data at this stage is mainly sampled from the target domain; second, the data at this stage is mainly input into the model in the form of a few samples. The few samples used in this specific embodiment use the N-way K-shot form to fine-tune the model, that is, the total number of samples for this fine-tuning is N×K.

[0056] like Figure 3-5 As shown, the experimental results show that the present invention can achieve very significant advantages compared with the existing methods.

[0057] Although the embodiments of the present invention have been described in conjunction with the accompanying drawings, those skilled in the art may make various modifications and variations without departing from the spirit and scope of the present invention, and such modifications and variations are all within the scope defined by the appended claims.

Claims

1. A self-supervised cross-domain few-shot classification method based on iterative pruning, characterized in that: The steps include: Step 1: Train the classification model using large amounts of labeled source domain data; Step 2: Retrain the classification model using a small amount of unlabeled target domain data; Step 3: Prune the classification model retrained in step 2 according to the parameter assignment of the classification model; Step 4: Repeat steps 2 and 3 until a classification model with excellent performance for the target domain task is obtained; Step 5: Fine-tune the classifier in the classification model obtained in step 4, and perform cross-domain few-sample classification using the fine-tuned classification model.

2. The self-supervised cross-domain few-sample classification method based on iterative pruning according to claim 1, characterized in that: In step 1, the model is trained by means of supervised loss guidance, and the specific process is as follows: Where W represents the model parameter; W s Represents data from the source domain Model parameters obtained by feature extraction on It is a loss of supervision.

3. The self-supervised cross-domain few-sample classification method based on iterative pruning according to claim 2, characterized in that: The supervised loss is the cross entropy loss.

4. The self-supervised cross-domain few-sample classification method based on iterative pruning according to claim 1, characterized in that: In step 2, the model is trained by constructing a contrastive learning loss function, and the specific process is as follows: in, In the formula, aug i (x) is the i-th data augmentation output of input x; f is the initialization parameter W s The feature extractor of is; h is the projection head; It is a pre-trained model for unlabeled target domain dataset; l sim is the self-supervised loss function.

5. The self-supervised cross-domain few-sample classification method based on iterative pruning according to claim 1, characterized in that: In step 3, the specific method of the pruning operation is: The model is pruned according to the absolute value of the classification model. The specific pruning method is: Among them, the hard threshold cutoff for: Where n is the total number of parameters in each layer of the model.

6. The self-supervised cross-domain few-sample classification method based on iterative pruning according to claim 1, characterized in that: In step 5, only the last layer classifier in the classification model obtained in step 4 is retrained to achieve fine-tuning.