Semi-Supervised Image Classification Method and System Based on Pseudo-Label and Embedding Clustering Matching
By introducing dynamic threshold adjustment strategy and embedded cluster matching module in semi-supervised image classification, the problem of confirmation bias and neglected relationship between samples is solved, and more efficient and accurate image classification is achieved.
Patent Information
- Application Number
- CN202411705872.8
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2024-11-26
- Publication Date
- 2025-05-30
- Estimated Expiration
- 2044-11-26
AI Technical Summary
The existing semi-supervised learning methods have problems of confirmation bias and ignoring the relationship between samples in image classification.
A semi-supervised image classification method based on pseudo-label and embedded cluster matching is proposed. A dynamic threshold adjustment strategy is introduced through the prediction matching module to alleviate confirmation bias, and a target map is constructed through the graph matching module to strengthen the consistency of the relationship between samples using the clustering results.
It effectively alleviates the confirmation bias, improves training efficiency, and enhances the consistency of local distribution relationships of data points in the feature space through the graph matching module, and improves the accuracy of image classification.
Smart Images

Figure CN119579992B_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to the technical field of image processing, and in particular to a semi-supervised image classification method and system based on pseudo-label and embedding clustering matching. Background Art
[0002] Pseudo-label and consistency regularization are two main technical routes in the field of semi-supervised learning. Combining pseudo-label and consistency regularization with data augmentation is a current new trend. However, these two methods still have obvious deficiencies: the pseudo-label method cannot avoid confirmation bias because the model trained by this method tends to maintain or amplify the generated predictions, even if they are noisy predictions; while the existing consistency regularization methods only emphasize applying the consistency assumption on a single instance and ignore the relationship between aligned samples. Summary of the Invention
[0003] To solve the above problems, the present invention proposes a semi-supervised image classification method and system based on pseudo-label and embedding clustering matching, which includes two parts: prediction matching and graph matching. Prediction matching applies consistency regularization at the level of a single sample, and a dynamic threshold adjustment strategy is introduced in this module to alleviate confirmation bias. The threshold is relatively low at the beginning and gradually increases as the training process progresses. Graph matching constructs a contrast graph and a target graph based on the clustering results of different augmented sample embeddings, and achieves consistency constraints at the level of sample affinity relationships by optimizing the graph matching loss.
[0004] To achieve the above object, the present invention adopts the following technical solutions:
[0005] In a first aspect, the present invention provides a semi-supervised image classification method based on pseudo-label and embedding clustering matching, including:
[0006] Obtain labeled images and unlabeled images; for unlabeled images, perform weak augmentation and strong augmentation to obtain weakly augmented images and strongly augmented images;
[0007] For labeled images, calculate the supervised loss based on a semi-supervised image classification model;
[0008] Input the augmented unlabeled images into the semi-supervised image classification model, generate a dynamic threshold based on the model's learning effect on categories to filter pseudo-labels, and calculate the unsupervised loss;
[0009] Map the weakly augmented images and strongly augmented images into the embedding space respectively for K-Means clustering. The clustering results are respectively used to generate a target graph and a strongly augmented embedding clustering graph, and calculate the graph matching loss;
[0010] Establish a total loss based on the supervised loss, unsupervised loss, and graph matching loss, and improve the performance of the semi-supervised image classification model by minimizing the total loss.
[0011] Preferably, for the labeled images, the supervised loss is calculated based on a semi-supervised image classification model, which specifically includes: weakly augmenting the labeled images, and calculating the supervised loss between the prediction of the weakly augmented labeled images input to the semi-supervised image classification model and their corresponding true labels.
[0012] Preferably, a dynamic threshold is generated based on the learning effect of the model on categories to filter pseudo-labels, which specifically includes:
[0013] In this round, the unlabeled images are input into the semi-supervised image classification model to obtain predicted classification labels;
[0014] Using the dynamic threshold updated in the previous round, filter the predicted classification labels to obtain the pseudo-labels for this round;
[0015] Evaluate the learning effect of the model on categories according to the pseudo-labels and predicted classification labels for this round;
[0016] Normalize the learning effect and use it to update the dynamic threshold for the next round.
[0017] Preferably, the process of using the dynamic threshold updated in the previous round to filter the predicted classification labels to obtain the pseudo-labels for this round specifically includes:
[0018]
[0019] wherein, is the discriminant function, is the unlabeled image 's maximum class probability prediction, that is, the pseudo-label for this round; is the dynamic threshold updated in the previous round; the filtering process associates the threshold with the learning effect of the model on each class, and is used to provide a lower threshold at the initial stage of training and gradually increase the threshold as the training process progresses.
[0020] Preferably, the process of evaluating the learning effect of the model on categories according to the pseudo-labels and predicted classification labels for this round specifically includes:
[0021]
[0022] wherein, reflects the learning effect of the model on class c at time step t; is the predicted vector of the weakly augmented unlabeled image by the model, is the unlabeled image 's maximum class probability prediction; the symbol represents a fixed threshold, which is used to select samples with high confidence to train the model; Indicates a pseudo-label belonging to class c; is the discriminant function; is an unlabeled sample The probability on class c.
[0023] Preferably, normalizing the learning effect and using it to update the dynamic threshold for the next round specifically includes:
[0024]
[0025] Among them, for perform maximum normalization to obtain ;
[0026]
[0027] Among them, scale the predefined threshold with the normalized class learning effect , and generate a dynamic threshold for each class at each time step t; to avoid generating extremely low resulting in many noisy samples participating in training, further adjust as follows:
[0028]
[0029] Among them, C represents the number of classes.
[0030] Preferably, mapping the weakly augmented image and the strongly augmented image to the embedding space respectively for K-Means clustering, and using the clustering results to generate the target graph and the strongly augmented embedding clustering graph respectively, and calculating the graph matching loss specifically includes:
[0031] Map the weakly augmented image and the strongly augmented image to the embedding space respectively, and perform K-Means clustering respectively; the clustering result of the weakly augmented image is used to generate the target graph, and the clustering result of the strongly augmented image is used to generate the strongly augmented embedding clustering graph, and calculate the graph matching loss between the target graph and the strongly augmented embedding clustering graph.
[0032] In a second aspect, the present invention provides a semi-supervised image classification system based on pseudo-labels and embedding clustering matching, including:
[0033] A data augmentation module, configured to obtain labeled images and unlabeled images; for unlabeled images, perform weak augmentation and strong augmentation to obtain weakly augmented images and strongly augmented images;
[0034] A prediction matching module, configured to calculate a supervised loss for labeled images based on a semi-supervised image classification model; input the augmented unlabeled images into the semi-supervised image classification model, generate a dynamic threshold based on the learning state of the model for each class to filter pseudo-labels, and calculate an unsupervised loss;
[0035] A graph matching module, which is used to map the weakly enhanced image and the strongly enhanced image into the embedding space respectively for K-Means clustering. The clustering results are respectively used to generate the target graph and the strongly enhanced embedding clustering graph, and calculate the graph matching loss.
[0036] A loss calculation module, which is used to establish the total loss based on the supervised loss, the unsupervised loss and the graph matching loss, and improve the performance of the semi-supervised image classification model by minimizing the total loss.
[0037] In a third aspect, the present invention provides a computer-readable storage medium, on which a computer program is stored. When the program is executed by a processor, the steps in a semi-supervised image classification method based on pseudo-label and embedding clustering matching described in the first aspect are implemented.
[0038] In a fourth aspect, the present invention provides a computer device, including a memory, a processor and a computer program stored on the memory and executable on the processor. When the processor executes the program, the steps in a semi-supervised image classification method based on pseudo-label and embedding clustering matching described in the first aspect are implemented.
[0039] Compared with the prior art, the beneficial effects of the present invention are as follows:
[0040] The semi-supervised learning framework ClusMatch proposed by the present invention includes two components: a prediction matching module and a graph matching module, which can simultaneously strengthen the prediction consistency at the instance level and the similarity consistency at the graph level. The prediction matching module is used to generate pseudo-labels, thereby strengthening the prediction consistency at the instance level. A simple dynamic threshold adjustment strategy is introduced in this module, which is used to provide a lower threshold at the initial stage of training, and then gradually increase the threshold as the training process progresses, so as to adaptively screen sufficient and high-quality samples according to the training process, which can greatly improve the training efficiency. The graph matching module can utilize the relationship between data points to generate the target graph and the comparison graph based on the clustering results of the low-dimensional embedding, so as to construct the graph matching target and make the local distribution relationship of the data points in the feature space consistent under different perturbations. In addition, introducing pseudo-labels when constructing the target graph can improve the accuracy of the target, thereby promoting graph learning.
[0041] The advantages of the additional aspects of the present invention will be partially given in the following description, partially become obvious from the following description, or be understood through the practice of the present invention. BRIEF DESCRIPTION OF THE DRAWINGS
[0042] The specification drawings constituting a part of the present disclosure are used to provide a further understanding of the present disclosure. The schematic embodiments of the present disclosure and their descriptions are used to explain the present disclosure and do not constitute a limitation to the present disclosure.
[0043] Figure 1 The main flowchart of a semi-supervised image classification method based on pseudo-label and embedded clustering matching provided by an embodiment of the present invention;
[0044] Figure 2 The framework schematic diagram of ClusMatch provided by an embodiment of the present invention;
[0045] Figure 3 The schematic diagram of dynamic threshold adjustment provided by an embodiment of the present invention;
[0046] Figure 4 The schematic diagram of the performance comparison of ClusMatch with FixMatch and FlexMatch in terms of confidence threshold provided by an embodiment of the present invention;
[0047] Figure 5 The schematic diagram of the performance comparison of ClusMatch with FixMatch and FlexMatch in terms of sampling rate provided by an embodiment of the present invention;
[0048] Figure 6 The schematic diagram of the performance comparison of ClusMatch with FixMatch and FlexMatch in terms of the correctness rate of pseudo-labels provided by an embodiment of the present invention;
[0049] Figure 7 The schematic diagram of the performance comparison of ClusMatch with CCSSL and CoMatch in terms of Top-1 correctness rate provided by an embodiment of the present invention;
[0050] Figure 8 The schematic diagram of the performance comparison of ClusMatch with CCSSL and CoMatch in terms of the loss about the graph provided by an embodiment of the present invention;
[0051] Figure 9 The schematic diagram of the performance comparison of ClusMatch with CCSSL and CoMatch in terms of the overall loss provided by an embodiment of the present invention. Detailed implementation manners
[0052] The present invention will be further described below in conjunction with the accompanying drawings and embodiments.
[0053] Technical term explanations
[0054] (1) Pseudo-label based semi-supervised learning: It is the simplest and most effective semi-supervised method. Its core idea is to use self-generated predictions as proxy labels for unlabeled data to expand the labeled training set. The model trained by this method cannot correct its own errors. Even though the threshold can filter out some low-confidence noisy labels, high-confidence samples may also contain systematic errors. These errors accumulate continuously and will eventually result in confident but incorrect pseudo-labels on unlabeled sample points, which is the confirmation bias.
[0055] (2) Consistency regularization: Based on the consistency assumption that the model should have the same or similar outputs for the same input before and after perturbation. In other words, consistency regularization encourages similar samples to have similar predictions. Most current methods only perform consistency regularization in the probability space, while ignoring the consistency of the local relationship between sample points before and after perturbation.
[0056] (3) Graph-based semi-supervised learning: Data points are usually represented as nodes in a graph, and the edges connecting each pair of nodes reflect their similarity or distance. Based on the popular assumption that high-dimensional data can be represented as a low-dimensional manifold, two types of semi-supervised graph learning methods have been derived from this assumption. One is label propagation, which propagates labels from labeled data to unlabeled data according to the manifold structure of the data and the similarity of intermediate nodes. This type of method is too dependent on the quality of nodes. If there is noise in neighboring nodes, it may reduce the effect of graph learning. The other is node embedding, whose purpose is to learn a low-dimensional vector representation for each node. These low-dimensional embeddings can effectively capture the features of nodes and the structural information of the graph, thus being more beneficial for downstream tasks. However, the similarity measurement methods between nodes for this type of method are usually computationally complex.
[0057] Example 1
[0058] As Figure 1 shown, this example discloses a semi-supervised image classification method based on pseudo-label and embedding clustering matching, including the following steps:
[0059] S1: Obtain labeled images and unlabeled images; for unlabeled images, perform weak augmentation and strong augmentation to obtain weakly augmented images and strongly augmented images;
[0060] S2: For labeled images, calculate the supervised loss based on the semi-supervised image classification model;
[0061] S3: Input the augmented unlabeled images into the semi-supervised image classification model, generate a dynamic threshold based on the learning effect of the model on categories to filter pseudo-labels, and calculate the unsupervised loss;
[0062] S4: Map the weakly augmented image and the strongly augmented image to the embedding space respectively for K-Means clustering. The clustering results are respectively used to generate the target graph and the strongly augmented embedding clustering graph, and calculate the graph matching loss.
[0063] S5: Establish the total loss based on the supervised loss, the unsupervised loss, and the graph matching loss, and improve the performance of the semi-supervised image classification model by minimizing the total loss.
[0064] Next, combined with Figure 2 , a semi-supervised image classification method based on pseudo-label and embedding clustering matching disclosed in this embodiment will be described in detail.
[0065] This embodiment proposes a semi-supervised learning framework - ClusMatch, which can effectively solve the above problems. ClusMatch consists of a prediction matching module and a graph matching module. The former is designed to generate pseudo-labels during the iterative process and provide classification predictions during the inference process. The latter uses the pseudo-labels and the clustering results on the embedding set to construct an embedding graph, which reflects the similarity relationship between samples. By using a variable threshold to obtain the classification loss, it can not only accelerate the model convergence but also alleviate the confirmation bias. At the same time, a simple method for measuring sample similarity is defined, and graph learning is carried out based on this to achieve multi-level consistency regularization.
[0066] The architecture diagram of ClusMatch provided in this embodiment is as Figure 2 shown. For the convenience of elaboration, the basic settings of the semi-supervised classification problem are introduced here first.
[0067] For a semi-supervised image classification task with C classes, usually let denote one batch of labeled data, where is the label of the b-th image . denotes one batch of unlabeled data, where µ is a hyperparameter used to control the relative size of X and U. A convolutional neural network-based encoder is used to extract the feature r, and at the same time a fully connected classifier is used to generate predictions for the images. General threshold-based semi-supervised learning methods, such as FixMatch, aim to obtain an encoder and a classifier with good performance, which is achieved by optimizing a supervised loss on the labeled samples and an unsupervised loss on the unlabeled samples.
[0068]
[0069] Among them, represents the cross-entropy function, Indicates weak augmentation, such as flipping or cropping.
[0070] To obtain the unsupervised loss, FixMatch applies one weak augmentation and one strong augmentation to the unlabeled data, and then calculates the cross-entropy between the maximum probability prediction (i.e., pseudo-label) of the weak augmentation and the prediction of the strongly augmented version. According to the consistency assumption, the pseudo-label should be similar or consistent with the strongly augmented prediction, so the unsupervised loss is defined in the following form:
[0071]
[0072] where and are abbreviations of and respectively; . The symbol represents the threshold used to select samples with high confidence for training the model. In algorithms using a fixed threshold, usually is set, which is a very high value. Although it can filter out most low-quality samples, it also results in low data utilization, especially in the initial stage of training, making the model training cycle very long.
[0073] (I) Prediction Matching Module
[0074] By mapping the image features to the probability space to obtain the pseudo-label, and then minimizing the cross-entropy between the pseudo-label and the prediction of the strongly augmented version of the same sample to achieve prediction alignment. This alignment applies the consistency assumption at the level of a single sample and can promote the encoder to extract features that are robust to perturbations. Therefore, the quality of the pseudo-label is crucial for the model training effect. To effectively improve the quality of the pseudo-label, this embodiment proposes a prediction matching module.
[0075] First, following FixMatch, this embodiment defines the supervised loss as the cross-entropy between the true label and the model prediction.
[0076] After that, to obtain the unsupervised loss, a predefined high threshold is introduced in FixMatch to screen out the pseudo-labels. Although it can ensure that the pseudo-labels are confident, it also results in a slow training speed.
[0077] Therefore, this embodiment introduces a dynamic threshold adjustment strategy to improve Equation (2). This strategy associates the threshold with the learning effect of the model for each class, and can provide a lower threshold at the beginning of training, and then gradually increase the threshold as the training process progresses. This can not only ensure that the samples selected in each iteration are relatively confident to reduce confirmation bias, but also prevent most samples from being discarded at the beginning of training to ensure training efficiency.
[0078] The process of dynamic class threshold adjustment is as Figure 3 shown. At time step t, the model first outputs the classification predictions of a batch of unlabeled data, and then uses the most recently updated dynamic class threshold to filter out low-confidence samples (represented by gray dots). The learning effect of each class is estimated as the mean of the sum of the probabilities of the samples belonging to that class and exceeding the dynamic class threshold. Finally, the learning effect of each class is maximally normalized, and the fixed threshold is scaled to obtain the dynamic class threshold to be used in the next time step.
[0079] Specifically, this embodiment uses a heuristic method to quantify the learning effect of the model for each class under the current threshold condition, and then uses the quantization value to scale the fixed threshold to obtain the new threshold for a specific class. In this module, the mean of the sum of the probabilities of all pseudo-labels that exceed the threshold and belong to class c on class c is used to represent the learning state of the model for class c, formalized as:
[0080]
[0081] where, reflects the learning effect of class c at time step t, is the prediction vector of the weak augmentation of the model for the unlabeled sample , is the unlabeled sample , is the maximum class probability prediction of the unlabeled sample , that is, the pseudo-label; the symbol represents the fixed threshold for selecting high-confidence samples to train the model; represents the pseudo-label belonging to class c; is the discriminant function;
[0082]
[0083] To smooth misclassifications and extreme samples, more precisely represent the confidence of the model for each class, and thus bring more accurate adjustment, this embodiment performs maximum normalization on
[0083]
[0084] However, in this case, the situation of may occur, especially in the initial stage of training, which may lead to a large number of low-quality samples participating in the training. Therefore, in this embodiment, a minimum limit is added to the threshold:
[0085]
[0086] where C represents the number of categories.
[0087] The unsupervised loss of ClusMatch calculated using this threshold can be expressed as:
[0088]
[0089] It can be seen that only when the maximum prediction probability of a sample exceeds the dynamic threshold of its corresponding category, the cross-entropy loss of the sample is calculated. In the initial stage of training, since the model has a poor learning effect on each category, the dynamic threshold is relatively low, allowing more samples to participate in the training and ensuring the training efficiency; as the training progresses, the learning effect of the model on each category gradually improves, and the dynamic threshold also gradually increases, thus ensuring that the samples selected in each iteration are relatively confident (i.e., have a high prediction probability).
[0090] (2) Graph Matching Module
[0091] To further alleviate the confirmation bias, graph learning is introduced in this embodiment. The clustering results on the low-dimensional embedding of the image features are used to generate a target graph and a contrast graph to construct a graph matching objective.
[0092] The similarity relationship between samples is represented by an affinity graph, aiming to apply consistency regularization in the low-dimensional embedding space. According to the popular hypothesis, high-dimensional data is roughly located on a low-dimensional manifold, so this embodiment generalizes the consistency hypothesis to the embedding space. The consistency hypothesis requires that similar samples have similar outputs. This embodiment considers that the relative distribution of a sample point to other sample points in the embedding space is similar to the relative distribution of its similar points to other sample points in the embedding space.
[0093] Formally, a batch of unlabeled data U has a strongly augmented version and a weakly augmented version . Assume that represents the relationship vector between the sample point of sample and , and at the same time represents the relationship vector between the sample point of sample and . Then, the distance between and should be very small.
[0094] For the measurement of similarity between samples, traditional methods use cosine similarity with high computational complexity or utilize the inner product of vectors. These two methods focus on fine-grained information and are easily affected by boundary values.
[0095] This embodiment believes that on a balanced dataset, focusing on the global distribution structure of samples in the embedding space while ignoring unimportant details may produce better results. Therefore, this embodiment expects to have a method that is simple, emphasizes the class boundaries in the embedding space more, and can ensure the numerical stability when calculating the loss.
[0096] Based on the above expectations, considering that K-Means clustering perfectly meets the above requirements, the division of clusters not only reflects the intra-class aggregation of samples of the same class but also can generate clear class boundaries, which helps to enhance the reliability of pseudo-labels. At the same time, clustering can suppress noise samples and reduce the negative impact of incorrect pseudo-labels on model training. In addition, since the number of classes C is known, this embodiment does not need to conduct additional exploration on the number of clusters k, and k is directly determined by C.
[0097] Specifically, this embodiment sets the number of clusters K = C, and conducts K-Means clustering on from and on from . The strong enhancement relationship graph constructed based on the clustering results on can be written as:
[0098]
[0099] where, represents the matrix element in . For two strongly enhanced samples belonging to the same cluster, this embodiment sets the edge between them to 1 and the rest to 0. That is, for any pair of embedding vectors and of strongly enhanced samples, if their clustering labels are the same, the similarity between them is 1, otherwise it is 0.
[0100] In graph matching, this embodiment hopes to use the weak enhancement graph as the target to guide generation. The ultimate goal is to obtain a more discriminative encoder and a more robust and efficient mapping head. Therefore, the more accurate is, the more it can improve the effect of graph learning. So this embodiment combines the cluster label and pseudo-label information to construct the target matrix , which is formalized as:
[0101]
[0102] Among them, represents the matrix element in is the low-dimensional embedding vector of the weakly augmented sample and is the low-dimensional embedding vector of the weakly augmented sample For any pair of embedding vectors of weakly augmented samples and , if their clustering labels are the same and their respective pseudo-labels exceed their corresponding dynamic class thresholds, the similarity between them is 1, otherwise it is 0.
[0103] The loss of the graph matching module is calculated by minimizing the cross-entropy loss between the two graphs in equations (9) and (10). It can be defined as:
[0104]
[0105] Since and the elements in are 0 or 1, binary cross-entropy is selected as H in this embodiment.
[0106] As Figure 2 shown, in the ClusMatch framework of this embodiment, given a batch of unlabeled data, the weakly augmented images enter a prediction matching module based on dynamic thresholds to generate pseudo-labels. The pseudo-labels are not only used to calculate the classification loss, but also combined with clustering on the weakly augmented embeddings to generate an accurate target graph. The graph matching module strengthens the consistency of the relationships between samples by minimizing the cross-entropy between the target graph and the strongly augmented embedding clustering graph.
[0107] Finally, the ClusMatch provided in this embodiment jointly optimizes a supervised loss , an unsupervised loss , and a graph matching loss . The overall training objective can be written as:
[0108]
[0109] Among them, and are both weight coefficients of the loss.
[0110] (III) Experiments
[0111] CIFAR10 is a balanced dataset consisting of 60K 32×32 color images from 10 classes. It contains 50K training images and 10K test images. In this embodiment, experiments are performed in a class-balanced manner under settings with 40, 250, and 4000 labels respectively.
[0112] CIFAR100 consists of 60K 32×32 color images and contains 100 classes. Each class has 500 images for training and 100 images for testing. In this embodiment, experiments are performed in a class-balanced manner under settings with 400, 2500, and 10000 labels respectively.
[0113] STL10 is derived from ImageNet and consists of 100K unlabeled images and 13K labeled images. The labeled data contains 10 classes, and each class provides 500 training images and 800 test images. All images in STL10 are 96×96 color images.
[0114] ImageNet-1k is a subset of the ImageNet dataset. It contains 1k classes, and each class has approximately 1.3k training images and 50 validation images.
[0115] Baseline methods: In this embodiment, methods based on pseudo-labeling and consistency regularization in the past five years are considered, including MixMatch, ReMixMatch, UDA, FixMatch, FlexMatch, Dash, FreeMatch, SimMatch. In addition, several state-of-the-art methods involving graph regularization, such as CoMatch, CCSSL, and CHMatch, are also included in the comparison in this embodiment. It should be noted that the method in this embodiment is implemented based on CCSSL. The code in this embodiment is integrated into SemiCLS.
[0116] Experimental settings: For CIFAR10, WideResNet-28-2 is used in this embodiment, and WideResNet-28-8 is used for CIFAR100. For STL10, ResNet-18 is used in this embodiment. A 2-layer MLP mapping head is used to obtain 64-dimensional embeddings. SGD with a momentum of 0.9 and a weight decay of 0.001 is used to optimize the model. All methods use an initial learning rate of 0.03 and a cosine learning rate decay strategy, and are trained for 1024 epochs. In all experiments on these three datasets, this embodiment uniformly sets , , , ,K , ,where μ is the sample coefficient and B is the batch size. is the unsupervised loss coefficient, is the graph loss coefficient, and K is a parameter of the k-means algorithm.
[0117] For ImageNet-1k, 1% or 10% of the labeled images are randomly sampled in a class-balanced manner, and the remaining images are unlabeled. ResNet-50 is used as the encoder, and a two-layer MLP is used as the projection head to obtain 128-dimensional embeddings. Except for B = 160 and τ = 0.6, other hyperparameters on ImageNet-1k are the same as those on the previous three datasets.
[0118] As a specific implementation manner, Figure 4 , Figure 5 , Figure 6 reflects the performance comparison of ClusMatch and two threshold-based methods, FixMatch and FlexMatch, on CIFAR10-40. They are the curves of the confidence threshold, sampling rate, and pseudo-label accuracy during the training process. As can be seen from Figure 4 and Figure 5 , ClusMatch shows superior thresholds and sampling rates, indicating that the dynamic threshold adjustment strategy is effective. As shown in Figure 6 , after 100k steps, the pseudo-label accuracies of the three methods fluctuate within a small range of 0.88 to 0.9.
[0119] As a specific implementation manner, Figure 7 , Figure 8 , Figure 9 reflects the performance of ClusMatch and two graph-based semi-supervised methods, CCSSL and CoMatch, on CIFAR100-400. They are the curves of the Top-1 accuracy, graph loss, and overall loss during the training process. As shown in Figure 7 , due to the different ways of calculating the graph loss, ClusMatch shows a higher graph loss compared with CCSSL and CoMatch. As shown in Figure 8 and Figure 9 , ClusMatch achieves the highest accuracy and the lowest overall loss, which proves the effectiveness of the graph matching module described in this embodiment.
[0120] Table 1 below shows the error rates of ClusMatch provided in this embodiment and other benchmark methods under different label numbers settings for the three small datasets of CIFAR10, CIFAR100, and STL10. ClusMatch achieved the optimal performance under four settings and was close to the optimal method in the other three settings. Even in the case of extremely sparse labels, ClusMatch still maintained a very stable effect, indicating that ClusMatch is very suitable for handling classification tasks on small datasets.
[0121]
[0122]
[0123] Table 2 below shows the error rates of ClusMatch provided in this embodiment on ImageNet-1K. Under the settings of 1% and 10% labeled numbers, ClusMatch generally outperformed CoMatch, indicating that ClusMatch can handle classification tasks on large datasets.
[0124]
[0125] In this embodiment, WideResNet-28-2 was used to conduct experiments on the CIFAR10 dataset with 40 labeled samples to verify the effectiveness of each component. The results are shown in Table 3. According to the comparison results of the error rates in the 5th and 6th rows of the table, the dynamic threshold adjustment scheme proposed in this embodiment is significantly better than the predefined threshold. Both the 1st and 2nd rows combine a fixed threshold with the graph matching module of this embodiment, and the difference is whether to introduce pseudo-label information to fine-tune the target graph; among them, the error rate of the 2nd row is lower than that of the 1st row, indicating that it is necessary to use pseudo-labels to fine-tune the target graph. The 3rd row shows the combination of the dynamic threshold adjustment scheme and the graph matching module, and the 4th row is ClusMatch proposed in this embodiment, which introduces pseudo-labels on the basis of the 3rd row. Their comparison also illustrates the importance of fine-tuning the target graph with pseudo-labels. The error rate of the 4th row is lower than that of the 5th row, indicating the promoting effect of the graph matching module of this embodiment on correct classification.
[0126]
[0127] This embodiment addresses the challenges in semi-supervised image classification by proposing a ClusMatch semi-supervised learning framework that integrates two major modules: prediction matching and graph matching. ClusMatch not only strengthens the prediction consistency at the instance level but also ensures the similarity consistency at the graph level. The prediction matching module introduces a dynamic threshold adjustment strategy, setting a lower threshold at the initial stage of training and gradually increasing it to adapt to the training process, effectively alleviating the confirmation bias and significantly improving the training efficiency. The graph matching module cleverly utilizes the relationships between data points to construct a target graph and a contrast graph based on the clustering results of low-dimensional embeddings, ensuring the consistency of the local distribution relationships of data points in the feature space under different perturbations. Additionally, by introducing pseudo-labels, the accuracy of the target graph is further improved, promoting the effect of graph learning. The proposed ClusMatch framework not only provides a new solution idea for semi-supervised image classification but also demonstrates significant beneficial effects in improving classification accuracy and training efficiency.
[0128] Embodiment Two
[0129] This embodiment provides a semi-supervised image classification system based on pseudo-labels and embedding clustering matching, including:
[0130] A data augmentation module for obtaining labeled images and unlabeled images; for unlabeled images, performing weak augmentation and strong augmentation to obtain weakly augmented images and strongly augmented images;
[0131] A prediction matching module for, for labeled images, calculating the supervised loss based on a semi-supervised image classification model; inputting the augmented unlabeled images into the semi-supervised image classification model, generating a dynamic threshold based on the model's learning effect on categories to filter pseudo-labels, and calculating the unsupervised loss;
[0132] A graph matching module for respectively mapping the weakly augmented images and strongly augmented images to the embedding space for K-Means clustering, and using the clustering results to generate a target graph and a strongly augmented embedding clustering graph respectively, and calculating the graph matching loss;
[0133] A loss calculation module for establishing a total loss based on the supervised loss, unsupervised loss, and graph matching loss, and improving the performance of the semi-supervised image classification model by minimizing the total loss.
[0134] Embodiment Three
[0135] This embodiment provides a computer-readable storage medium with a computer program stored thereon, and when the program is executed by a processor, it implements the steps in a semi-supervised image classification method based on pseudo-labels and embedding clustering matching as described in Embodiment One above.
[0136] Embodiment Four
[0137] This embodiment provides a computer device, including a memory, a processor, and a computer program stored on the memory and executable on the processor. When the processor executes the program, it implements the steps in a semi-supervised image classification method based on pseudo-label and embedding clustering matching as described in the above-mentioned Embodiment 1.
[0138] The steps or modules involved in the above Embodiments 2 to 4 correspond to those in Embodiment 1. For specific implementation manners, reference may be made to the relevant description part of Embodiment 1. The term "computer-readable storage medium" should be understood to include a single medium or multiple media including one or more instruction sets; it should also be understood to include any medium that can store, encode, or carry an instruction set for execution by a processor and enable the processor to execute any method in the present invention.
[0139] The above are only the preferred embodiments of the present invention and are not intended to limit the present invention. For those skilled in the art, the present invention may have various changes and modifications. Any modification, equivalent replacement, improvement, etc. made within the spirit and principle of the present invention shall be included within the protection scope of the present invention.
Claims
1. A semi-supervised image classification method based on pseudo-label and embedded cluster matching, characterized in that: include: Obtain labeled images and unlabeled images; For unlabeled images, weak enhancement and strong enhancement are performed to obtain weakly enhanced images and strongly enhanced images; For labeled images, the supervised loss is calculated based on the semi-supervised image classification model; The enhanced unlabeled images are input into the semi-supervised image classification model, and a dynamic threshold is generated based on the model's learning effect on the category to filter out pseudo labels and calculate the unsupervised loss; The weakly enhanced image and the strongly enhanced image are respectively mapped to the embedding space for K-Means clustering, and the clustering results are respectively used to generate the target graph and the strongly enhanced embedded clustering graph, and the graph matching loss is calculated; A total loss is established based on supervised loss, unsupervised loss and graph matching loss, and the performance of the semi-supervised image classification model is improved by minimizing the total loss.
2. A semi-supervised image classification method based on pseudo-label and embedded cluster matching as claimed in claim 1, characterized in that: The method of calculating the supervised loss for the labeled image based on the semi-supervised image classification model specifically includes: weakly enhancing the labeled image, and calculating the supervised loss between the prediction of the input weakly enhanced labeled image and its corresponding true label based on the semi-supervised image classification model.
3. A semi-supervised image classification method based on pseudo-label and embedded cluster matching as claimed in claim 1, characterized in that: The dynamic threshold is generated based on the learning effect of the model on the category to filter the pseudo-labels, specifically including: In this round, the unlabeled images are input into the semi-supervised image classification model to obtain the predicted classification labels; Using the dynamic threshold updated in the previous round, the predicted classification labels are screened to obtain the pseudo labels for this round; Based on the pseudo labels and predicted classification labels of this round, evaluate the learning effect of the model on the category in this round; The learning effect is normalized and used to update the dynamic threshold for use in the next round.
4. A semi-supervised image classification method based on pseudo-label and embedded cluster matching as claimed in claim 3, characterized in that: The method uses the dynamic threshold updated in the previous round to filter the predicted classification labels to obtain the pseudo labels of this round, which specifically includes: in, is the discriminant function, is an unlabeled image The maximum class probability prediction of , that is, the pseudo label of this round; is the dynamic threshold updated in the previous round; the screening process associates the threshold with the learning effect of the model on each class, which is used to provide a lower threshold in the early stage of training and gradually increase the threshold as the training process progresses.
5. A semi-supervised image classification method based on pseudo-label and embedded cluster matching as claimed in claim 3, characterized in that: The learning effect of the model on the category is evaluated based on the pseudo labels and predicted classification labels of this round, specifically including: in, It reflects the learning effect of the model on class c at time step t; is the model for unlabeled images The weakly enhanced prediction vector, is an unlabeled image The maximum class probability prediction of Represents a fixed threshold, which is used to select samples with high confidence to train the model; represents the pseudo label belonging to class c; is the discriminant function; It is an unlabeled sample The probability of being in class c.
6. A semi-supervised image classification method based on pseudo-label and embedded cluster matching as claimed in claim 3, characterized in that: The learning effect is normalized and used to update the dynamic threshold for use in the next round, specifically including: Among them, Perform maximum normalization to obtain ; Among them, the normalized category learning effect To scale the predefined threshold , generates a dynamic threshold for each category at each time step t; in order to avoid generating extremely low This results in many noise samples being involved in training. Further adjustments are made as follows: Among them, C represents the number of categories.
7. A semi-supervised image classification method based on pseudo-label and embedded cluster matching as claimed in claim 1, characterized in that: The weakly enhanced image and the strongly enhanced image are respectively mapped to the embedding space for K-Means clustering, and the clustering results are respectively used to generate the target graph and the strongly enhanced embedded clustering graph, and calculate the graph matching loss, specifically including: The weakly enhanced image and the strongly enhanced image are mapped to the embedding space respectively, and K-Means clustering is performed respectively; the clustering result of the weakly enhanced image is used to generate the target graph, and the clustering result of the strongly enhanced image is used to generate the strongly enhanced embedded clustering graph, and the graph matching loss between the target graph and the strongly enhanced embedded clustering graph is calculated.
8. A semi-supervised image classification system based on pseudo-label and embedded cluster matching, characterized in that: include: Data augmentation module, used to obtain labeled images and unlabeled images; For unlabeled images, weak enhancement and strong enhancement are performed to obtain weakly enhanced images and strongly enhanced images; A prediction matching module is used to calculate the supervision loss for the labeled image based on the semi-supervised image classification model; The enhanced unlabeled images are input into the semi-supervised image classification model, and a dynamic threshold is generated based on the model's learning effect on the category to filter out pseudo labels and calculate the unsupervised loss; A graph matching module is used to map the weakly enhanced image and the strongly enhanced image to the embedding space for K-Means clustering, and the clustering results are used to generate the target graph and the strongly enhanced embedded clustering graph, and calculate the graph matching loss; The loss calculation module is used to establish the total loss based on the supervised loss, unsupervised loss and graph matching loss, and improve the performance of the semi-supervised image classification model by minimizing the total loss.
9. A computer-readable storage medium having a computer program stored thereon, characterized in that: When the program is executed by a processor, the steps in a semi-supervised image classification method based on pseudo-labels and embedded cluster matching as described in any one of claims 1 to 7 are implemented.
10. A computer device comprising a memory, a processor, and a computer program stored in the memory and executable on the processor, characterized in that: When the processor executes the program, the steps in the semi-supervised image classification method based on pseudo-label and embedded cluster matching as described in any one of claims 1-7 are implemented.
Citation Information
Patent Citations
Network supervision fine-grained image recognition method and system based on deep learning
CN115496948A
Semi-supervised image classification method and semi-supervised image classification system
CN116894985A