Supervised prototype contrastive learning method based on prototype generator

By adopting a supervised prototype contrastive learning method based on prototype generators, the problem of insufficient integration of global contextual information in medical image segmentation is solved, which improves the segmentation performance and robustness of the model, especially showing stronger generalization ability in complex scenes.

CN120726335BActive Publication Date: 2025-11-18TIANJIN UNIV
View PDF 2 Cites 0 Cited by

Patent Information

Application Number
CN202511249856.7
Authority / Receiving Office
CN · China
Patent Type
Patents(China)
Current Assignee / Owner
Filing Date
2025-09-03
Publication Date
2025-11-18
Estimated Expiration
2045-09-03

AI Technical Summary

Technical Problem

Existing medical image segmentation techniques struggle to effectively integrate global contextual information when processing complex medical images, resulting in insufficient contextual modeling capabilities. This leads to significant performance bottlenecks, particularly in scenarios with uneven data distribution and small sample learning.

Method used

A supervised prototype contrastive learning method based on prototype generator is adopted. Features are extracted through backbone network to generate pixel-wise cross-entropy loss and prototype contrastive loss. Combined with weighted mask attention Transformer and multilayer perceptron, class prototypes are generated and stored in prototype memory. The joint loss function is calculated to improve segmentation performance.

Benefits of technology

It significantly improves the performance of medical image segmentation models, enhances the attention to local and global information, improves the class imbalance problem, and enhances segmentation results and robustness, especially showing stronger generalization ability in complex scenes.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN120726335B_ABST
    Figure CN120726335B_ABST
Patent Text Reader

Abstract

The application discloses a supervised prototype contrast learning method based on a prototype generator, and solves the technical problem that it is difficult to effectively integrate global context when a complex medical image is processed in the prior art. It comprises the following steps: a backbone network is used to extract features from an input image; after feature extraction, the network generates two calculation branches, which are respectively used to calculate a pixel-by-pixel cross-entropy loss and calculate a prototype contrast loss; a prototype generator is constructed by using a downsampling process, a weighted mask attention Transformer and an MLP to generate a class prototype; the generated class prototype is stored in a prototype memory, and the prototype contrast loss is calculated by calculating the similarity between the pixel embedding and the stored prototype; and a joint loss function is formed by combining the pixel-by-pixel cross-entropy loss and the prototype contrast loss. The application can effectively capture the global attributes of the embedding space, accurately reflect the internal structure of the training data, and better process the problems of uneven intra-class pixel distribution and neglect of spatial information.
Need to check novelty before this filing date? Find Prior Art

Description

TECHNICAL FIELD

[0001] The present application relates to the technical field of medical image segmentation, and particularly relates to a supervised prototype contrast learning method based on a prototype generator. BACKGROUND

[0002] In recent years, with the rapid development of deep learning technology, medical image segmentation has become an important research direction in medical image analysis. The goal of medical image segmentation is to accurately separate the target regions (such as organs, lesions, etc.) in medical images from the background, so as to facilitate subsequent diagnosis and treatment. However, existing image segmentation techniques still face some challenges and shortcomings in dealing with complex medical images.

[0003] The U-shaped network (U-Net) is a widely used deep learning architecture in the field of medical image segmentation. The design of this architecture aims to perform end-to-end pixel-level segmentation, with good performance. However, one of the main drawbacks of the U-shaped network is its spatial invariance, which limits its ability to model the context information between pixels. Context information is crucial for image segmentation tasks, as pixels of the same class in medical images usually have similar features and are spatially related to each other.

[0004] To enhance the context modeling ability of the U-shaped network, researchers have developed various context aggregation modules. For example, dilated convolution can increase the receptive field to capture a larger range of context information; spatial pyramid pooling can obtain context features at different scales through multi-scale feature fusion; and multi-level feature fusion technology tries to combine features at different levels to improve segmentation accuracy. These methods have improved the learning ability of local context to some extent, but still fail to effectively integrate global context, i.e., the semantic relationship between different images. Learning global context is particularly important in medical image segmentation, especially when dealing with multiple diseases and different cases.

[0005] In the case of uneven data distribution, the model may tend to learn the dominant class and ignore other classes, which significantly affects the segmentation effect. Currently, learning global context usually requires a large amount of computational resources and data. For example, in some cases, in order to extract global context features, training on a large dataset may be required, which undoubtedly increases the training cost and time. In addition, in the scenario of small sample learning, the model often has difficulty in learning enough global features, resulting in performance bottlenecks.

[0006] In general, the current medical image segmentation field is facing challenges such as insufficient context modeling capability, class imbalance, and computational resource consumption. Although existing various context aggregation modules improve the modeling capability of local context to some extent, it is still necessary to effectively integrate global context information to improve the overall performance of the model. SUMMARY

[0007] The purpose of the present application is to provide a supervised prototype contrast learning method based on a prototype generator to solve the technical problem that the prior art is difficult to effectively integrate global context when processing complex medical images.

[0008] To achieve the above-mentioned purpose, the present application provides the following technical solutions:

[0009] The present application provides a supervised prototype contrast learning method based on a prototype generator, comprising the following steps:

[0010] S1, using a backbone network to extract features from an input image to generate a pixel representation R∈R H×W×C , which lays the foundation for subsequent calculations, wherein R represents image features, H and W are the spatial resolution of R, and C represents the pixel dimension;

[0011] S2, after feature extraction, the backbone network generates two calculation branches for calculating the pixel-wise cross-entropy loss and calculating the prototype contrast loss, respectively, wherein the cross-entropy loss measures the difference between the model prediction and the true label, and the prototype contrast loss helps the model learn the global context and the semantic relationship between pixels;

[0012] S3, using a down-sampling process, a weighted mask attention Transformer, and a multi-layer perception mechanism to build a prototype generator to generate class prototypes from the extracted features;

[0013] S4, storing the generated class prototypes in a prototype memory, calculating the prototype contrast loss by calculating the similarity between the pixel embedding and the stored prototype;

[0014] S5, combining the pixel-wise cross-entropy loss and the prototype contrast loss to ensure that the model can effectively learn the pixel and class information, forming a joint loss function for training the image segmentation model, thereby realizing efficient and accurate medical image segmentation.

[0015] Further, step S3 comprises the following steps:

[0016] S31, using down-sampling to reduce the resolution of the feature map to obtain R'∈R H'×W'×C' , wherein R' represents the image features after down-sampling, H' and W' are the spatial resolution of R', and C' represents the pixel dimension after down-sampling;

[0017] S32. The weighted masked attention Transformer calculates the output from the downsampled image features and the learnable query features. Extract global information of categories, where Q represents the global information of all categories in the image features, and k represents the number of categories. The encoding length of global information for each category;

[0018] S33. Output the weighted masked attention result. Input into a multilayer perceptron to generate class prototypes Where N represents the number of types and C represents the dimension of the prototype vector, which is the same as the pixel dimension of the image feature R extracted by S1.

[0019] Furthermore, in step S31, the downsampling process includes applying two 3×3 convolutions and a 2×2 pooling operation four times.

[0020] Further, in step S32, the weighted mask attention Transformer includes an encoder and a decoder. The encoder is used to convert the input sequence into a feature representation, and the decoder is used to generate an output sequence based on the feature representation generated by the encoder. The main structures of the encoder and decoder include a multi-head self-attention mechanism and a feedforward neural network. The decoder also includes weighted mask attention, located in the first layer of the decoder, which is used to initialize the global information state and reduce the interference of irrelevant features on the prototype.

[0021] Furthermore, the calculation formula for the weighted mask attention is shown in (1):

[0022] (1);

[0023] X1 represents the output result. , These are image features Transform encoding in Transformer, and and yes Spatial resolution. M0 represents... Let J represent the true class prediction probability at all positions in the predicted probability distribution, and let J be a tensor with a value of 1 and the same shape as M0. Indicates the first The query vector generated by the layer, where X0 represents the input query vector of the Transformer decoder.

[0024] Furthermore,

[0025] The formula for calculating the prototype contrast loss function is shown in (2):

[0026] (2);

[0027] where P + represents the set of indices of positive samples related to the current pixel i, P - represents the set of indices of negative samples related to the current pixel i, i + represents the embedding representation of positive samples related to the current pixel i, i - represents the embedding representation of negative samples related to the current pixel i, and τ represents the temperature coefficient.

[0028] Further, the calculation formula of the joint loss function is shown in (3):

[0029] (3);

[0030] where L SEG represents the joint loss function, i represents the pixel index, P represents the set of all pixel indices that need to calculate the loss, represents the cross-entropy loss of the i-th pixel, represents the prototype contrast loss of the i-th pixel, and λ>0 is a coefficient.

[0031] Based on the above technical solutions, the embodiments of the present application can at least produce the following technical effects:

[0032] (1) The supervised prototype contrast learning method based on the prototype generator provided by the present application proposes a new prototype generator structure, which can effectively capture the global attributes of the embedding space and accurately reflect the internal structure of the training data by using the label classification information. Compared with the traditional prototype generation method, this method can better handle the problems of uneven intra-class pixel distribution and neglect of spatial information.

[0033] (2) The supervised prototype contrast learning method based on the prototype generator provided by the present application significantly improves the performance of the image segmentation model by combining the pixel-wise cross-entropy loss and the prototype contrast loss. The pixel-wise cross-entropy loss ensures the classification accuracy of each pixel, while the prototype contrast loss helps the model learn more representative feature representations through contrast learning. This combination enables the model to focus on both local and global information, enhancing its ability to distinguish different classes. Compared with traditional methods, the joint loss function effectively handles the class imbalance problem, and through the introduction of class prototypes, it improves the learning effect of minority class samples. In addition, the joint loss provides more comprehensive feedback during the training process, promoting the optimization of feature representations. Experimental results show that the model using this loss function performs excellently in multiple evaluation indicators, especially in complex scenarios, exhibiting stronger robustness and generalization ability, thereby achieving better segmentation results.

[0034] (3) The supervised prototype contrast learning method based on the prototype generator has higher accuracy and expressiveness compared to the global average and weighted average method. The method can adaptively represent the weight between the pixels, and does not need to manually set the weight or artificially define the prototype, and can automatically learn the class prototype from the pixel representation. Experimental results show that on the medical thyroid ultrasound data set, the contrast method based on the prototype generator automatically generates class prototypes, which shows better performance and expressiveness, further proving the superiority of the method. BRIEF DESCRIPTION OF DRAWINGS

[0035] In order to more clearly illustrate the technical solutions in the embodiments of the present application or the prior art, the following will briefly introduce the drawings needed to be used in the embodiments or prior art description. Obviously, the drawings in the following description are only some embodiments of the present application, and for those skilled in the art, other drawings can be obtained from the structures shown in the drawings without creative labor.

[0036] Figure 1 is the structure diagram of the supervised prototype contrast method based on the prototype generator of the present application;

[0037] Figure 2 is the down-sampling process of the present application;

[0038] Figure 3 is the structure diagram of the weighted mask attention Transformer of the present application. DETAILED DESCRIPTION

[0039] The technical solutions in the embodiments of the present application will be described clearly and completely below. Obviously, the described embodiments are only some of the embodiments of the present application, not all. Based on the embodiments in the present application, all other embodiments obtained by those skilled in the art without creative labor are within the scope of protection of the present application. In addition, the technical solutions of each embodiment can be combined with each other, but it must be based on the fact that those skilled in the art can realize it. When the combination of technical solutions appears contradictory or unachievable, it should be considered that the combination of technical solutions does not exist, and is not within the scope of protection claimed by the present application.

[0040] As shown in Figure 1 The supervised prototype contrast learning method based on the prototype generator includes the following steps:

[0041] S1, a backbone network is used to extract features from an input image to generate a pixel representation R e R H×W×C , wherein R represents image features, H and W are the spatial resolution of R, and C represents the pixel dimension;

[0042] S2, after feature extraction, the network generates two calculation branches for calculating the pixel-wise cross-entropy loss and calculating the prototype contrast loss, respectively, the cross-entropy loss measures the difference between the model prediction and the true label, and the prediction value generated by the pixel feature and the ground truth are calculated, and the prototype contrast loss helps the model to learn the global context and the semantic relationship between pixels, and on this basis, to learn a structured pixel semantic embedding space;

[0043] S3, using a downsampling process, a weighted mask attention Transformer and a multi-layer perception to build a prototype generator to generate class prototypes from the extracted features;

[0044] S4, store the generated class prototypes in the prototype memory, and calculate the prototype contrast loss by calculating the similarity between the pixel embedding and the stored prototype;

[0045] S5, combine the pixel-wise cross-entropy loss and the prototype contrast loss to ensure that the model can effectively learn the pixel and class information, form a joint loss function, and use it to train the image segmentation model, so as to realize efficient and accurate medical image segmentation.

[0046] Step S3 includes the following steps:

[0047] S31, use downsampling to reduce the resolution of the feature map to obtain R'∈R H'×W'×C' , where R' represents the image feature after downsampling, H' and W' are the spatial resolution of R', and C' represents the pixel dimension after downsampling;

[0048] S32, the weighted mask attention Transformer calculates the output from the downsampled image feature and the learnable query feature , which extracts the global information of the class, where Q represents the global information of all classes in the image feature, k represents the number of classes, and C Q represents the encoding length of the global information of each class;

[0049] S33, input the output result of the weighted mask attention into the multi-layer perception to generate class prototypes , where N represents the number of types, and C represents the dimension of the prototype vector, which is the same as the dimension represented by the input pixel.

[0050] Specifically, the downsampling process is as shown in Figure 2 , which includes two 3x3 convolutions and a 2x2 pooling operation applied repeatedly 4 times. After downsampling, the output feature map is sent to the weighted mask attention Transformer for processing, and the structure of the weighted mask attention Transformer is as shown in Figure 3As shown, the model consists of two parts: an encoder and a decoder. The encoder converts the input sequence into a series of feature representations, while the decoder generates the output sequence based on these representations. The main structures of the encoder and decoder include a multi-head self-attention mechanism and a feedforward network. The decoder also includes weighted mask attention, a mechanism that not only helps the prototype generator better initialize the global information state but also more accurately captures the importance of each position to the output, providing a better starting point for subsequent decoding and improving the quality of global information. After processing by the weighted mask attention Transformer and the multilayer perceptron (MLP), the generated class prototypes are stored in a prototype memory. The prototype memory is a memory structure used to store prototype vectors for each class during training. Based on pixel features, label information, and all prototypes in the memory, the model can calculate a prototype contrastive loss to help it learn a well-structured pixel semantic embedding space.

[0051] Specifically, weighted masked attention is a variant of cross-attention that appears only in the first layer of the Transformer decoder. It is used to initialize the state of global information while reducing the interference of irrelevant features on the prototype. The calculation formula for weighted masked attention is shown in (1):

[0052] (1);

[0053] Where X1 represents the output result. , These are image features In the Transformer, the transformation encoding, T, represents the transpose matrix, and and yes Spatial resolution. M0 represents... Let J represent the true class prediction probability at all positions in the predicted probability distribution, and let J be a tensor with a value of 1 and the same shape as M0. Indicates the first The query vector generated by the layer, X0 represents the input query vector of the Transformer decoder. To accommodate computational needs, the shape of M0 is changed from... It was expanded to This means the probability value is copied H′W′ times. Weighted masked attention is used to initialize the state of the global information. By considering the contribution of each position to the output, it provides a better starting point for the subsequent decoding process, improving the decoding quality. In the Transformer model, the initial state only provides one direction, while subsequent decoder layers gradually adjust the output based on the context of the input sequence to obtain a more accurate result. To avoid the influence of the initial state on the final result, weighted masked attention only appears in the first layer of the decoder. Then, the output of the weighted masked attention is... ∈ Input to a Multilayer Perceptron (MLP) to generate class prototypes Here, N represents the number of types, and C represents the dimension of the prototype vector, which is the same as the dimension of the input pixel representation. The class prototype represents the center point of the category and is a representative feature of the category. In contrastive learning, this embodiment uses class prototypes as contrast samples to provide information about category distribution and sample features. Specifically, this embodiment uses class prototypes of the same category as the pixel embedding as positive samples and class prototypes of different categories as negative samples. Finally, the generated class prototypes are... The data is placed in the prototype memory. With the help of this memory, a large number of comparative samples enable the model to better utilize the data to learn feature representations and improve robustness and generalization performance. The size of the memory is predefined. , where N M The number of samples that can be stored is represented by , and C represents the dimension of the prototype. A first-in, first-out (FIFO) update strategy is used in memory, meaning that a new class prototype replaces the oldest class prototype in memory. This ensures that the class prototypes in memory remain up-to-date and more relevant to the current task.

[0054] In step S5, the pixel-wise cross-entropy loss ensures that the pixel embedding can correctly predict the category by calculating the consistency between the pixel embedding and the label. The prototype contrastive loss further shapes the pixel embedding space by exploring the structural information of the labeled pixel samples and training a good feature representation. The combination of these two losses enables the pixel embedding to be both discriminative and predictive, thereby improving the class discrimination effect. The prototype contrastive loss uses all prototypes in the memory as positive and negative samples in the contrastive learning. At the same time, the features extracted by the backbone network are used as anchor samples. Due to the conditions of supervised contrastive learning, the label information can be used to confirm the true classification of the features. Finally, the loss between the anchor sample and the contrast sample is calculated by applying the pixel-prototype contrastive learning formula, as shown in formula (2):

[0055] (2);

[0056] Among them, P +P represents the set of indices of positive samples associated with the current pixel i. - Let i represent the set of indices of all negative samples associated with the current pixel i. + This represents the embedding representation of the positive samples associated with the current pixel i. - τ represents the embedding representation of the negative sample associated with the current pixel i, and τ represents the temperature coefficient.

[0057] Formula (2) forces each pixel embedding to be similar to its specified positive prototype, but different from other irrelevant negative prototypes. Compared with existing segmentation models based on inter-pixel contrast learning, pixel prototype contrast learning only requires a small number of prototypes for pixel prototype contrast calculation, which does not lead to large memory costs or require a large number of pixel pairs for comparison. The method uses both pixel-level cross-entropy loss and prototype contrast loss. Pixel-level cross-entropy loss calculates the loss by classifying each pixel and learning pixel features that are meaningful for classification. Contrast loss is used to learn an embedding space with good feature representation, which can ensure the consistency and generalization of feature representation and explore the global semantic relationship between pixel samples. The total loss function is shown in Formula (3):

[0058] (3);

[0059] Among them, L SEG Let represent the joint loss function, i represent the pixel index, and P represent the set of all pixel indices for which the loss needs to be calculated. This represents the cross-entropy loss of the i-th pixel. This indicates that this is the prototype contrast loss for the i-th pixel, where λ>0 is the coefficient.

[0060] The foregoing has shown and described the basic principles, main features, and advantages of the present invention. Those skilled in the art should understand that the present invention is not limited to the above embodiments. The embodiments and descriptions in the specification are merely illustrative of the principles of the invention. Various changes and modifications can be made to the invention without departing from its spirit and scope, and all such changes and modifications fall within the scope of the present invention as claimed. The scope of protection of the present invention is defined by the appended claims and their equivalents.

Claims

1. A supervised prototype comparison learning method based on a prototype generator, characterized in that, Includes the following steps: S1. Use a backbone network to extract features from the input image and generate pixel representations R∈R H×W×C Where R represents the image features, H and W are the spatial resolutions of R, and C represents the pixel dimension; S2. After feature extraction, the backbone network generates two computational branches, which are used to calculate the pixel-wise cross-entropy loss and the prototype contrast loss, respectively. S3. A prototype generator is constructed using the downsampling process, weighted mask attention Transformer, and multilayer perceptron to generate class prototypes from the extracted features; The weighted mask attention Transformer includes an encoder and a decoder. The encoder is used to convert the input sequence into a feature representation, and the decoder is used to generate an output sequence based on the feature representation generated by the encoder. The main structure of the encoder and decoder includes a multi-head self-attention mechanism and a feedforward neural network. The decoder also includes weighted mask attention, which is located in the first layer of the decoder and is used to initialize the global information state and reduce the interference of irrelevant features on the prototype. The formula for calculating the weighted mask attention is shown in (1): (1); Where X1 represents the final output of the weighted mask attention. , These are image features Transform encoding in Transformer, and and yes Spatial resolution, C Q M0 represents the encoding length of global information for each category. Let J represent the true class prediction probability at all positions in the predicted probability distribution, and let J be a tensor with a value of 1 and the same shape as M0. Indicates the first The query vector generated by the layer, where X0 represents the input query vector of the Transformer decoder; S4. Store the generated class prototype in the prototype memory, and calculate the prototype contrast loss by calculating the similarity between the pixel embedding and the stored prototype. S5. Combine pixel-wise cross-entropy loss and prototype contrast loss to construct a joint loss function for training the image segmentation model.

2. The supervised prototype comparison learning method based on a prototype generator according to claim 1, characterized in that, Step S3 includes the following steps: S31. Use downsampling to reduce the resolution of the feature map to obtain R'∈R H'×W'×C' Where R' represents the downsampled image features, H' and W' are the spatial resolutions of R', and C' represents the pixel dimension after downsampling; S32. The weighted masked attention Transformer calculates the initial output from the downsampled image features and the learnable query features. Extract global information of categories, where Q represents the global information of all categories in the image features, k represents the number of categories, and C... Q The encoding length of global information for each category; S33, the final output of the weighted mask attention. Input into a multilayer perceptron to generate class prototypes Where N represents the number of types and C represents the dimension of the prototype vector, which is the same as the pixel dimension of the image feature R extracted by S1.

3. The supervised prototype comparison learning method based on a prototype generator according to claim 2, characterized in that, In step S31, the downsampling process includes applying two 3×3 convolutions and a 2×2 pooling operation four times.

4. The supervised prototype comparison learning method based on a prototype generator according to claim 1, characterized in that, The formula for calculating the prototype contrast loss is shown in (2): (2); Among them, P + P represents the set of indices of positive samples associated with the current pixel i. - Let i represent the set of indices of all negative samples associated with the current pixel i. + The embedding representation of the positive samples related to the current pixel i, i - τ represents the embedding representation of the negative sample associated with the current pixel i, and τ represents the temperature coefficient.

5. The supervised prototype comparison learning method based on a prototype generator according to claim 4, characterized in that, The formula for calculating the joint loss function is shown in (3): (3); Among them, L SEG Let represent the joint loss function, i represent the pixel index, and P represent the set of all pixel indices for which the loss needs to be calculated. This represents the cross-entropy loss of the i-th pixel. This indicates that this is the prototype contrast loss of the i-th pixel, where λ>0 is the coefficient.

Citation Information

Patent Citations

  • Cross-modal video text retrieval method, system and equipment and medium

    CN116910307A

  • Multi-source target detection method based on prototype network dynamic balance optimization strategy

    CN118015342A