Unsupervised small sample medical image segmentation method and system based on prototype correction and application
By generating query prototype correction support prototypes and combining self-attention mechanism and prototype distance measurement, the medical image segmentation accuracy and multi-organ segmentation confusion problems under scarcity of labeled data are solved, and efficient unsupervised small sample medical image segmentation is achieved.
Patent Information
- Application Number
- CN202410177632.9
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2024-02-08
- Publication Date
- 2025-08-08
AI Technical Summary
In the case of scarce labeling data in the prior art, the segmentation accuracy of medical image segmentation models is low, and semantic confusion is prone to occur during multi-organ segmentation.
By generating query prototypes, correcting the support prototypes, using query features to enrich the semantic information of the support prototypes, and using the self-attention mechanism to generate segmentation abnormal thresholds, and combining the prototype distance metric for multi-class fusion segmentation.
It improves the generalization ability and robustness of the model, reduces the work burden of labeling, and can achieve accurate multi-category prediction segmentation in multi-organ segmentation scenarios.
Smart Images

Figure BDA0004702582170000041 
Figure BDA0004702582170000051 
Figure BDA0004702582170000052
Abstract
Description
Technical Field
[0001] The present invention belongs to the technical field of image segmentation and relates to a small-sample medical image segmentation method, specifically an unsupervised small-sample medical image segmentation method, system and application based on prototype correction, which can be applied to medical image segmentation in scenarios where labeled data is scarce. Background Art
[0002] Medical image segmentation is the task of dense pixel-level classification of medical images, such as computed tomography (CT) and magnetic resonance imaging (MRI). Medical image segmentation plays a vital role in providing important anatomical and lesion information for disease diagnosis, treatment planning, and surgical navigation during clinical procedures. Over the years, fully supervised deep learning methods based on large amounts of annotated data have achieved promising results. However, due to concerns about patient privacy, obtaining raw medical image data is often very difficult. Furthermore, the annotation of medical images requires professional medical knowledge and extensive clinical experience, so the entire annotation process requires a lot of manpower and material resources from doctors. These two reasons lead to the scarcity of annotated medical image data, which hinders the generalization and transfer capabilities of traditional data-driven models and reduces the segmentation accuracy of the models.
[0003] To address the limited availability of labeled data, researchers have proposed small-shot segmentation techniques in recent years, which aim to rapidly extract new patterns from a very limited number of samples belonging to a new category. Small-shot segmentation employs a learning paradigm called "meta-learning," the core idea of which is to enable the model to quickly adapt to new tasks or domains while maintaining efficient segmentation by simulating a set of similar training tasks in scenarios where labeled data is scarce. However, due to the low resolution and intricate organizational structure of medical images, it is often difficult for researchers to obtain training data for similar tasks. Recent work relies on prototype learning methods, which divide the data into a support set and a query set. After training a set of representation prototypes on the support set, images in the query set are segmented by measuring the pixel-level distance between the images in the query set and these prototypes. The support set typically consists of only a few images, simulating the real-world scenario of scarce labeled data.
[0004] Fully convolutional networks (FCNs) solve pixel-level classification problems by replacing the last fully connected layer of a convolutional neural network with a convolutional layer. Building on this framework, U-Net (Convolutional networks for biomedical image segmentation) introduces symmetrical encoder and decoder architectures to ensure consistency and coherence of deep semantic information. This design demonstrates superior performance in low-resolution medical image segmentation. Leveraging the advantages of U-Net, Dense U-Net (Bidirectional conv1stmu-net with densley-connected convolutions) enhances feature propagation and reuse through densely connected convolutions. UNet++ (Anestedu-net architecture for formed medical image segmentation) bridges the semantic gap between encoder and decoder feature maps, simultaneously capturing both shallow, simple features and deep, abstract features. Advanced techniques such as dilated convolutions and pyramid pooling can also capture a wider range of contextual semantic information and multi-scale feature representations. Although these models show impressive performance, their segmentation accuracy is often low when labeled data is scarce, making it impractical to directly use them for medical image segmentation in data-scarce scenarios.
[0005] OSLSM (One-shot learning for semantic segmentation) laid the foundation for early exploration of few-shot segmentation through a shared-weight two-branch network. Recent work has mostly adopted prototype-based approaches, mapping support features to prototype vectors in semantic space and performing pixel-wise matching between query features and the prototype vectors to achieve segmentation. PANet (Panet: Few-shot image semantic segmentation with prototype alignment) proposes prototype alignment regularization between the support and query sets to fully leverage knowledge from the support set. SSL-ALPNet (Self-supervision with superpixels: Training few-shot medical image segmentation without annotation) proposes an adaptive local prototype pooling module that constructs multiple local prototypes to mitigate the loss of intra-class local information caused by mask average pooling. Compared to SSL-ALPNet, ADNet (Anomaly detection-inspired few-shot medical image segmentation through self-supervision with supervoxels) uses only a single prototype to model homogeneous foreground classes and preserves local information by measuring pixel outliers.
[0006] However, due to the limited number of images in the support set, it is unrealistic to obtain prototypes from the support set that cover diverse semantic information. Given that the images in the query set and the support set differ significantly in color, composition, and spatial layout, relying solely on information-limited support prototypes for segmentation guidance may lead to semantic confusion, that is, the foreground class or the target organ class to be segmented and the background class are incorrectly classified. Although researchers have proposed many techniques and methods to address this shortcoming, such as introducing attention mechanisms and various prototype generation methods to generate more representative support prototypes, they have not fundamentally solved the problem. Summary of the Invention
[0007] In order to address the deficiencies in the prior art, the present invention aims to provide an unsupervised small-sample medical image segmentation method, system and application based on prototype correction.
[0008] Unlike prior art research, the present invention addresses the fundamental problem of semantic confusion, which is caused by the deviation between the supporting prototype and the query features. Therefore, the present invention considers using query features to generate a query prototype to correct the supporting prototype, generating a corrected prototype that is more suitable for the image features in the query set, thereby reducing the deviation of the supporting prototype from the query set image. The corrected prototype can more accurately match the images in the query set and generate more accurate segmentation results. Meanwhile, prior art research performs segmentation predictions on a slice-by-slice and category-by-category basis. The present invention provides a multi-class fusion segmentation method based on prototype distance metrics in multi-organ segmentation scenarios.
[0009] In this paper, the prototype refers to the category prototype, that is, the representative feature points of each organ category, represented as a vector in the feature space, representing the distribution of the category in the feature space. The support prototype refers to the category prototype learned from the support image and its mask, and the query prototype refers to the category prototype learned from the query image and its pseudo mask.
[0010] The present invention provides an unsupervised small-sample medical image segmentation method based on prototype correction. The present invention can meet the original design requirements. Compared with other small-sample medical image segmentation methods, the present invention has stronger generalization ability and robustness. The present invention does not require additional labeled data during training, which can greatly reduce the time-consuming and labor-intensive pixel-by-pixel labeling work of medical staff. At the same time, for multi-organ segmentation scenarios, the present invention proposes a simple but effective method based on prototype distance metric, which can complete multi-class fusion prediction segmentation with very little time cost.
[0011] The present invention can realize the prediction and segmentation of single-class or multi-class samples through prototype-based distance measurement; when predicting a single class, only one correction prototype is generated, and the distance measurement is performed between the pixel point and this correction prototype to determine whether it belongs to the foreground class (i.e. the currently predicted organ class) or the background class; when predicting multiple classes, multiple correction prototypes are generated, each organ class corresponds to a correction prototype, and the pixel point is compared with the correction prototypes of these organ classes to determine which organ class it belongs to.
[0012] The present invention is achieved in that:
[0013] This invention proposes a prototype correction network for small-sample medical image segmentation and a multi-class segmentation method based on the prototype correction network. The first aspect of the invention is the prototype correction network. For this prototype correction network, the invention proposes two key technologies: a pseudo-mask generation method and a prototype correction method. The core concept of the invention is to use query prototypes generated from query features to improve support prototypes generated from support features. First, a query pseudo-mask is generated, which serves as the basis for generating the query prototype. Similar to the method for generating the support prototype, the query pseudo-mask and query features are combined through masked average pooling to generate the query prototype. Masked average pooling in the invention involves using bilinear interpolation to resize the query feature map to the same size as the query pseudo-mask, then performing a Hadamard product between the query feature map and the query pseudo-mask and averaging the results. The query prototype and support prototypes are then combined through a prototype correction method, which uses a learnable hyperparameter λ to weight the query and support prototypes to generate a final corrected prototype. This produces a higher-quality and more representative corrected prototype. Finally, the corrected prototype is used to perform feature matching with images in the query set to generate a segmentation prediction. To address the problem of local information loss caused by the masked average pooling operation used in prototype generation, the present invention also proposes a segmentation anomaly threshold generated based on a self-attention mechanism to perform the segmentation process. In the multi-organ segmentation scenario, a simple but effective multi-organ segmentation method based on prototype distance metric is proposed, namely, the cosine distance calculated by the query feature vector and the corrected prototype. Specifically, the cosine similarity value of each pixel point is calculated with multiple class prototypes, and the class assignment of the pixel point is performed by selecting the highest cosine similarity value.
[0014] Specifically, the unsupervised small sample medical image segmentation method based on prototype correction of the present invention includes the following steps:
[0015] Step 1: extract slices from the 3D data image as support images and query images;
[0016] Step 2: Input the support image and query image obtained in step 1 into the shared weight encoder network, extract the deep semantic feature map, and obtain the support feature map and query feature map;
[0017] Step 3: Perform Hadamard product on the support feature map and the foreground mask of the support image;
[0018] Step 4: Calculate the query feature vector v pixel by pixel q and support eigenvector v s The cosine similarity between
[0019] Step 5: Select the maximum similarity value between the query feature vector of a certain pixel position in the query feature map and the support feature vector of each pixel point in the support feature map as the cosine similarity value of the pixel position in the query feature map; calculate the cosine similarity value of each pixel position in the query feature map, and together form a query pseudo mask; perform a dimensionality transformation operation and a maximum-minimum regularization on the obtained query pseudo mask to obtain a final query pseudo mask;
[0020] Step 6: Perform a mask average pooling operation based on the query feature map and the query pseudo mask to generate a query prototype; perform a mask average pooling operation based on the support feature map and the support mask to generate a support prototype;
[0021] Step 7: Add the weights of the query prototype and the support prototype generated in step 6 to generate a correction prototype;
[0022] Step 8. Calculate the cosine distance between the query feature map and the corrected prototype to obtain the similarity feature map, and use the function to generate the final segmentation prediction foreground class and segmentation prediction background.
[0023] In step 1, the slice is a two-dimensional image randomly extracted from the three-dimensional data image, and the support image and the query image are different images;
[0024] In step 2, the support image and the query image are fed into the same shared weight encoder network f θ , the deep semantic feature maps of the support image and the query image are represented as F s and F q ;
[0025] F s =f θ (x s )∈R C×H×W ,
[0026] F q =f θ (x q )∈R C×H×W ,
[0027] Among them, x s Indicates support for images, x q Represents the query image. After the feature extraction by the encoder, F s and F q They are all of C×H×W dimensions, where C is the channel size, H and W are the height and width of the feature space respectively;
[0028] In step 3, the foreground mask of the support image is the annotation of the support image; the calculation of the Hadamard product can set the background in the support feature map to 0; the calculation process of the Hadamard product is as follows:
[0029]
[0030] in, represents the foreground mask of the support image, F s′ Indicates the support feature map after processing. At this time, the background pixel value in the support feature map is already 0.
[0031] Between step 3 and step 4, there is also the query feature vector v q and support eigenvector v s The query feature vector v q From the query feature graph F q , support eigenvector v s From the processed support feature map F s′ Specifically, the feature map is of H×W×C dimensions, with a total of H×W pixels. Each pixel can extract a 1×C dimension feature vector. The query feature vector is obtained from the query feature map F. q Extract the support feature vector from the processed support feature map F s′ Extraction
[0032] In step 4, the cosine similarity is calculated as follows:
[0033]
[0034] The cosine similarity represents the similarity between two pixels. The cosine similarity value is calculated by comparing the query feature vector of a pixel position in the query feature map with the support feature vector of each pixel in the processed support feature map.
[0035] In step 5, the cosine similarity value of the pixel position is expressed as:
[0036]
[0037] The cosine similarity value of each pixel position in the query feature map is used to form the query pseudo mask as follows:
[0038]
[0039] Perform dimension transformation and max-min regularization on the query pseudo mask to convert the query pseudo mask into H×W×1 dimensions, and limit the similarity value to between 0 and 1 to obtain the final query pseudo mask:
[0040]
[0041] Wherein, ∈ is set to a minimum value to prevent the algorithm from dividing by 0; in a specific embodiment, it can be set to 1e-7;
[0042] In step 6, the process of generating the query prototype or the support prototype by the mask average pooling operation includes the following: using bilinear interpolation to adjust the size of the feature map to the same as the mask, performing Hadamard product on the feature map and the mask and then taking the average. The generation of the query prototype and the support prototype is respectively expressed as follows:
[0043]
[0044]
[0045] In step seven, the corrected prototype is represented by the following formula:
[0046] p=(1-λ)×p s +λ×p q ,
[0047] Among them, λ is a learnable hyperparameter used to balance the degree of correction of the query prototype to the support prototype.
[0048] In step eight, the cosine distance between the query feature map and the corrected prototype is expressed as follows:
[0049]
[0050] Wherein, α represents a multiplier, which helps the gradient to be back-propagated during training. In the present invention, the α is fixed to 20;
[0051] The final generated segmentation prediction foreground class is represented as follows:
[0052]
[0053] Where T represents the threshold of abnormal pixels, which is obtained by setting the initial value and adding the loss function; the σ function is used to convert the similarity feature map into the probability that each pixel belongs to the foreground class;
[0054] The threshold of the abnormal pixel point is dynamically generated through the self-attention mechanism. The output feature map of the fourth layer of the encoder network is extracted, flattened into a one-dimensional vector after the average pooling operation, and applied to the support feature map and the query feature map to obtain Q, K, and V respectively. The weighted V is calculated using the attention algorithm, and the above calculation result is mapped to a scalar through a full connection operation. The scalar is the initial value of T:
[0055] T = FC(S(Q,K,V));
[0056] The segmentation prediction background is expressed as:
[0057]
[0058] During image segmentation, the loss function of model training includes query loss L Q , its loss L A and outlier loss L T The query loss is obtained by calculating the binary cross entropy segmentation loss, which is expressed as:
[0059]
[0060] Where H' and W' represent the height and width of the image when input to the encoder, respectively. and Represent the true foreground mask and background mask of the query image respectively. The same strategy as the general prototype-based method is adopted, and the prototype loss L is added. A , that is, using the prediction results of the query image, the query prototype is calculated in the same way to guide the segmentation of the support image. The formula can be expressed as:
[0061]
[0062] L T The threshold used to learn segmentation is expressed as:
[0063] L T =T / α.
[0064] Therefore, the total loss function is L = L Q +L A +L T .
[0065] The present invention also provides a system for implementing the above-mentioned image segmentation method, the system comprising: an encoder module, a pseudo mask generation module, a prototype correction module and a prediction module;
[0066] The encoder module contains a shared weight encoder network f θ , used to extract deep semantic feature maps of images;
[0067] The pseudo mask generation module is used to generate a pseudo foreground mask of the query image;
[0068] The prototype correction module is used to balance the correction degree of the query prototype extracted from the query feature to the support prototype obtained from the support feature, and to perform weighted summation of the support prototype and the query prototype to generate a correction prototype;
[0069] The prediction module is used to perform image segmentation on the image to be segmented based on the image segmentation method.
[0070] The present invention also provides the application of the above-mentioned image segmentation method or the above-mentioned image segmentation system in unsupervised small sample image segmentation.
[0071] The beneficial effects of the present invention include:
[0072] 1) This paper proposes a novel unsupervised small-sample medical image segmentation method based on prototype correction, which alleviates the prototype bias problem by using query prototypes to enrich the semantic information of supporting prototypes.
[0073] 2) The present invention uses the self-attention mechanism to generate a segmentation anomaly threshold that is closely related to the query image instead of a fixed anomaly initial value, making the network more robust to changes in the foreign trade differences of the query image.
[0074] 3) We design a simple yet effective prototype-based multi-class prediction method to alleviate the class confusion problem during multi-organ prediction.
[0075] 4) The present invention proposes to add boundary loss to the total loss function to help the neural network model better learn the boundary details and accuracy of small organs, so as to alleviate the problem of huge differences in the areas occupied by organ classes and background classes in medical image segmentation. BRIEF DESCRIPTION OF THE DRAWINGS
[0076] In order to more clearly illustrate the embodiments of the present invention or the technical solutions in the prior art, the following briefly introduces the drawings required for use in the embodiments or the description of the prior art. Obviously, the drawings described below are only some embodiments of the present invention. For those skilled in the art, other drawings can be obtained based on these drawings without paying any creative work.
[0077] Figure 1 This is a diagram of the overall network architecture.
[0078] Figure 2 Schematic diagram of the pseudo-mask generation module.
[0079] Figure 3 Segmentation visualization result diagram (the organs in the image are from left to right: yellow represents the liver, light blue represents the right kidney, pink represents the left kidney, orange represents the spleen, and the dark blue area represents the area with category confusion). DETAILED DESCRIPTION
[0080] The present invention is further described in detail with reference to the following specific examples and accompanying drawings. The processes, conditions, experimental methods, etc. for implementing the present invention, except for those specifically mentioned below, are common knowledge and common common sense in the art and are not particularly limited by the present invention.
[0081] This paper proposes an unsupervised small-sample medical image segmentation method based on prototype correction. This method addresses the prototype bias problem caused by variations in appearance between support and query images, which can occur when matching support prototypes with query features. By cleverly designing a query prototype from the query image, the support prototype is corrected, making it more suitable for query image segmentation. The method introduces a pseudo-mask generation module that generates a query pseudo-mask through a many-to-many approach, which serves as the basis for generating the query prototype. The method then designs a prototype correction module with a learnable parameter λ to balance the degree of correction between the query prototype extracted from the query features and the support prototype obtained from the support features. Furthermore, the method designs a prototype-based multi-class segmentation method to alleviate the problem of confusing region prediction between query images of different organs in multi-organ segmentation scenarios. This method effectively mitigates the prototype bias problem caused by intra-class appearance variations in small-sample medical image segmentation, enhances the adaptability of the corrected prototype to changes in query features, and improves the robustness and generalization of the model. Extensive experimental validation on two widely used medical image datasets demonstrates that the proposed method achieves leading-edge performance in multiple organ segmentation scenarios, demonstrating its effectiveness.
[0082] In order to verify the effectiveness of the present invention, the present invention selects two commonly used medical image datasets CHAOST2 and CMR for testing the performance of medical image segmentation algorithms for testing. Of course, the present invention can also be extended to other medical image datasets, such as blood vessel segmentation datasets, tooth segmentation datasets, bone segmentation datasets, etc. Among them, CHAOST2 is an image for abdominal organ segmentation, including 20 3D CT scans, and each 3D CT image contains an average of 36 2D slice images. The pre-labeled organs include 12 organ classes, including liver (Liver), spleen (Spleen), left kidney (LK) and right kidney (RK). The first four organ classes are used for training and testing when verifying the present invention; CMR is an MRI image for heart segmentation, including 35 3D cardiac MRI images, and each 3D image contains an average of 13 2D slices. The pre-labeled organs include 34 organ classes, including left ventricular cistern (LVBP), left ventricular myocardium (LV-MYO) and right ventricular myocardium (RV). In this method, the present invention uses the first three organ classes. Original medical images are in DICOM format, containing information about all tissues. Directly inputting them into a neural network for learning may yield poor results. To facilitate neural network learning, the present invention converts the original images into NII format and preprocesses the data. First, the data is normalized. To mitigate the impact of non-resonance effects, the present invention truncates the values of voxels exceeding the 95th percentile intensity value. The voxel values are then mean-normalized to maintain a range between 0 and 255. Following normalization, the present invention resamples each slice of the three-dimensional image data to ensure uniform voxel spacing between slices. Finally, each slice is cropped to a 256*256 resolution with reference to the center. The annotation, or mask, corresponding to each DICOM-formatted medical image is a PNG format image. Referring to the above method, the mask undergoes NII format conversion, slice resampling, and size cropping. Since the mask voxel values range from 0 to the number of annotated organ categories, the mask does not require normalization.
[0083] To address the limitation of neural network learning, which requires a large amount of labeled data, this paper uses supervoxel clustering to generate the labels needed for neural network training. Supervoxel clustering involves clustering voxels with similar attributes into small, compact regions. The basic principle is to fuse voxel data that meet similarity constraints within a local area, generating a pseudo-class label for each supervoxel.
[0084] The network proposed in this invention can be divided into four parts: encoder module, pseudo mask generation module, prototype correction module and final prediction module.
[0085] The present invention randomly extracts a slice as a support image and a slice as a query image from the same volume, i.e., 3D image data, and inputs them into the shared weight encoder network f θ , to extract the deep semantic feature map of the support image and the query image, which is represented as F s and F q The process can be expressed as follows:
[0086] F s =f θ (x s )∈R C×H×W ,
[0087] F q =f θ (x q )∈R C×H×W ,
[0088] where x s Indicates support for images, x q Represents the query image. After the feature extraction by the encoder, F s and F q They are all of dimension C×H×W, where C is the channel size, H and W are the height and width of the feature space respectively.
[0089] The pseudo foreground mask of the query image is generated by the pseudo mask generation module used in the present invention. The specific implementation process is as follows: first, the support feature map and the foreground mask of the support image, i.e., the annotation of the support image, are Hadamard products (mathematically represented by ⊙) to set the background in the support feature map to 0. The formula for this process can be expressed as:
[0090]
[0091] in represents the foreground mask of the support image, F s′ Represents the processed support feature map, where the background pixel value in the support feature map is already 0. Then the query feature vector v extracted from the feature map is calculated pixel by pixel q and support eigenvector v s The cosine similarity between q and v s A feature vector of 1×C dimensions, where the query feature vector comes from the query feature graph F q , the support feature vector comes from the processed support feature map F s′ , the formula indicates that this process can be:
[0092]
[0093] The cosine similarity value cos(v q ,v s ) is a real number that represents the similarity between two vectors. In the segmentation scenario, it can be understood as the similarity between two pixels. According to this method, the query feature vector of a pixel position in the query feature map is calculated with the support feature vector of each pixel in the processed support feature map. The largest similarity value is selected as the cosine similarity value of the current pixel position in the query feature map. The formula is expressed as:
[0094]
[0095] By calculating the cosine similarity value of a certain pixel position in the query feature map pixel by pixel as described above, a query mask, i.e., the annotation of the query image, can be obtained. Since it is not a real annotation, this invention calls it a query pseudo mask. The formula for this process can be:
[0096]
[0097] The medical image segmentation method of the present invention is an unsupervised segmentation method. The support mask of the present invention is generated by a supervoxel clustering algorithm and serves as an annotation of the image, thereby achieving the purpose of unsupervised training.
[0098] Finally, the dimension transformation operation will be Converted into H×W×1 dimension, the present invention expresses it as Performing maximum-minimum regularization limits the similarity value in the query pseudo mask to between 0 and 1. The formula can be expressed as:
[0099]
[0100] At this point, the final query pseudo mask is obtained Here, ∈ is set to 1e-7 to prevent division by zero in the algorithm. Since the background pixel values of the support feature map are set to 0 when calculating similarity, the value of each pixel in the generated query pseudo-mask only reflects the degree of correlation with the foreground target class region of the support feature. A higher correlation coefficient indicates that the pixel is more likely to belong to the target class region in the query image. This is a simple but extremely effective method for calculating the correlation between the query feature map and the support feature map, while preserving the spatial and appearance differences between the query feature map and the support feature map, providing a key guarantee for a more accurate prototype matching process.
[0101] The query prototype is generated in the same way as the support prototype, which is to perform mask average pooling on the feature map and mask. In order to avoid unnecessary details, the present invention uses the generation of the query prototype as an example to illustrate the process of generating the prototype by mask average pooling. First, the query feature map is resized to the same size as the query pseudo mask by bilinear interpolation. The query feature map and the query pseudo mask are Hadamard-multiplied and then averaged. The formula represents this process as follows:
[0102]
[0103] Similarly, the generated support prototype is marked as p s , The prototype correction module adds the weights of the support prototype and the query prototype to generate a correction prototype. The formula of this process is expressed as:
[0104] p=(1-λ)×p s +λ×p q ,
[0105] Here, λ is a learnable hyperparameter used to balance the degree of correction applied to the query prototype against the support prototype. The correction prototype is a vector representation of the foreground class. By using the query prototype to correct the support prototype, a more adaptable and flexible, high-quality correction prototype is produced. A key reason for this is that the correction prototype, with the support of the query features, can achieve more accurate and effective matching when the query features differ from the support features in appearance. In particular, when the support features contain limited semantic information, the correction prototype demonstrates strong generalization capabilities.
[0106] By calculating the cosine distance between the query feature map and the corrected prototype, a similarity feature map can be obtained, in which the value of each pixel represents the degree of similarity with the target class prototype. The target class prototype is the corrected prototype, that is, the corrected prototype represents an organ class prototype, that is, the currently segmented organ class (target class); this process is expressed by the formula:
[0107]
[0108] Where α represents a multiplier that helps the gradient to be back-propagated during training. In the experiment, α is set to 20. Finally, the Sigmoid activation function symbolically represented as σ(·) is used to generate the final segmentation prediction foreground class:
[0109] The predicted foreground after segmentation is expressed as:
[0110] Where T represents the threshold of abnormal pixels. Previous work set the initial value of T to 10, and then learned an abnormal threshold applicable to the segmentation of all query images by adding the T loss function. However, since different organ slices may be very different, it is not reasonable to use a unified abnormal threshold to perform segmentation. For this reason, the present invention proposes to dynamically generate abnormal values of pixels through a self-attention mechanism. This abnormal value will be closely related to the slice currently predicted to be segmented. The specific implementation process is to extract the output feature map of the fourth layer of the encoder network, flatten it into a one-dimensional vector after the average pooling operation. This operation is applied to the support features twice to obtain Q and K, and applied to the query features to obtain V. The weighted V is calculated using the ScaledDotProductAttention algorithm (symbolized as S(·)). Finally, the above calculation result is mapped to a scalar through a fully connected operation (FC(·)). This scalar is used as the initial value of T. The formula is expressed as follows:
[0111] T=FC(S(Q,K,V)).
[0112] The initial value of the anomaly obtained in this way is more reasonable because it is closely related to the query image. It can better adapt to changes in the query image when performing segmentation, making the algorithm more robust.
[0113] The predicted background can be expressed as:
[0114]
[0115] Through prototype matching, we can clearly see that the more similar a point is to the prototype, the closer the predicted value is to 1, that is, the more likely this point is to belong to the foreground class, and vice versa.
[0116] In the image segmentation method, the loss function of the model training includes three components: query loss L Q , its loss L A and outlier loss L T The query loss is obtained by calculating the binary cross entropy segmentation loss, which is expressed as:
[0117]
[0118] Where H' and W' represent the height and width of the image when input to the encoder, respectively. and Represent the true foreground mask and background mask of the query image respectively. The same strategy as the general prototype-based method is adopted, and the prototype loss L is added. A , that is, using the prediction results of the query image, the query prototype is calculated in the same way to guide the segmentation of the support image. The formula can be expressed as:
[0119]
[0120] L T The threshold used to learn segmentation is expressed as:
[0121] L T =T / α.
[0122] Therefore, the total loss function is L = L Q +L A +L T .
[0123] When verifying the effectiveness of the present invention, the average Dice coefficient is used as an evaluation indicator. The average Dice coefficient is a commonly used indicator for evaluating medical image segmentation algorithms. It is calculated by comparing the degree of overlap between the algorithm segmentation results and the true segmentation results. Its value range is 0 to 1, where 1 indicates complete overlap and 0 indicates no overlap. The formula is:
[0124]
[0125] Where A and B represent the predicted value and the true annotation respectively. The higher the DSC, the closer the segmentation result is to the true annotation, and the better and more accurate the effect.
[0126] The present invention is implemented using PyTorch (v2.0.0). The encoder is taken from the ResNet-101 network pre-trained on the MS-COCO dataset. The optimization algorithm adopts stochastic gradient descent (SGD), and its momentum is set to 0.9. The initial learning rate is set to 0.001, and the learning rate reduction coefficient is set to 0.95. It decreases once every 1000 iterations. The entire training process is set to 50,000 iterations. NVIDIA RTX 2070Ti GPU is used for training.
[0127] The method proposed in the present invention adopts a supervoxel-based training method, that is, no labeled data is required in the training stage, which solves the limitations of data-driven neural networks and reduces the dependence of neural networks on the amount of labeled data. In the testing stage, the present invention adopts two experimental settings for sufficient verification. In setting one, weak labels are provided, that is, both the support slice and the query slice contain the target class. In setting two, the middle slice of a volume is extracted as the support image, and the remaining slices are used as the query image. Under this setting, both the support image and the query image may not contain the target class. This is an experimental setting that is more in line with the real-world medical image segmentation scenario, and of course it is more challenging.
[0128] To validate the effectiveness of our method, we conducted extensive experiments in two experimental settings and compared it with state-of-the-art prototype-based models from recent years. These models include PANet, ALPNet, PPNet, and ADNet. The results for PANet, ALPNet, and PPNet are from the original papers, while the data for ADNet comes from executing the original paper's code on a server. To eliminate random factors in the experimental results, each dataset was evenly divided into five parts for five-fold cross-validation. This was repeated three times for a total of three times, and the average standard deviation was calculated. The proposed method significantly outperformed other state-of-the-art networks on both datasets.
[0129] Table I
[0130]
[0131]
[0132] Table II
[0133]
[0134] Tables I and II show the DSC scores on two datasets under experimental setup 1, namely, weak labeling. As shown in the tables, our method achieves the highest DSC scores on the CHAOST2 and MS-CMR datasets, with average DSC scores of 80.62% and 77.52%, respectively. This achievement is particularly significant compared to the performance of ADNet, which significantly improves by 1.85% and 2.37% on the respective datasets, and has a lower average standard deviation, demonstrating that our method is more stable in segmentation and more robust to image variations.
[0135] Table III
[0136]
[0137] Table IV
[0138]
[0139] In the more challenging setting 2, where the query image does not necessarily contain the target class, the proposed method demonstrates excellent performance. As shown in Tables III and IV, the average DSC score is improved by more than 2.5% compared to ADNet on both datasets. It is worth noting that the significant enhancement was observed in the segmentation of the left kidney. Given that the left kidney is similar in size to the spleen and similar in shape to the right kidney, the left kidney is a more complex and difficult organ to segment in the CHAOST2 dataset. The proposed method significantly improved the recognition rate of this organ category by 6.04%. This highlights the effectiveness of the query prototype of the present invention, which can cleverly capture the appearance changes and spatial inconsistencies between the query image and the support image, while also successfully correcting the support prototype. Therefore, the corrected prototype can better adapt to changes in the query image, further enhancing the robustness and generalization ability of the model.
[0140] In the multi-organ segmentation scenario, this paper proposes a prototype-based multi-class segmentation, a simple but efficient method for multi-organ prediction. Specifically, the paper calculates the voxel cosine distance between the query feature and the prototypes of different categories. For each voxel, the class prediction (c q ) is achieved by selecting the prototype with the smallest cosine distance to the query feature, and the formula is expressed as:
[0141] c q =argmax(cos(v q ,p i )) i∈1,2,3,…n,
[0142] Among them, p i represents the prototype of the i-th organ category, and n represents the number of categories. Figure 3 The visualization results of prototype-driven fusion prediction on the CHAOST2 dataset are presented, which clearly shows that the fusion prediction method of the present invention effectively solves the inter-class confusion problem where the organ boundaries are fuzzy.
[0143] Furthermore, in deep learning, boundary loss is a loss function used for segmentation tasks, which aims to help neural network models better learn the details and accuracy of object boundaries, especially at target edges and subtle structures. Boundary loss is different from region-based Dice loss. It increases attention to boundary pixels by integrating on the boundaries between regions, improving the model's prediction accuracy for target boundaries, which helps to solve the related problems of regional loss in highly unbalanced segmentation problems, thereby producing more accurate segmentation results and avoiding the predicted segmented areas being too fuzzy or unclear. Since the present invention is directed to medical image segmentation, and the organ pixels in medical images are often very different from those in the background, such as the size of the left and right kidneys is very different from the background, and the area occupied by the same organ class in different slices will also be very different. Using only region-based Dice loss may encounter difficulties when encountering very small areas. In view of this, boundary loss can be added to the total loss function to alleviate the problem of huge differences in the areas occupied by organ classes and background classes in medical image segmentation.
[0144] The protection content of the present invention is not limited to the above embodiments. Without departing from the spirit and scope of the present invention, changes and advantages that can be thought of by those skilled in the art are included in the present invention and are protected by the appended claims.
Claims
1. An unsupervised small sample medical image segmentation method based on prototype correction, characterized in that: The method comprises the following steps: Step 1: extract slices from the 3D data image as support images and query images; Step 2: Input the support image and query image obtained in step 1 into the shared weight encoder network, extract the deep semantic feature map, and obtain the support feature map and query feature map; Step 3: Perform Hadamard product on the support feature map and the foreground mask of the support image; Step 4: Calculate the query feature vector v pixel by pixel q and support eigenvector v s The cosine similarity between Step 5: Select the maximum similarity value between the query feature vector of a certain pixel position in the query feature map and the support feature vector of each pixel point in the support feature map as the cosine similarity value of the pixel position in the query feature map; Calculating the cosine similarity value of each pixel position in the query feature map to form a query pseudo mask; performing a dimensionality transformation operation and a maximum-minimum regularization operation on the obtained query pseudo mask to obtain a final query pseudo mask; Step 6: Perform a mask average pooling operation based on the query feature map and the query pseudo mask to generate a query prototype; Perform mask average pooling operation based on the support feature map and support mask to generate support prototype; Step 7: Add the weights of the query prototype and the support prototype generated in step 6 to generate a correction prototype; Step 8. Calculate the cosine distance between the query feature map and the corrected prototype to obtain the similarity feature map, and use the function to generate the final segmentation prediction foreground class and segmentation prediction background.
2. The method according to claim 1, wherein In step 1, the slice is a two-dimensional image randomly extracted from the three-dimensional data image, and the supporting image and the query image are different images; and / or, In step 2, the support image and the query image are fed into the same shared weight encoder network f θ , the deep semantic feature maps of the support image and the query image are represented as F s and F q ; F s =f θ (x s )∈R C×H×W , F q =f θ (x q )∈R C×H×W , Among them, x s Indicates support for images, x q Represents the query image. After the feature extraction by the encoder, F s and F q They are all of dimension C×H×W, where C is the channel size, H and W are the height and width of the feature space respectively.
3. The method according to claim 1, wherein In step 3, the foreground mask of the support image is the annotation of the support image; the calculation of the Hadamard product can set the background in the support feature map to 0; the calculation process of the Hadamard product is as follows: in, represents the foreground mask of the support image, F s′ represents the processed support feature map, where the background pixel values in the support feature map are already 0; and / or, In step 4, the cosine similarity is calculated as follows: The cosine similarity represents the degree of similarity between two pixels, and the cosine similarity value is calculated for the query feature vector of a pixel position in the query feature map and the support feature vector of each pixel point in the processed support feature map; and / or, Between step 3 and step 4, there is also the query feature vector v q and support eigenvector v s The query feature vector v q From the query feature F q , support eigenvector v s From the processed support feature F s′ .
4. The method according to claim 1, wherein In step 5, the cosine similarity value of the pixel position is expressed as: The cosine similarity value of each pixel position in the query feature map is used to form the query pseudo mask as follows: Perform dimension transformation and max-min regularization on the query pseudo mask to convert the query pseudo mask into H×W×1 dimensions, and limit the similarity value to between 0 and 1 to obtain the final query pseudo mask: Among them, ∈ is set to a minimum value to prevent the algorithm from dividing by 0.
5. The method according to claim 1, wherein In step 6, the process of generating the query prototype or the support prototype by the mask average pooling operation includes the following: using bilinear interpolation to adjust the size of the feature map to the same as the mask, performing Hadamard product on the feature map and the mask and then taking the average. The generation of the query prototype and the support prototype is respectively expressed as follows:
6. The method according to claim 1, wherein In step seven, the corrected prototype is represented by the following formula: p=(1-λ)×p s +λ×p q , Among them, λ is a learnable hyperparameter used to balance the degree of correction of the query prototype to the support prototype.
7. The method according to claim 1, wherein In step eight, the cosine distance between the query feature map and the corrected prototype is expressed as follows: Among them, α represents a multiplier that helps the gradient to be back-propagated during training; The final generated segmentation prediction foreground class is represented as follows: Where T represents the threshold of abnormal pixels, which is obtained by setting the initial value and adding the loss function; The threshold of the abnormal pixel point is dynamically generated through the self-attention mechanism. The output feature map of the fourth layer of the encoder network is extracted, flattened into a one-dimensional vector after the average pooling operation, and applied to the support feature map and the query feature map to obtain Q, K, and V respectively. The weighted V is calculated using the attention algorithm, and the above calculation result is mapped to a scalar through a full connection operation. The scalar is the initial value of T: T = FC(S(Q,K,V)); The segmentation prediction background is expressed as:
8. The method according to claim 7, wherein During image segmentation, the loss function of model training includes query loss L Q , its loss L A and outlier loss L T ; The query loss is obtained by calculating the binary cross entropy segmentation loss, and the formula is expressed as: Where H' and W' represent the height and width of the image when input to the encoder, respectively. and Represent the true foreground mask and background mask of the query image respectively; The prototype loses L A , that is, using the prediction results of the query image, the query prototype is calculated to guide the segmentation of the support image. The formula can be expressed as: L T The threshold used to learn segmentation is expressed as: L T =T / a; The total loss function is L = L Q +L A +L T .
9. An image segmentation system for implementing the image segmentation method according to any one of claims 1 to 8, characterized in that: The image segmentation system includes an encoder module, a pseudo mask generation module, a prototype correction module and a prediction module; The encoder module contains a shared weight encoder network f θ , used to extract deep semantic feature maps of images; The pseudo mask generation module is used to generate a pseudo foreground mask of the query image; The prototype correction module is used to balance the correction degree of the query prototype extracted from the query feature to the support prototype obtained from the support feature, and to perform weighted summation of the support prototype and the query prototype to generate a correction prototype; The prediction module is used to perform image segmentation on the image to be segmented based on the image segmentation method.
10. Application of the image segmentation method according to any one of claims 1 to 8, or the image segmentation system according to claim 9, in unsupervised small sample image segmentation.