A Few-Shot Medical Image Segmentation Method Based on Bidirectional Guidance Prototype Alignment
By introducing two-way guided prototype alignment and adaptive prototype modules in small sample medical image segmentation, the problems of prototype deviation and local information loss are solved, and high-quality medical image segmentation under small sample conditions are achieved.
Patent Information
- Application Number
- CN202311428266.1
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2023-10-31
- Publication Date
- 2025-06-13
- Estimated Expiration
- 2043-10-31
AI Technical Summary
The existing small sample medical image segmentation method has differences between the medical image mode and the natural image mode, resulting in prototype deviation and local information loss, and the direct transplantation of natural image segmentation method is not ideal.
A small sample medical image segmentation method based on bidirectional guided prototype alignment is proposed. By querying prototype matching query features, the correlation consistency within the image is maintained, and a dual-boot optimization process is introduced between the query guide branch and the support guide branch, combining the adaptive prototype module and the adaptive prototype alignment loss, and improving prototype learning.
It effectively alleviates the problem of prototype deviation, improves the accuracy and performance of the segmentation model, and can perform high-quality medical image segmentation under small sample conditions, which is suitable for medical fields where data is scarce.
Smart Images

Figure CN117314884B_ABST
Abstract
Description
Technical Field
[0001] The present invention belongs to the technical fields of computer vision and medical artificial intelligence, and particularly relates to a small-sample medical image segmentation method based on bidirectional guidance prototype alignment. Background Art
[0002] Medical image segmentation is the focus of medical image analysis, and its goal is to depict regions such as organs and lesions, which is crucial for computer-aided treatment and intelligent medicine. With the rapid development of deep learning, many supervised medical image segmentation models have emerged. However, these models rely heavily on widely labeled datasets, and it is particularly challenging to obtain a large number of medical pixel-level labels. Since few-shot learning (FSL) reduces data dependence, the challenges emphasized above have prompted researchers to integrate FSL into medical image segmentation. Although few-shot segmentation (FSS) technology has made considerable progress, many works are mainly used for natural image segmentation. Because medical image patterns are very different from general domain images, it has been proven that directly transplanting these methods into medical image segmentation does not yield ideal results. Despite the great potential of medical image segmentation, few-shot medical image segmentation remains largely unknown and nascent.
[0003] In the field of medical image segmentation, most existing FSS methods are based on prototype learning, and the support prototypes extracted by them often lack consistency with query features, and prototype deviation problems are prone to occur. In addition, although these methods have proposed many innovative ideas, there has been no improvement in the query-guided branch (using the query image to reverse-segment the support image). However, most prototype methods use mask average pooling, which also causes a large amount of loss of local information. Summary of the Invention
[0004] In view of the above problems existing in the prior art, the present invention proposes a small-sample medical image segmentation method based on bidirectional guidance prototype alignment, which uses query prototypes to match query features, so as to maintain the correlation consistency within the image. In addition, a dual-guidance optimization process is proposed between the query-guided branch (query to support) and the support-guided branch (support to query), thereby further alleviating the prototype deviation problem. And an Adaptive Prototype Module (APM) is proposed, which organizes similar local information on the basis of extracting global information and retaining valuable local information; finally, an adaptive prototype alignment loss is introduced to better improve the prototype.
[0005] To achieve the above object, the technical solution of the present invention is as follows:
[0006] A small-sample medical image segmentation method based on bidirectional guidance prototype alignment, comprising the following steps:
[0007] Step 1: Use a guiding-supported branch as the main branch, and follow the traditional paradigm to use the support set as a guide to segment the query image, and obtain the feature representations F of the support set and the query set corresponding medical images from the ResNet 101 encoder network S and F Q , where F S is the support feature and F Q is the query feature;
[0008] Step 2: Use the masked average pooling operation to learn the support prototype P S from F S , and estimate the initial query prediction mask by calculating the cosine similarity between P S and F Q . Use the initial query prediction mask M q and F Q to generate the adaptive query prototype P Q through the adaptive prototype module;
[0009] Step 3: The adaptive query prototype P Q and the support prototype P S are weighted aggregated to form the aggregated prototype P. We calculate the cosine similarity between the aggregated prototype P and the query feature F Q to obtain the final query prediction;
[0010] Step 4: The final query prediction of the main branch and the feature representations F S 、F Q are jointly input into another query-guided branch to obtain the support prediction result; then use the few-shot medical image segmentation model composed of the bidirectional-guided double-branch structure for training;
[0011] Step 5: Load the model in Step 4, input the required medical image into the trained few-shot medical image segmentation model, and obtain the corresponding segmentation mask image.
[0012] Based on the above scheme, this method is applicable to the medical image segmentation task with small samples. By introducing the support set and the query set, it can effectively perform image segmentation without a large number of labeled samples, which is very useful for the situation of scarce data in the medical field. Adopting a two-way guidance method, combining the main branch and the query-guided branch, can improve the performance of the segmentation model. The main branch uses the support set for segmentation, while the query-guided branch provides more information to improve the segmentation result by using the result of the main branch for support prediction. Through the prototype learning of support features and query features, it can better capture the features and structures in medical images, thereby improving the segmentation accuracy. Prototype learning can help the model better understand the features of different categories. Using an adaptive prototype module can further improve the segmentation performance. It allows the model to dynamically generate query prototypes according to the features of each query image, thus better adapting to different image contents. By aggregating the adaptive query prototype and the support prototype into an aggregated prototype, different information can be comprehensively utilized, thereby improving the accuracy of segmentation prediction. This aggregation can effectively combine the information of the two guidance branches. Using cosine similarity to estimate and calculate the similarity between prediction results helps to reduce the noise in the feature space and improve the recognition of tiny details in medical images. Generally speaking, this method can effectively handle the medical image segmentation task with small samples, and through techniques such as prototype learning, adaptive modules, and two-way guidance, it improves the accuracy and performance of the segmentation model, and is expected to provide better results in the field of medical image analysis.
[0013] Furthermore, step 2 specifically includes:
[0014] Step 2.1: First, use the traditional mask average pooling (MAP) operation to generate the support prototype P S , as shown in formula (1):
[0015] P S = MAP(M S , F S ) (1)
[0016] Where, M S represents the ground truth mask;
[0017] Step 2.2: For the negative cosine similarity -S between the query feature vector F Q and the support prototype P S , use the shifted sigmoid function σ; then, adopt a soft threshold operation with the learned threshold τ to generate the foreground probability map M q,f , as shown in formula (2) and formula (3):
[0018] M q,f = 1 - σ(0.5(-S(P S , F Q)-τ)) (2)
[0019] where
[0020]
[0021] in the formula, κ = 20 is the introduced scaling factor; · represents the dot product operation; ||·|| represents the second norm; in addition, through the cat operation, the initial foreground probability map M q,f and the initial background prediction map are concatenated to generate the initial query prediction mask M q ={1 - M q,f , M q,f};
[0022] Step 2.3: Jointly input the generated initial query prediction mask M q and the query feature F Q into our Adaptive Prototype Module (APM) to obtain the adaptive query prototype P Q ={P Q,f , P Q,b , P Q,lb};
[0023] Considering the obvious similarity among foreground pixels and the less obvious semantic consistency among background pixels, different threshold filters are established for the foreground and background prediction masks, as shown in formulas (4a) and (4b):
[0024]
[0025]
[0026] in the formula, ψ represents the threshold filter function. ψ retains the pixels greater than the thresholds T q and T f and T b in M ; through the threshold filtering operation, new foreground mask
[0027] and background mask Q,f and global background prototype P Q,b are generated, as shown in formulas (5a) and (5b):
[0028]
[0029]
[0030] The query feature is reconstructed by performing mask multiplication between the filtered background mask and the query feature. Then, based on the reconstructed query feature With the original query feature F Q Construct the similarity matrix by matrix multiplication between Its structure is expressed by formula (6):
[0031]
[0032] in ⊙ represents mask multiplication operation; T represents matrix transposition operation; Represents matrix multiplication operation;
[0033] The similarity matrix is normalized by the softmax operation along the first dimension, and the normalized similarity matrix is used together with the reconstructed query features. The local background prototype P is generated by matrix multiplication between the transpose of Q,lb , expressed by formula (7):
[0034]
[0035] Where φ represents the softmax function.
[0036] Furthermore, the step 3 specifically includes:
[0037] Step 3.1: The adaptive query prototype P output by the adaptive prototype module (APM) in step 2.3 Q With the support prototype P in step 2.1 S Perform weighted aggregation; generate aggregation prototype P = {P f ,P b}, using formula (8a) and formula (8b):
[0038] P f =αP S +(1-α)P Q,f (8a)
[0039] P b =βP Q,lb +(1-β)P Q,b (8b)
[0040] Among them, α and β are the corresponding aggregation weights;
[0041] Step 3.2: Calculate the aggregation prototype P and query features F Q to generate the final query prediction mask It is expressed by formula (9):
[0042]
[0043] Where φ represents the sofmtx function and S is the negative cosine similarity.
[0044] Furthermore, step 4 specifically includes:
[0045] Step 4.1: For a given query prediction mask, initially use masked average pooling to generate query prototypes; subsequently, utilize the Adaptive Prototype Module (APM) to extract adaptive support prototypes; to further eliminate prototype bias, merge the query prototypes and the adaptive support prototypes;
[0046] Step 4.2: Calculate the cosine similarity between the merged prototypes and the support features to generate the final support prediction mask
[0047] Step 4.3: Train the proposed model by calculating two types of losses, including the segmentation loss and the alignment loss The segmentation loss is measured by calculating the cross-entropy between the predicted query segmentation mask and the corresponding ground truth, while the alignment loss is used to measure the segmentation error between the support prediction mask and the corresponding ground truth of the support features; next, the overall loss of the model is expressed by formula (10):
[0048]
[0049] Furthermore, step 4.3 specifically includes:
[0050] Step 4.3.1: First, the segmentation loss consists of two loss terms: the initial prediction loss and the final prediction loss Calculate the segmentation loss expressed by formula (11):
[0051]
[0052] where λ 1 and λ 2 control the balance of these two loss terms, which are set to 0.7 and 0.3 respectively;
[0053] Obtain the final prediction loss by calculating the binary cross-entropy loss BCE between the final query prediction mask and the corresponding ground truth mask M Q of the query image, expressed by formula (12):
[0054]
[0055] Obtain the initial prediction loss by calculating the binary cross-entropy loss BCE between the initial prediction mask M q and the corresponding ground truth mask M QThe binary cross - entropy loss BCE between them to obtain the initial prediction loss It is expressed by formula (13):
[0056]
[0057] Step 4.3.2: Then calculate the alignment loss It consists of three loss terms, including the support prediction loss of the query - guiding branch the support prototype loss and the proposed adaptive prototype alignment loss It is expressed by formula (14):
[0058]
[0059] where λ 3 、λ 4 and λ 5 represent weights, which are set to 0.2, 0.4, and 0.4 respectively according to experience;
[0060] The support prediction mask generated by the query - guiding branch is calculated and the corresponding ground - truth mask M of the query image S The binary cross - entropy loss BCE between them is used to obtain the support prediction loss It is expressed by formula (15):
[0061]
[0062] The support prediction result generated by using the support prototype P S to match the support feature F S is calculated and M S The binary cross - entropy loss BCE between them is used to obtain the support prototype loss It is expressed by formula (16):
[0063]
[0064] where σ is the shifted sigmoid function; S is the negative cosine similarity, and τ is the learning threshold;
[0065] The support feature F S is matched with the adaptive aggregation prototype P to generate the support prediction result; then the binary cross - entropy loss BCE between the support prediction result and the ground - truth mask M S is calculated to obtain the adaptive prototype alignment loss It is expressed by formula (17):
[0066]
[0067] where φ represents the sofmtx function.
[0068] Further, step 5 specifically includes:
[0069] Load the trained model model best in step 4, input the medical image into the model, and output the segmentation result and corresponding metrics for the medical image.
[0070] Advantages of the present invention: The present invention proposes a dual-guided prototype alignment network, which effectively optimizes the network by swapping the support and query roles; secondly, it alleviates the inherent data scarcity problem of the FSS method and provides a more powerful optimization strategy; by using different thresholds for different situations to aggregate similar pixels to adaptively generate high-quality prototypes, thereby focusing on more useful semantic features; in addition, the present invention further optimizes the adaptive prototype through the adaptive prototype alignment loss, thereby better improving the prototype quality; by minimizing the cross-entropy loss between the predictions of the two branches and the ground truth labels to optimize the network, a comprehensive two-way optimization process is established, effectively improving the accuracy of the entire model. BRIEF DESCRIPTION OF THE DRAWINGS
[0071] Figure 1 It is a framework diagram of a few-shot medical image segmentation network based on bidirectional guidance prototype alignment;
[0072] Figure 2 It is a schematic diagram of the adaptive prototype module in the network. DETAILED DESCRIPTION OF THE EMBODIMENTS
[0073] The embodiments of the present invention are implemented on the premise of the technical solution of the present invention, and detailed implementation manners and specific operation processes are given, but the protection scope of the present invention is not limited to the following embodiments.
[0074] The present invention provides a few-shot medical image segmentation method based on bidirectional guidance prototype alignment, which uses query prototypes to match query features to maintain the correlation consistency within the image. By swapping the support and query roles to effectively optimize the network, a dual-guided optimization process is proposed between the query-guided branch (query to support) and the support-guided branch (support to query), thereby further alleviating the prototype deviation problem. In addition, the present invention proposes an adaptive prototype module, which adaptively generates prototypes by aggregating similar pixels using different thresholds for different situations, and it organizes similar local information on the basis of extracting global information and retaining valuable local information. In addition, an adaptive prototype alignment loss is also proposed to further optimize the adaptive prototype, thereby better improving the prototype quality.
[0075] Embodiment 1
[0076] This embodiment uses the Windows system as the development environment, Pycharm as the development platform, and Python as the development language. The few-shot medical image segmentation method based on bidirectional guidance prototype alignment of the present invention is adopted to complete the segmentation mask prediction for the region of interest in medical images.
[0077] In this embodiment, the few-shot medical image segmentation method based on bidirectional guidance prototype alignment includes the following steps:
[0078] Step 1: Divide the dataset into a query set and a support set, preprocess the original data in the dataset, and construct paired support-query image pairs as training data;
[0079] Step 2: Construct a few-shot medical image segmentation network based on bidirectional guidance prototype alignment, and input the preprocessed training data and the pre-trained weights of the ResNet 101 encoder into the few-shot medical image segmentation network (as Figure 1 shown);
[0080] Step 3: Use binary cross-entropy loss as the training loss function for training to obtain the final trained few-shot medical image segmentation model;
[0081] Step 4: Take the medical image to be segmented as the input, load the model saved after training in Step 3, and obtain the segmentation mask image corresponding to the input medical image and the corresponding evaluation metrics. The average Dice score (DSC) will be used as the evaluation metric. Its calculation method can be represented by formula (18), where A and B represent the final segmentation prediction mask of the model and the corresponding ground truth respectively. The higher the Dice score, the better the segmentation performance.
[0082]
[0083] According to the above steps, the present invention is compared with the SE-Net model, ALPNet model, SSL-ALPNet model, SSL-PANet model, ADNet model, CRAPNet, and Q-Net model, etc. on three different test sets (Abdominal-CT, Abdominal-MRI, and Cardiac-MRI) under two different experimental settings (Setting: background slices are not excluded and Setting 2: background slices are excluded). It can be seen from Table 1, Table 2, and Table 3 that the method proposed by the present invention is basically superior to other methods in terms of segmentation accuracy on these three common test sets.
[0084] Table 1 Quantitative comparison with state-of-the-art models on the Abdominal-CT and Abdominal-MRI datasets under Setting 1 (Note: The best and second-best results are highlighted and underlined, respectively)
[0085]
[0086]
[0087] Table 2 Quantitative comparison with state-of-the-art models on the Abdominal-CT and Abdominal-MRI datasets under Setting 2 (Note: The best and second-best results are highlighted and underlined, respectively)
[0088]
[0089] Table 3 Quantitative comparison with state-of-the-art models on the Cardiac-MRI dataset under Setting 1 (Note: The best and second-best results are highlighted and underlined, respectively)
[0090]
[0091] The foregoing description of specific exemplary embodiments of the present invention is for purposes of illustration and exemplification. These descriptions are not intended to limit the invention to the precise forms disclosed, and it is apparent that many changes and variations are possible in light of the above teaching. The purpose of selecting and describing the exemplary embodiments is to explain the specific principles of the invention and its practical applications, so that those skilled in the art can implement and utilize various different exemplary embodiments of the invention, as well as various different selections and changes. The scope of the invention is intended to be defined by the claims and their equivalents.
Claims
1. A few-shot medical image segmentation method based on bidirectional guidance prototype alignment, characterized in that, the method comprises the following steps: Step 1: Use a guided branch as the main branch, and follow the traditional paradigm to split the query image using the support set as a guide, and obtain the feature representations F of the support set and the query set corresponding medical images from the ResNet 101 encoder network S and F Q , where F S is the support feature and F Q is the query feature; Step 2: Use masked average pooling operation to learn the support prototype P from the support feature F S ; Estimate the initial query prediction mask M by calculating the cosine similarity between the support prototype P S and the query feature F S ; Generate the adaptive query prototype P through the adaptive prototype module APM using the initial query prediction mask M Q and the query feature F q ; q ; Q ; Q ; Step 3: Aggregate the adaptive query prototype P Q and the support prototype P S to form an aggregated prototype P through weighted aggregation, and calculate the cosine similarity between the aggregated prototype P and the query feature F Q to obtain the final query prediction; Step 3 specifically comprises: Step 3.1: Weightedly aggregate the adaptive query prototype P output by the adaptive prototype module APM Q with the support prototype P S to generate an aggregated prototype P = {P f , P b}, using formulas (8a) and (8b): P f = αP S + (1 - α)P Q,f (8a) P b = βP Q,lb + (1 - β)P Q,b (8b) where α and β are corresponding aggregation weights; P Q,f represents the global foreground prototype, P Q,b represents the global background prototype, P Q,lb represents the local background prototype; Step 3.2: Calculate the cosine distance between the aggregated prototype P and the query feature F Q to generate the final query prediction mask which is expressed by formula (9): where φ represents the softmax function and S is the negative cosine similarity; Step 4: Co-input the final query prediction and feature representation F of the main branch S and F Q into another query-guided branch to obtain the support prediction result; then use the few-shot medical image segmentation model composed of the bidirectional-guided double-branch structure for training; Step 5: Load the model in Step 4, input the required medical image into the trained few-shot medical image segmentation model, and obtain the corresponding segmentation mask image.
2. The few-shot medical image segmentation method based on bidirectional guidance prototype alignment according to claim 1, characterized in that, Step 2 specifically comprises: Step 2.1: First, generate the support prototype P using the traditional masked average pooling MAP operation, as shown in Equation (1): S , as shown in Equation (1): P S = MAP(M S , F S ) (1) where M S represents the ground truth mask; Step 2.2: For the negative cosine similarity -S between the query feature F Q and the support prototype P S use the shifted sigmoid function σ; then, perform a soft thresholding operation with the learned threshold τ to generate the foreground probability map M q,f , as shown in formulas (2) and (3): M q,f = 1 - σ(0.5(-S(P S , F Q )) - τ)) (2) where where κ = 20 is the introduced scaling factor; · represents the dot product operation; ||·|| represents the second norm; in addition, through the cat operation, the initial foreground probability map M q,f and the initial background prediction map are concatenated to generate the initial query prediction mask M q = {1 - M q,f , M q,f}; Step 2.3: Input the generated initial query prediction mask M q and the query feature F Q jointly into the adaptive prototype module APM to obtain the adaptive query prototype P Q ={P Q,f ,P Q,b ,P Q,lb}; Among them, P Q,f represents the global foreground prototype, P Q,b represents the global background prototype, P Q,lb represents the local background prototype; different threshold filters are established for the foreground and background prediction masks, as shown in formulas (4a) and (4b): where ψ represents the threshold filter function, and ψ retains the pixels greater than the threshold T q in M f and T b ; through the threshold filtering operation, new foreground mask and background mask Then, use the Masked Average Pooling (MAP) operation to generate the global foreground prototype P Q,f and the global background prototype P Q,b , as shown in Equations (5a) and (5b): The query feature is reconstructed by performing masked multiplication between the filtered background mask and the query feature, and then, based on the reconstructed query feature and the original query feature F Q a similarity matrix is constructed through matrix multiplication whose structure is represented by formula (6): Among them ⊙ represents the masked multiplication operation; T represents the matrix transpose operation; represents the matrix multiplication operation; Normalize the similarity matrix through a softmax operation along the first dimension, and generate the local background prototype P by matrix multiplication between the normalized similarity matrix and the transpose of the reconstructed query feature as shown in Equation (7): Q,lb which is expressed by formula (7): where φ represents the softmax function.
3. The few-shot medical image segmentation method based on bidirectional guidance prototype alignment according to claim 1 or 2, characterized in that, Step 4 specifically comprises: Step 4.1: For the given query prediction mask, initially use mask average pooling to generate the query prototype; subsequently, use the adaptive prototype module APM to extract the adaptive support prototype, and merge the query prototype and the adaptive support prototype; Step 4.2: Calculate the cosine similarity between the merged prototype and the support feature F to generate the final support prediction mask S Step 4.3: Train the proposed model by calculating two types of losses, including the segmentation loss and the alignment loss The segmentation loss is measured by calculating the cross-entropy between the predicted query segmentation mask and the corresponding ground truth, and the alignment loss is used to measure the segmentation error between the support prediction mask and the corresponding ground truth of the support features. Next, the overall loss of the model is expressed by Equation (10):
4. The medical visual question answering method based on global visual information intervention according to claim 3, characterized in that, Step 4.3 specifically comprises: Step 4.3.1: First, split the loss It consists of two loss terms: the initial prediction loss and the final prediction loss Calculate the split loss It is expressed by formula (11): where λ 1 and λ 2 control the balance of these two loss terms; By calculating the binary cross-entropy loss BCE between the final query prediction mask and the corresponding ground truth mask M of the query image Q the final prediction loss is obtained which is expressed by Equation (12): By calculating the binary cross-entropy loss BCE between the initial prediction mask M q and the corresponding ground truth mask M of the query image Q the initial prediction loss is obtained which is expressed by Equation (13): Step 4.3.2: Then calculate the alignment loss which consists of three loss terms, including the support prediction loss of the query guidance branch the support prototype loss and the adaptive prototype alignment loss which is expressed by Equation (14): where λ 3 , λ 4 and λ 5 represent weights; Calculating the support prediction mask generated by the query guidance branch and the corresponding ground truth mask M of the query S to obtain the support prediction loss by the binary cross-entropy loss BCE which is expressed by Equation (15): By calculating the use of the support prototype P S to match the support prediction results generated by the support features with M S the binary cross-entropy loss BCE between them to obtain the support prototype loss is expressed by formula (16): where σ is the shifted sigmoid function; S is the negative cosine similarity, and τ is the learning threshold; Use the support feature F S Match with the adaptive aggregation prototype P to generate a support prediction result; then calculate the binary cross-entropy loss BCE between the support prediction result and the ground truth mask M S to obtain the adaptive prototype alignment loss Denoted by formula (17): where φ represents the softmax function.
5. The few-shot medical image segmentation method based on bidirectional guidance prototype alignment according to claim 4, characterized in that, The said λ 1 and λ 2 are respectively set to 0.7 and 0.3; λ 3 , λ 4 and λ 5 are respectively set to 0.2, 0.4 and 0.4.