Few-shot medical image segmentation method based on edge-aware multi-prototype learning
By using an edge-aware multi-prototype learning method, multi-scale prototypes are generated, which solves the problem of edge detail loss in the segmentation of medical images with few samples in existing models, and improves the segmentation accuracy and robustness, especially the ability to locate boundaries and preserve details in complex medical image scenes.
Patent Information
- Application Number
- CN202511086299.1
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2025-08-05
- Publication Date
- 2025-10-28
- Estimated Expiration
- 2045-08-05
AI Technical Summary
Existing few-sample medical image segmentation models tend to lose useful category information when generating category feature prototypes, resulting in inaccurate segmentation results. This is especially true when dealing with complex textures and structural features, where local detail information is lost, making it difficult to capture long-distance dependencies and affecting diagnostic and treatment outcomes.
We adopt an edge-aware multi-prototype learning approach, which generates multi-scale prototypes through a feature encoder, a local attention fusion prototype generator, a two-stage prototype optimization network, and a loss calculation module. Combined with an edge-aware loss function, this improves the model's segmentation performance under data-scarce conditions.
It significantly improves the segmentation performance of the model on the CHAOS, SABS and CMR datasets, especially in completely unseen class scenes, with an average Dice score improvement of 1.25%-3.05% compared to the second-best method, and improves the ability to localize boundaries and preserve details.
Smart Images

Figure CN120599269B_ABST
Abstract
Description
Technical Field
[0001] This invention relates to the field of medical image segmentation technology, specifically to a few-sample medical image segmentation method based on edge-aware multi-prototype learning. Background Technology
[0002] Medical image analysis technology plays an increasingly important role in the clinical field. Using methods such as MRI, CT, and ultrasound, anatomical structures and pathological regions can be accurately depicted. Computers then automatically process and analyze the image data to provide doctors with scientific and accurate decision support. With the advent of the big data era, especially the widespread application of deep learning, the efficiency and accuracy of medical image analysis have been significantly improved, directly impacting the accuracy of disease diagnosis and surgical planning. Although deep learning has made significant progress in image segmentation, existing segmentation models (such as U-Net and its variants) typically rely on large-scale labeled data for training. However, in real-world clinical scenarios, obtaining large amounts of high-quality labeled data often faces numerous challenges. Labeling medical images requires not only specialized domain knowledge but also significant time and economic costs, resulting in a scarcity of labeled data. Under such limited data conditions, traditional fully supervised segmentation models are prone to overfitting on small training sets, severely limiting their generalization performance and robustness in practical applications. To address this problem, few-shot learning has emerged. Few-shot learning is a machine learning paradigm designed to quickly adapt to new tasks using a very small number of labeled samples.
[0003] In the field of Few-Shot Medical Image Segmentation (FSMIS), several studies have proposed models based on prototype networks and two-branch structures. These models generate class feature prototypes using average pooling (MAP) operations and apply them to segment query images. However, simple prototype generation methods inevitably lose some useful class information, leading to decreased segmentation accuracy, especially in medical images where a single prototype cannot effectively capture class diversity due to differences in morphology, texture, and structure among different tissues and organs. Therefore, improving prototype generation methods and optimizing model performance have become current research hotspots to address this issue. Existing FSMIS models based on prototype networks and two-branch structures typically employ simple feature extraction and pooling operations to generate class prototypes or feature representations. When processing query images, these models often directly utilize the feature information of the support set for segmentation, failing to fully model the dynamic relationships between features. Furthermore, this simple pooling operation easily leads to the loss of local detail information, particularly noticeable when processing complex texture and structural features in medical images. For the detection of early lesions or small lesions, the lack of local details can lead to missed or misdiagnosed cases. In clinical applications requiring precise boundary delineation (such as surgical planning), this deficiency can severely impact diagnostic and treatment outcomes. Traditional models also struggle to capture long-distance dependencies, fail to effectively utilize global contextual information, and lack dynamic adjustment capabilities when processing feature interactions between the support and query sets, resulting in inaccurate segmentation results. Therefore, further processing of the extracted features is necessary to address the issue of lost edge details. Summary of the Invention
[0004] The purpose of this invention is to provide a few-shot medical image segmentation method based on edge-aware multi-prototype learning, which can effectively solve the problem of edge detail loss in the background art.
[0005] To solve the above-mentioned technical problems, the present invention adopts the following technical solution:
[0006] A few-shot medical image segmentation method based on edge-aware multi-prototype learning is proposed. The edge-aware multi-prototype learning model includes five key modules: a feature encoder, a local attention fusion prototype generator, a two-stage prototype optimization network, prototype prediction, and loss calculation. The method includes the following steps:
[0007] S1. Input the support image and query image into the feature encoder to extract support feature maps and query feature maps of different sizes;
[0008] S2. The support feature map is adjusted to the same size as the mask using a bilinear interpolation algorithm, and then element-wise multiplication is performed. Next, the pixels in the target area are summed by mask weights and normalized using a mask normalization factor to generate a support foreground prototype. At the same time, the mask image is processed by dynamic erosion to obtain the inner boundary mask. The feature map is then weighted by weighted average pooling to generate the inner boundary prototype.
[0009] S3. Multiply the interpolated feature map with the mask, remove zero values and sample an equal number of background points to fill, reshape the feature vector, and perform weighted optimization. Through the cross-weighting mechanism, capture the local and global correlations and generate a multi-foreground local prototype.
[0010] S4. Perform multi-scale transformation on the support feature map and its corresponding binary mask to extract and optimize local and global information to obtain a multi-scale prototype; then fuse the support foreground prototype, inner boundary prototype, multi-foreground local prototype and multi-scale prototype to obtain a multi-prototype foreground prototype.
[0011] S5. The two-stage prototype optimization network is used to dynamically calculate and weight the multi-prototype foreground prototypes and perform automatic calibration to obtain the weighted aggregated multi-prototype foreground prototypes.
[0012] S6. Calculate the negative cosine similarity of the query feature map, and input the negative cosine similarity of the query feature map and the weighted aggregated multi-prototype foreground prototype into the prototype prediction module to make predictions by balancing semantics and details. Finally, perform collaborative optimization through the loss calculation module.
[0013] The edge-aware multi-prototype learning-based few-shot medical image segmentation method provided in the above technical solution significantly improves the model's segmentation performance under data-scarce conditions by using an innovative Local Attention Fusion Prototype Generator (LAF) and a Two-Stage Prototype Optimization Network (DPO), combined with an edge-aware loss function and a multi-scale feature fusion strategy. Experiments on three public datasets—CHAOS, SABS, and CMR—show that this method outperforms existing state-of-the-art methods in two different experimental settings. Particularly in the more challenging scenario of completely unseen categories, the average Dice score is improved by 1.25%-3.05% compared to the second-best method, validating its advantages in boundary localization and detail preservation. This invention provides an effective solution for few-shot medical image segmentation and has significant application value in clinical auxiliary diagnosis. Attached Figure Description
[0014] Figure 1 This is a diagram illustrating the overall architecture of the edge-aware multi-prototype learning method of this invention.
[0015] Figure 2 This is a flowchart illustrating the feature encoder process.
[0016] Figure 3 This is a flowchart illustrating the channel weight aggregation module.
[0017] Figure 4 This is a flowchart illustrating the spatial weight aggregation module.
[0018] Figure 5 This is a flowchart illustrating the cross-attention module.
[0019] Figure 6 This is a box plot showing the results of this embodiment;
[0020] Figure 7 The segmentation results for the CHAOS dataset;
[0021] Figure 8 The segmentation result of the SABS dataset;
[0022] Figure 9 The segmentation results for the CMR dataset;
[0023] Figure 10 This is the loss curve. Detailed Implementation
[0024] To make the objectives and advantages of this invention clearer, the invention will be specifically described below with reference to embodiments. It should be understood that the following text is merely used to describe one or more specific embodiments of the invention and does not strictly limit the scope of protection specifically claimed by the invention.
[0025] like Figure 1 As shown, the edge-aware multi-prototype learning of this invention consists of five key modules: feature encoder, locality-attentive fusion prototyper, dual-stage prototype optimization network, prototype prediction, and loss calculation.
[0026] First, the support and query images are input into a parameter-shared feature encoder to extract support and query features of different sizes, generating support and query feature maps. Then, the support mask is processed through dynamic erosion to obtain an inner boundary mask, which is then weighted by average pooling to generate an inner boundary prototype. Next, the interpolated feature map is multiplied by the mask, zero values are removed, and an equal number of background points are sampled for filling, reshaping the background feature vector and performing weighted optimization. A cross-weighting mechanism is used to capture local and global correlations, and then an MLP is used to generate multiple foreground local prototypes. The support feature map and its corresponding binary mask are then subjected to multi-scale transformation to extract and optimize local and global information, resulting in multi-scale prototypes. These prototypes are then fused, and a lightweight attention network is used to dynamically calculate the weighted prototype, which is automatically calibrated using cosine similarity. Finally, hierarchical prototype matching and multi-scale feature fusion are used to balance semantics and details for prediction, and alignment loss and edge-aware loss are used for collaborative optimization.
[0027] The feature encoder used in this embodiment is based on the DeepLabv3 pre-trained ResNet-101 architecture, achieving efficient feature extraction by loading pre-trained weights from the COCO dataset. In terms of network structure, the original five layers of ResNet-101 are retained, while two key modules are added: firstly, 1×1 convolutional dimensionality reduction layers are added after layer 3 and layer 4, compressing the number of channels from 1024 / 2048 to a uniform 512 dimensions; secondly, a global average pooling layer and a fully connected layer are introduced to extract image-level global features. A layered initialization strategy is adopted during training, where the pre-training part retains its original parameters, the newly added convolutional layers are initialized using the Kaiming normal distribution, and fine-tuned using a progressive unfreezing strategy; for example... Figure 2 As shown, LAFP generates multiple representative descriptors to comprehensively represent the commonalities of the category distribution; the final output contains a 512-dimensional multi-scale feature map (1 / 8 and 1 / 16 resolution) and a scalar value representing global scene information.
[0028] In tasks such as medical image segmentation, accurate extraction of foreground regions is crucial for model performance. Especially in few-shot learning scenarios, due to limited labeled data, effectively learning and accurately segmenting foreground regions from a small number of support samples becomes a major challenge in model design. To address this issue, this invention proposes a local attention fusion prototype generator. This module, through sophisticated feature processing, including feature map interpolation, dynamic dilation / erosion masks, feature weighting, multi-scale feature extraction, and contextual information fusion, can accurately capture foreground regions in images, thereby improving segmentation accuracy and enhancing the model's adaptability in complex image scenes.
[0029] First, this invention employs a bilinear interpolation algorithm to adjust the spatial dimension of the input support feature map, ensuring it maintains the same size as the mask matrix. Then, by performing element-wise multiplication between the upsampled interpolated feature map and the mask, a weighted summation of each region is achieved. To extract region prototype features, this invention performs a mask-weighted summation operation on all pixels within the target region and standardizes them using a normalization factor from the mask matrix, ultimately generating a support foreground prototype. This process can be formally represented as:
[0030] ;
[0031] in This represents the value of the interpolated feature map at position (h, w). This represents the value of the mask at position (h, w).
[0032] The kernel size is dynamically calculated based on the mask size, ensuring it is no less than 3×3 to effectively control the intensity of the erosion operation. Subsequently, a morphological erosion operation is applied to shrink the target region in the mask image, eliminating edge noise and small false targets, thus generating an eroded mask. Finally, the foreground inner boundary mask is obtained by calculating the difference between the original mask and the eroded mask. The mathematical expression for this process is:
[0033] ;
[0034] ;
[0035] This represents the mask after two erosion processes, where Here, (a, b) and (a′, b′) are the pixel positions in the original mask, respectively, and the structuring element K represents the local coordinates of the structuring element K during the two erosion processes. During the first erosion, for each pixel (h, w), the minimum value within the neighborhood of the structuring element K is calculated, i.e., the minimum value within the neighborhood of the structuring element K is calculated. This shrinks the foreground region. The second erosion repeats this operation based on the first result, further refining the target region, ultimately yielding... By adaptively adjusting the kernel size, this method can maintain stable erosion effects at different resolutions, improving the accuracy of inner boundary detection to serve subsequent foreground prediction.
[0036] After generating the interpolated feature map and inner boundary mask, the norm of each pixel in the interpolated feature map is calculated as a weighting coefficient. A feature norm weighting strategy is used to suppress interference from high-intensity feature points. Subsequently, a weighted average pooling operation is performed on the interpolated feature map using the inner boundary mask and weighting coefficients to generate the inner boundary prototype. This approach highlights key features while also capturing fine-grained information about boundary regions. Mathematically, this can be represented as:
[0037] ;
[0038] ;
[0039] Where ⊙ represents the Hadamard product, To support the inner boundary prototype of the image, Weights controlled by the feature norm, This is the inner boundary mask.
[0040] To effectively focus on and utilize the key features of the foreground region, this invention first performs element-wise multiplication of the interpolated feature map F with the mask M to extract foreground region features, thus obtaining supporting foreground features. :
[0041] ;
[0042] A background feature map is then initialized, and a masking operation is used to ensure that the background region does not interfere with the extraction of foreground features. Next, a set of feature points for the foreground region is extracted from the feature map, and feature points with a mask value of 1 are selected. Mathematically, this can be represented as:
[0043] ;
[0044] in For the set of foreground feature points, only select A pixel with a value of 1 The spatial extent of the input feature map.
[0045] To maintain feature diversity in the foreground region, this invention randomly selects a specified number of feature vectors from the foreground region. Mathematically, this can be expressed as:
[0046] ;
[0047] in For the final randomly selected set of foreground features, Foreground feature set This indicates the number of foreground features that need to be sampled. This indicates the random selection of a required number of non-repeating features from a sequence;
[0048] However, if the number of available foreground pixels is insufficient, an oversampling strategy is employed, supplementing feature points by repeatedly selecting a subset of samples. Randomly sampled foreground features are filled into the original feature map, and positions where the background mask is zero are replaced, thereby optimizing the feature representation in the segmentation task. Mathematically, this can be represented as:
[0049] ;
[0050] This involves n original foreground feature stitching operations, where n represents... , This indicates that the data is rounded down, and some foreground features are randomly selected to supplement the insufficient number of features.
[0051] Finally, the zero-value regions of the foreground features are filled with the sampled feature points to obtain the final foreground features. Mathematically, this can be represented as:
[0052] ;
[0053] If and only if the background point is 0, these background point positions are replaced with randomly sampled prototype feature points. For logical judgment Is it 0?
[0054] To effectively capture the interaction between local and global information, while highlighting key regions and suppressing background noise interference, this invention proposes Spatial-Channel Cross Attention (SCCA). SCCA optimizes spatial features (foreground regions in the image) and channel features (the contributions of each feature channel) by simultaneously weighting spatial and channel dimensions, thereby improving the model's ability to extract important foreground features. SCCA mainly consists of three parts: (1) Channel-weighted pooling: enhancing the representation of discriminative feature channels through adaptive channel weight learning; (2) Spatial-weighted pooling: focusing on key regions using spatial attention mechanisms to improve local feature representation capabilities; (3) Cross attention: effectively capturing global contextual information by modeling long-range dependencies through bidirectional spatial attention.
[0055] like Figure 3As shown, channel-weighted pooling dynamically optimizes the importance representation of each channel by adaptively learning the weight distribution along the channel dimension of the feature map. This method first uses a 1×1 convolution to decompose the supporting foreground features into query features and value features. The query features undergo global average pooling to extract channel-level global statistics, and then a tensor reshaping operation is performed to conduct matrix multiplication with the value features to construct a correlation matrix between channels. This matrix is then processed by SoftMax normalization and Sigmoid activation to generate a channel-weighted mask. Finally, feature enhancement is achieved by multiplying the mask with the original supporting foreground features channel by channel. This mechanism effectively models the dependencies between channels, adaptively enhances the feature representation of important channels, and significantly improves the discriminative ability of the features. Channel attention is represented mathematically as follows:
[0056] ;
[0057] in These are foreground features extracted through random sampling of feature points. and They are 1 1 convolutional layer, , and It involves reshaping three tensors. It is a global average pooling operation. , It is a SoftMax operation. “ "This is the matrix dot product operation." It is the Sigmoid operation. , It is a channel multiplication operator.
[0058] like Figure 4 As shown, spatially weighted pooling dynamically learns the weight distribution across the spatial dimension of the support foreground features, adaptively optimizing the importance representation of features at each location. This method first decomposes the input support foreground features into query features and value features using 1×1 convolution. After tensor reshaping, the SoftMax-normalized query features and value features are subjected to matrix operations to generate spatial context features. Then, a spatial attention mask is generated using 1×1 convolution and Sigmoid activation. Finally, feature enhancement is achieved through element-wise multiplication of the mask with the original support foreground features. This mechanism effectively models the dependencies between spatial locations, adaptively strengthens the representation of key regions, and significantly improves the discriminative ability of features. Its mathematical representation is as follows:
[0059] ;
[0060] in Indicates spatial attention; These are foreground features extracted through random sampling of feature points. , and They are 1 1 convolutional layer, and It involves reshaping two tensors. It is a SoftMax operation. “ "This is the matrix dot product operation." It is the Sigmoid operation. , It is a space multiplication operator.
[0061] like Figure 5 As shown, the cross-attention module implements an efficient bidirectional spatial attention mechanism that enhances feature representation capabilities by simultaneously capturing long-range dependencies supporting foreground features in both height and width directions. This method first uses three 1×1 convolutions to generate query, key, and value feature projections, respectively. Then, it decomposes the feature map into two orthogonal directions—height and width—through dimensionality transformation. In the height branch, after feature reshaping, attention weights along the height direction are calculated using batch matrix multiplication, and an INF mask is introduced to ensure that only valid regions are considered. The width branch employs a symmetrical processing approach. After concatenation and SoftMax normalization, the attention maps from both directions are multiplied with their corresponding value features to achieve feature aggregation. Finally, the results from both directions are weighted and fused, and a learnable scaling factor is introduced. This design enables the module to effectively model global spatial relationships while maintaining low computational complexity. Residual connections ensure training stability, making it particularly suitable for capturing dependencies across space and channels. Its mathematical representation is as follows:
[0062] ;
[0063] ;
[0064] ;
[0065] ;
[0066] in and This represents the key features and value features being convolved with X using a 1×1 method, with processing applied to the height direction. and This represents the key features and value features being convolved with X using a 1×1 method, with the width direction being processed. and This indicates that a 1×1 convolution is performed on the query features, and the dimensions are transformed after processing in the height and width directions. The INF mask generates a matrix with negative infinity on the diagonal to avoid self-attention. “ "This is batch matrix multiplication." It decouples and separates the characteristics of sotfmax into height and width. It is a scaling factor that can dynamically control attention-enhanced features; This indicates a high level of attention. Note the width. Indicates cross attention; This refers to the characteristics of attention across spatial-channel intersections.
[0067] To comprehensively represent the feature information of each category, this invention designs a feature descriptor generation framework based on a multilayer perceptron (MLP). This network takes foreground and background padding features as input and generates highly discriminative category representation descriptors through deep nonlinear transformations. Specifically, this invention constructs a deep feedforward neural network with multiple hidden layers, its hierarchical structure being: an input layer (receiving flattened spatial features), a 2048-dimensional expansion layer (equipped with LeakyReLU activation and Dropout regularization), a 1024-dimensional compression layer (also employing LeakyReLU and Dropout), and finally, an output layer. This expansion-then-compression architecture, combined with a Dropout probability of 0.3 and a LeakyReLU activation function with a negative slope of 0.01, effectively mines high-order nonlinear patterns in the input features while ensuring that the generated representative descriptors are both discriminative and robust. Through an end-to-end training process, this MLP can automatically learn the optimal feature transformation, generating multiple semantically representative feature descriptors for each category, thereby significantly improving the performance of subsequent classification or segmentation tasks. (Multi-foreground local prototype) The mathematical representation is as follows:
[0068] ;
[0069] in, Flatten the input features into , for , where in represents the input features H and W. for , LeakyReLU(x) is equivalent to max(0.01x,x). It is a regularization layer with a probability of 0.3. for , for , for , for , out is the number of custom output prototypes.
[0070] This invention proposes an innovative multi-scale prototype generation framework that constructs a robust feature representation system through a hierarchical feature fusion mechanism. The method first processes the input features that have undergone spatial-channel cross-attention. The binary mask M is used for multi-scale transformation, processed in parallel across three typical scale spaces. For each scale space, a mask-guided feature aggregation strategy is employed to generate discriminative prototype vectors, where fine-grained scales focus on local structural features, and coarse-grained scales capture global contextual information. By adaptively fusing multi-scale prototype representations, the final feature representation system achieves synergistic optimization of local detail preservation and global semantic understanding. This end-to-end learnable framework significantly improves the model's feature discrimination capability in complex scenes while maintaining computational efficiency. Its mathematical representation is as follows:
[0071] ;
[0072] ;
[0073] Where s represents the three scales , It is a prototype generated at three scales. It is a feature map after scaling the original features. For binary masking, , Used to prevent division by zero errors. This represents a stacked prototype feature across three scales; This represents the height of the feature map at scale s. This represents the width of the feature map at scale s. Indicates the index in the height direction of the feature map. Indicates the index in the width direction of the feature map.
[0074] By integrating support for future prototypes Inner boundary prototype Multi-foreground local prototype and multi-scale prototypes The resulting foreground prototype set significantly improves medical image segmentation, especially on small sample datasets. This fusion strategy effectively integrates diverse feature information, including foreground, boundaries, local details, and features at different scales, enhancing the model's focus on boundary regions and improving its ability to capture details. Common problems such as boundary blurring and detail loss are effectively mitigated through this strategy. The combination of boundary enhancement and multi-scale information further improves segmentation accuracy and robustness. Fine-grained target region segmentation enables the model to achieve accurate segmentation at multiple scales, thus avoiding errors caused by boundary blurring and scale differences. Therefore, the fused foreground prototype set provides a more comprehensive and efficient solution for medical image segmentation, particularly demonstrating significant advantages in organ segmentation tasks.
[0075] ;
[0076] This indicates a serial operation.
[0077] The Adaptive Prototype Aggregation Module is a prototype optimization method based on a lightweight attention mechanism, aiming to improve the prototype representation capability of supporting set sample generation. Taking a multi-prototype foreground prototype P as input, it adaptively fuses multiple prototypes within the same category through learnable attention weights, thereby generating more discriminative and robust category representations. Specifically, the module first ensures that all prototype tensors reside on a unified computing device and preserves the feature structure inherent in few-shot tasks. The core component is a lightweight attention network containing only one linear transformation layer that maps dimensional features to scalar attention weights, and then uses Softmax normalization to ensure that the weights satisfy convex combination constraints. Finally, the module outputs a weighted sum of all prototypes within the same category, with weights dynamically calculated by the attention mechanism, allowing the model to automatically focus on more reliable or discriminative prototypes while suppressing noise or outlier samples. (Prototype set) The mathematical representation is as follows:
[0078] ;
[0079] in , representing the learnable weight matrix, maps multidimensional features to one dimension, and K is the number of supporting sample prototypes for the current class. , is the bias term (scalar). Transpose of the weight matrix It is a SoftMax operation. This indicates element-wise multiplication.
[0080] The adaptive prototype aggregation module fully considers the data scarcity characteristic of few-shot learning, boasting advantages such as small parameter count and computational efficiency, effectively avoiding overfitting. Its attention mechanism is end-to-end differentiable, enabling collaborative optimization with the entire model to adapt to different task data distributions. This adaptive aggregation strategy significantly improves prototype quality, resulting in stronger generalization ability in subsequent meta-testing stages. For example, in noisy support set samples, this module can automatically reduce the influence of low-quality prototypes while strengthening the representation of key features; even with significant intra-class differences, it can generate more stable prototypes through weighted fusion. Overall, the adaptive prototype aggregation module achieves prototype optimization at extremely low computational cost, providing an efficient and scalable solution for few-shot learning tasks.
[0081] In few-shot learning frameworks, the Prototype Calibration Module (PCM) is a lightweight post-processing module designed to optimize initial prototype representations through a data-driven, dynamically weighted strategy, significantly improving the model's discriminative performance and robustness under limited sample conditions. This module receives a set of prototypes from a pre-attention aggregation layer. First, it calculates the arithmetic mean of prototypes across categories as initial class centers. Then, it measures the directional consistency between each prototype and its class center using cosine similarity and transforms the similarity into normalized attention weights using the Softmax function. Based on these weights, the module performs weighted aggregation of prototypes, making the optimized class representations more focused on high-confidence sample regions while effectively suppressing noise or outliers. The entire process is fully differentiable and requires no learnable parameters, maintaining computational efficiency (adding only <1% to inference time) while avoiding overfitting risks in small-sample scenarios. Compared to traditional mean pooling or linear transformation methods, this module achieves non-linear calibration of the prototype space through cosine similarity weights, maintaining the compactness of intra-class features while adaptively filtering low-quality samples. Its core value lies in achieving automatic calibration of prototype representations through a purely data-driven approach, without requiring additional supervisory signals or complex parameter learning. This has significant practical implications for data-scarce, few-shot learning scenarios. Its mathematical representation is as follows:
[0082] ;
[0083] ;
[0084] in , Let K represent the prototype feature of the i-th supporting sample, and K represent the number of supporting sample prototypes for the current category. This is the formula for calculating cosine similarity. It is an exponential function; Represents the i-th prototype Cosine similarity to the arithmetic mean class center of all prototypes in the current category; Indicates the original prototype The prototype obtained after weighted aggregation; and Let represent the cosine similarity between the k-th and j-th prototypes and the class center, respectively; This represents the prototype feature vector extracted by the network for the k-th supporting sample.
[0085] This prototype prediction module implements a few-shot semantic segmentation framework based on multi-scale prototype matching. Its core idea is to achieve pixel-level classification by comparing the similarity between query image features and prototypes in the support set. Specifically, the system first constructs a hierarchical prototype matching mechanism: for each query image feature map, it calculates the negative cosine similarity between the feature map and the foreground prototypes of each category, and introduces a learnable scaling factor and a category-specific threshold to adjust the similarity metric space. The adjusted distance is converted into a probability value using the sigmoid function, forming the initial segmentation response map. Considering that deep features contain rich semantic information while shallow features retain more details, this method adopts a multi-scale feature fusion strategy—feature maps extracted from different network layers are prototype matched separately, then uniformly upsampled to the original image size using bilinear interpolation, and adaptively weighted using learnable weight coefficients (alpha) for the prediction results of different levels. This design allows the model to automatically balance the contributions of high-level semantics and low-level details; for example, it gives higher weight to shallow features in object boundary regions, while relying more on the discriminative power of deep features in homogeneous regions. The final prediction result is represented by a binary classification output in the form of softmax, concatenating the foreground and background probabilities to form a dual-channel output tensor. The entire process fully embodies the key characteristics of few-shot learning: on the one hand, knowledge transfer is achieved through prototype matching, using class prototypes of supporting samples as "conceptual anchors" to guide the segmentation of the query image; on the other hand, the effectiveness of feature representation is maximized under limited sample conditions through an end-to-end multi-scale fusion mechanism. Its mathematical representation is as follows:
[0086] ;
[0087] ;
[0088] in It's the Sigmoid function, which compresses the output to [0, 1]. This is the formula for calculating cosine similarity. It is a tensor concatenation function. These are learnable weights. It is bilinear interpolation upsampling; This represents the predicted probability value of the nth query sample belonging to a certain category, which is compressed to the [0,1] interval by the Sigmoid function; This represents the depth features of the nth query sample; The prototype feature representing a certain category is generated from the feature mean of similar samples in the support set or other aggregation methods; This represents the prediction result for all N samples. The final segmentation probability map is generated by stitching together weighted aggregation and upsampling. This represents the original prediction result for the nth sample.
[0089] The alignment loss function establishes consistency constraints between the support set and the query set through a prototype-based bidirectional verification mechanism. This loss function first extracts class prototype features using the query set prediction mask, and then performs reverse verification by calculating the cross-entropy loss between the support set prediction mask and the true label. By forcing spatial consistency between prototype and support set features, this loss function effectively alleviates the feature distribution shift problem in small-sample scenarios. Finally, it normalizes the data using the number of classes and samples to ensure gradient balance across different training episodes.
[0090] ;
[0091] ;
[0092] ;
[0093] in, It is a predictive mask. It is to obtain query features. It supports image features. It's a similarity calculation. It is bilinear interpolation upsampling. This is the category threshold, where K is the number of categories and S is the number of shots. It is an indicator function. It is cross-entropy loss. It supports set prediction masks; Represents the prototype of feature c; This represents the set of all pixels belonging to category c, typically defined by a predicted mask or a ground truth mask. The logical sigmoid function compresses real number inputs into the (0,1) interval, and its output can be interpreted as a probability. This represents a learnable weight parameter used to weight feature similarity at different levels or locations; Indicates the category label currently being processed; This represents the prediction mask for class c of the s-th sample in the support set, obtained through prototype... With supporting features Similarity calculation is used to generate the result; This represents a loss function that measures the difference between the predicted mask and the true mask in the support set, promoting alignment between the prototype and the features of the support set.
[0094] This paper proposes an edge-aware loss function that enhances boundary localization capabilities in few-sample semantic segmentation by explicitly modeling geometric structural features. The loss function employs a multi-stage collaborative optimization strategy: first, an edge detection module is constructed based on the classic Sobel operator, and then a horizontal convolutional kernel is used... and vertical convolution kernel Calculate the spatial gradient of the support set mask to generate an edge ground map with pixel-level accuracy. Subsequently, a lightweight edge prediction network was designed, which replicates the foreground probability map output by the model to form a three-channel input, and predicts the edge response map. Finally, end-to-end optimization is performed using weighted binary cross-entropy loss. This design has three significant advantages: (1) it accurately captures the geometric features of the target boundary through the anisotropic gradient calculation of the Sobel operator; (2) it adopts a prediction mask-guided edge detection mechanism, enabling the network to adaptively focus on semantic boundary regions; and (3) the collaborative optimization strategy with the main segmentation loss significantly improves the boundary quality while maintaining the segmentation accuracy of the main region. Its mathematical representation is as follows:
[0095] ;
[0096] ;
[0097] ;
[0098] ;
[0099] ;
[0100] in Limit the value to the range [0, 1]. The predicted mask is copied three times for a specific channel and then stitched together along the channel dimension. It is a lightweight edge detection network. It is a binary cross-entropy loss function. This represents a function that performs a convolution operation on the input feature map; This represents the depth feature representation extracted from the query image; This indicates the detection of vertical edges; Indicates the detection of horizontal edges; The edge probability map represented by the network is generated by a lightweight edge detection network; This represents the binary cross-entropy loss, which measures the difference between the predicted edge and the actual edge. Represents the true edge label at position (i,j); This represents the probability value that the model predicts the position (i,j) to be on the edge.
[0101] experiment:
[0102] 1. Datasets: A comprehensive method evaluation was conducted on three publicly available medical imaging datasets, covering a variety of anatomical structures and imaging modalities:
[0103] 1) CHAOS: This abdominal MRI dataset is taken from the ISBI 2019 Healthy Abdominal Organ Joint Segmentation Challenge. The dataset contains 20 3D T2-SPIR MRI scan images, each containing approximately 36 slices. This invention selects four common categories shared with the Synapse-CT dataset: left kidney, right kidney, liver, and spleen.
[0104] 2) SABS: This abdominal CT dataset is taken from the MICCAI 2015 Multi-Map Abdominal Organ Annotation Challenge. The dataset contains 30 3D abdominal CT scan images, from which four specific organs (i.e., left kidney, right kidney, liver, and spleen) were selected for evaluation.
[0105] 3) Cardiac MRI Dataset (CMR): This cardiac MRI dataset is taken from the MICCAI 2019 Multi-Sequence Cardiac MRI Segmentation Challenge. The dataset contains 35 3D cardiac MRI scan images, each image is divided into approximately 13 slices, and includes 3 different cardiac annotation labels: left ventricular blood pool (LV-BP), left ventricular myocardium (LV-MYO), and right ventricular myocardium (RV).
[0106] Dataset processing: First, pseudo-labels are generated based on 3D superpixel clustering for training. Then, 3D medical images undergo preprocessing such as normalization (zero mean, unit variance) and resampling to ensure consistency in image intensity and spatial resolution. Next, the 3D volume data is divided into 2D slices of 256×256 resolution, and channel duplication is used to adapt the model input. During training, geometric enhancement techniques such as random affine transformations (rotation, translation, scaling, shearing) and elastic deformation are used to improve model robustness, and optional Gamma correction can be used to adjust image contrast. During testing, training / validation data are strictly isolated, and support and query sets are sampled from different slices to simulate small sample scenarios.
[0107] 2. Implementation Details: The edge-aware multi-prototype learning (EML) model is implemented on a single NVIDIA GeForce RTX 3090 GPU configured with Python 3.9 and PyTorch 2.1.0. The model supports processing multiple medical image datasets such as CHAOS, CMR, and SABS, employs a 5-fold cross-validation strategy, and uses the SGD optimizer (initial learning rate 1e). -3 The model was trained using a momentum of 0.9 and a learning rate decay of 2% every 1000 iterations. Key few-shot learning parameters were set to a single support set (n_shot=1) and a single class (n_way=1), and model performance was improved through a dual-scale feature fusion mechanism (α=0.9). The model performed feature extraction based on hypervoxels (5000 hypervoxels in the abdominal dataset / 1000 hypervoxels in the heart dataset), employing a prototype optimization algorithm with 10 iterations. To address the common class imbalance problem in medical images, this invention specifically designed a background class weight (0.1) for optimization and supports joint analysis of up to 3 adjacent slices. The evaluation phase primarily focused on the segmentation performance of abdominal organs, with a model snapshot saved every 1000 iterations. The entire experimental process was managed using the Sacred framework to ensure reproducibility.
[0108] 3. Evaluation Metric: This paper adopts the Sorensen-Dice similarity coefficient, widely used in the FSMIS task, as the evaluation metric. This metric effectively measures the spatial overlap between the model's predicted segmentation results and the ground truth annotations. Its calculation formula is as follows:
[0109] ;
[0110] Here, A and B represent two sets, and the Dice score, ranging from 0 to 1, measures the similarity between the predicted and actual labels. A value closer to 1 indicates better model performance, and vice versa. It is worth noting that the average Dice score obtained from 5x cross-validation is the final result of the experiments in this invention.
[0111] 4. Experimental setup: This invention uses two setups, setup 1 and setup 2, to challenge the proposed method in order to improve its ability to generalize to new data.
[0112] In setting 1, test category objects may appear as background in the training dataset, but they are unlabeled. Since this invention uses a self-supervised supervoxel segmentation method to process the complete image, the obtained supervoxels are similar in shape and size to the test category objects. This similarity may lead to the implicit learning of test category object features during training, making these categories not entirely "unseen" by the algorithm.
[0113] In setting 2, novel categories appear neither in the foreground nor in the background of the training images, making all novel categories completely unknown to the model during the testing phase. For CT and MRI datasets, since the left and right kidneys typically appear in the same slice, this invention groups them together (lower abdomen); similarly, the liver and spleen are grouped into another (upper abdomen). However, this setting is not suitable for CMR datasets because it is impossible to completely exclude test categories from a single slice.
[0114] 5. Comparison of Results with State-of-the-Art (SOTA) Methods: To fully verify the effectiveness and advancement of the proposed method, this invention conducts a systematic performance comparison analysis with current mainstream FSMIS models, including PA-Net, SE-Net, SSL-ALPNet, ADNet, AAS-DCL, SR&CL, CRAPNet, Q-Net, CAT-Net, RPT, GMRD, and PAMI. Table 1 shows the performance evaluation results of the current mainstream models on the CHAOS and SABS datasets, respectively, which presents the average Dice score of the five cross-validation folds of the four organs (Spleen, Liver, LK, and RK) and the average Dice score of these four organs under two experimental settings (Set 1 and Set 2).
[0115] Table 1. Quantitative comparison of the mean Dice values of the CHAOS and SABS datasets under settings 1 and 2.
[0116]
[0117]
[0118] The best result is marked in bold, and the second best result is marked with an underline.
[0119] As shown in Table 1, the method of the present invention exhibits significant performance advantages in both scenario 1 and scenario 2.
[0120] Experimental results on the CHAOS dataset demonstrate that the proposed model exhibits significant performance advantages in both settings. In setting 1, the proposed method ranks first with an average Dice score of 84.79%, improving upon the second-best method, GMRD, by 1.89%. The improvement in segmentation performance for the spleen and left kidney (LK) is particularly significant, with Dice scores increasing by 4.38% and 2.17%, respectively. In the more challenging setting 2, the proposed method achieves the best results on all four target organs (spleen, liver, left kidney, and right kidney) segmentation tasks, with corresponding Dice scores of 78.17%, 82.73%, 82.23%, and 87.86%, respectively, outperforming the second-best methods for each organ by 2.37%, 1.64%, 3.46%, and 1.13%. Ultimately, the proposed method significantly surpasses the GMRD method by 3.05% with an average Dice score of 82.75%, demonstrating a comprehensive advantage in segmentation performance.
[0121] Experimental results on the SABS dataset demonstrate that the proposed model achieves significant performance improvements in both settings. In setting 1, the proposed method achieves a new state-of-the-art average Dice score of 79.88%, a 1.36% improvement over the previous best method, GMRD. Its performance is particularly outstanding in the spleen segmentation task, where it achieves a 4.78% improvement over the existing best. In the more challenging setting 2, the proposed method maintains leading or near-optimal segmentation performance across all four target organs, achieving a significant 7.36% improvement in the spleen segmentation task. Ultimately, it establishes a new performance benchmark with an average Dice score 1.25% higher than the GMRD model.
[0122] Table 2. Quantitative comparison of the mean Dice of the CMR dataset
[0123]
[0124] The best result is marked in bold, and the second best result is marked with an underline.
[0125] Experimental results demonstrate that the proposed method exhibits superior performance on the CMR dataset, surpassing all existing state-of-the-art (SOTA) methods with an average dice score of 79.14%. Particularly noteworthy is the significant performance improvement achieved in the left ventricular blood pool (LV-BP) segmentation task, outperforming the previous best Q-Net model by 1.18%. Overall, the proposed method further improves the average dice score by 0.03 percentage points.
[0126] Figure 6This presentation compares the DSC scores of 13 segmentation methods across five datasets (CHAOS-1, SABS-1, CHAOS-2, SABS-2, and CMR). Box plots use color to distinguish different methods, with the boxes displaying the interquartile range (IQR) and the median marked by the midline. This provides a clearer and more intuitive view of how EML maintains the highest performance across all datasets.
[0127] To intuitively evaluate the image segmentation performance of this method, this invention... Figure 7 , Figure 8 and Figure 9 This paper presents a visual comparison of our proposed method with mainstream baseline methods, including semi-supervised learning adaptive local networks (SSL-ALPNet), anomaly detection heuristic networks (ADNet), query information networks (QNet), prototype Transformer networks for region augmentation (RPT), generation of multiple representative descriptors (GMRD), and intelligent medical image segmentation techniques (PAIM), on three publicly available medical image datasets: CHAOS, SABS, and CMR. Experimental results demonstrate that our proposed method exhibits superior performance in multi-organ segmentation tasks. On the CHAOS dataset, our segmentation results for the spleen, left kidney, right kidney, and liver closely match the ground truth annotations. The segmentation accuracy for the spleen and left kidney is significantly better than other comparative models, fully demonstrating our outstanding ability in organ edge detail recognition and contour refinement. On the SABS dataset, our method achieves a qualitative improvement in spleen and left kidney segmentation tasks. Particularly in spleen segmentation, our method can accurately reconstruct the complete anatomical structure of the organ, resulting in clearer and sharper organ boundaries and more complete preservation of tissue texture features. These visual comparisons not only validate the effectiveness of this method in medical image segmentation but also highlight its strong generalization ability under limited labeled data conditions, demonstrating its enormous application potential as a clinical intelligent auxiliary diagnostic system. On the CMR dataset, the proposed method significantly improves the segmentation performance of LV-BP, exhibiting superior ability to segment clear boundaries. On the CMR dataset, this method achieves significant improvements in the left ventricular blood pool (LV-BP) segmentation task, with segmentation results not only highly consistent with the gold standard but also accurately capturing subtle boundary features of cardiac structures. This result further validates the clinical application value of this method in cardiac MRI image analysis.
[0128] In this invention, the controlled variable method is used to conduct ablation studies on relevant components under setting 1 in order to evaluate the effectiveness of each part of the model of this invention.
[0129] Table 3 Ablation studies of the impact of each component (measured by Dice score)
[0130]
[0131] Table 4 Ablation studies on the impact of intra-class and inter-class losses (measured by Dice score)
[0132]
[0133] Table 5 Ablation studies using different feature sizes (measured by Dice score %)
[0134]
[0135] Effects of Each Component: To verify the effectiveness of each component in the model of this invention, their positive contributions are shown in Table 3. The method of this invention significantly improves its performance on top of the baseline model (PANet), achieving an average dice score of 82.66%. This improvement is attributed to the Local Attention Fusion Prototype Generator (LAF), which captures the boundary regions of features, focusing on the foreground. Furthermore, the two-stage prototype optimization module facilitates the fusion and refinement of prototypes, further improving the model's performance by 2.13%.
[0136] Impact of Edge-Aware Loss: To verify the effectiveness of edge-aware loss, this invention conducted systematic ablation experiments as shown in Table 4. Experimental results show that introducing edge-aware loss can significantly improve the model's performance. This invention uses the Sobel operator to generate edge ground truth maps, predicts edge responses through a lightweight network, and combines it with weighted cross-entropy loss for joint optimization. This significantly enhances the model's ability to extract boundary features, enabling it to more accurately focus on semantically rich edge regions.
[0137] Selection of Encoded Input Features: To verify the impact of the size of the encoded input features on the model, an ablation experiment was conducted on the input feature sizes shown in Table 5. Using both 32×32 and 64×64 feature images as input features simultaneously yielded superior results compared to using either 32×32 or 64×64 feature images alone, achieving more comprehensive multi-scale information. Experimental results demonstrate that fusing encoded features of different scales significantly improves the model's representational ability. 32×32 features help capture global contextual information, while 64×64 features retain finer local details. This multi-scale feature fusion strategy achieves a better balance between semantic understanding and detail recovery, resulting in superior performance.
[0138] like Figure 10 As shown, after 20,000 iterations, the training loss and average loss of the model of this invention gradually flattened out, and no obvious overfitting was observed. The model is approaching convergence, which proves that the EML of this invention exhibits excellent convergence and can generalize well to the trained data.
[0139] This paper proposes an edge-aware multi-prototype learning framework (EML) for few-shot medical image segmentation. Through an innovative Local Attention Fusion Prototype Generator (LAF) and a two-stage Prototype Optimization Network (DPO), combined with an edge-aware loss function and a multi-scale feature fusion strategy, the model's segmentation performance is significantly improved under data-scarce conditions. Experiments on three public datasets—CHAOS, SABS, and CMR—show that this proposed method outperforms state-of-the-art methods in two different experimental settings. Particularly in the more challenging scenario of completely unseen categories, the average Dice score is improved by 1.25%–3.05% compared to the second-best method, validating its advantages in boundary localization and detail preservation. This research provides an effective solution for few-shot medical image segmentation and has significant application value in clinical auxiliary diagnosis.
[0140] The embodiments of the present invention have been described in detail above with reference to the examples. However, the present invention is not limited to the above embodiments. For those skilled in the art, after learning the contents described in the present invention, several equivalent changes and substitutions can be made without departing from the principle of the present invention. These equivalent changes and substitutions should also be considered to fall within the protection scope of the present invention.
Claims
1. A few-shot medical image segmentation method based on edge-aware multi-prototype learning, characterized in that, The edge-aware multi-prototype learning model comprises five key modules: a feature encoder, a local attention fusion prototype generator, a two-stage prototype optimization network, prototype prediction, and loss calculation; the method includes the following steps: S1. Input the support image and query image into the feature encoder to extract support feature maps and query feature maps of different sizes; S2. The support feature map is adjusted to the same size as the mask using bilinear interpolation algorithm, and then element-wise multiplication is performed. Next, the pixels in the target area are summed by mask weights and normalized using mask normalization factor to generate the support foreground prototype. At the same time, the mask image is processed by dynamic erosion operation to obtain the inner boundary mask. The feature map is then weighted by weighted average pooling to generate the inner boundary prototype. S3. Multiply the interpolated feature map with the mask, remove zero values and sample an equal number of background points to fill, reshape the feature vector, and perform weighted optimization. Through the cross-weighting mechanism, capture the local and global correlations and generate a multi-foreground local prototype. S4. Perform multi-scale transformation on the support feature map and its corresponding binary mask to extract and optimize local and global information to obtain a multi-scale prototype; then fuse the support foreground prototype, inner boundary prototype, multi-foreground local prototype and multi-scale prototype to obtain a multi-prototype foreground prototype. S5. The two-stage prototype optimization network is used to dynamically calculate and weight the multi-prototype foreground prototypes and perform automatic calibration to obtain the weighted aggregated multi-prototype foreground prototypes. S6. Calculate the negative cosine similarity of the query feature map, and input the negative cosine similarity of the query feature map and the weighted aggregated multi-prototype foreground prototype into the prototype prediction module to make predictions by balancing semantics and details. Finally, perform collaborative optimization through the loss calculation module.
2. The few-shot medical image segmentation method based on edge-aware multi-prototype learning according to claim 1, characterized in that, Step S2 is as follows: The spatial dimension of the input support feature map is adjusted by a bilinear interpolation algorithm to keep it consistent with the size of the mask matrix; then the obtained interpolated feature map is multiplied element-wise with the mask to realize the weighted summation calculation of each region. A mask-weighted summation operation is performed on all pixels within the target region, and then normalized using a normalization factor of the mask matrix to ultimately generate a foreground prototype. The mathematical expression is: ; in This represents the value of the interpolated feature map at pixel location (h, w). This represents the value of the mask at pixel position (h, w); Subsequently, morphological erosion is applied to shrink the target region in the mask image, eliminating edge noise and small false targets, and generating an eroded mask. ; The foreground inner boundary mask is obtained by calculating the difference between the original mask and the eroded mask. ; The mathematical expression is: ; ; This represents the mask after two erosions, where (a,b) and (a′,b′) are the local coordinates of the structuring element K in each erosion. During the first erosion, the minimum value within the neighborhood of the structuring element K is calculated for each pixel location; that is, the minimum value within the neighborhood of the structuring element K is calculated. The first erosion shrinks the foreground region; the second erosion repeats the calculation based on the first result, refining the target region, and finally obtaining... ; After generating the interpolated feature map and inner boundary mask, the norm of each pixel in the interpolated feature map is calculated as a weight coefficient, and the interference of high-intensity feature points is suppressed by the feature norm weighting strategy. Then, by combining the inner boundary mask and weight coefficients, a weighted average pooling operation is performed on the interpolated feature map to generate the inner boundary prototype. ; indicates as: ; ; Where ⊙ represents element-wise multiplication. Weights controlled by the feature norm, This is the inner boundary mask.
3. The few-shot medical image segmentation method based on edge-aware multi-prototype learning according to claim 2, characterized in that, Step S3 is as follows: First, the interpolated feature map F is multiplied element-wise with the mask M to extract foreground region features, thus obtaining the supporting foreground features. : ; Then, a background feature map is initialized, and a masking operation is used to ensure that the background region does not interfere with the extraction of foreground features. Next, a set of feature points of the foreground region is extracted from the interpolated feature map, and feature points with a mask value of 1 are selected, represented as: ; in For the set of foreground feature points, only select A pixel with a value of 1 The spatial extent of the input feature map; To maintain feature diversity in the foreground region, a specified number of feature vectors are randomly selected from the foreground region; mathematically, this can be expressed as: ; in For the final randomly selected set of foreground features, Foreground feature set This indicates the number of foreground features that need to be sampled. This indicates the random selection of a required number of non-repeating features from a sequence; Finally, the zero-value regions of the foreground features are filled with the sampled feature points to obtain the final foreground features; mathematically, this can be represented as: ; If and only if the background point is 0, these background point positions are replaced with randomly sampled prototype feature points. Supporting foreground features for logical judgments Is it 0? We introduce spatial-channel cross-attention, which optimizes spatial and channel features by simultaneously weighting spatial and channel dimensions. Spatial-channel cross-attention comprises three parts: channel weighted pooling, spatial weighted pooling, and cross-attention. Channel-weighted pooling dynamically optimizes the importance representation of each channel by adaptively learning the weight distribution along the channel dimension of the feature map. A 1×1 convolution decomposes the supporting foreground features into query features and value features. The query features are then subjected to global average pooling to extract channel-level global statistics, followed by a tensor reshaping operation and matrix multiplication with the value features to construct a correlation matrix between channels. This correlation matrix is then normalized using SoftMax and activated by Sigmoid to generate a channel-weighted mask. Finally, feature enhancement is achieved by multiplying the mask with the original supporting foreground features channel by channel. Channel attention, mathematically represented as follows: ; and They are Convolutional layer , and It involves reshaping three tensors. It is a global average pooling operation. , It is a SoftMax operation. " "This is the matrix dot product operation." It is the Sigmoid operation. , It is element-wise multiplication; Spatial weighted pooling dynamically learns the weight distribution across the spatial dimension of the support foreground features, adaptively optimizing the importance representation of features at each location. First, a 1×1 convolution decomposes the input support foreground features into query features and value features. After tensor reshaping, the SoftMax-normalized query and value features are subjected to matrix operations to generate spatial context features. Then, a 1×1 convolution and Sigmoid activation are used to generate a spatial attention mask. Finally, feature enhancement is achieved through element-wise multiplication of the mask with the original support foreground features. The mathematical representation is as follows: ; yes Convolutional layer; Indicates spatial attention; The cross-attention module enhances feature representation by simultaneously capturing long-range dependencies of supporting foreground features in both height and width directions. First, three 1×1 convolutions are used to generate query, key, and value feature projections, respectively. Then, the feature map is decomposed into two orthogonal directions (height and width) for processing through dimensionality transformation. In the height branch, after feature reshaping, attention weights along the height direction are calculated using batch matrix multiplication, and an INF mask is introduced to ensure that only valid regions are considered. The width branch employs a symmetrical processing approach. After concatenation and SoftMax normalization of the attention maps from both directions, matrix multiplication is performed with the corresponding value features to achieve feature aggregation. Finally, the results from both directions are weighted and fused, and a learnable scaling factor is introduced. Its mathematical representation is as follows: ; ; ; ; in and This represents the key features and value features being convolved with X using a 1×1 method, with processing applied to the height direction. and This represents the key features and value features being convolved with X using a 1×1 method, with the width direction being processed. and This indicates that a 1×1 convolution is performed on the query features, and the dimensions are transformed after processing in the height and width directions. The INF mask generates a matrix with negative infinity on the diagonal to avoid self-attention. " "This is the matrix dot product operation." It decouples and separates the characteristics of sotfmax into height and width. It is a scaling factor that can dynamically control attention-enhanced features; This indicates a high level of attention. Note the width. Indicates cross attention; Features of attention across spatial-channel intersections; To comprehensively represent the feature information of each category, a feature descriptor generation framework based on a multilayer perceptron is adopted. First, a deep feedforward neural network with multiple hidden layers is constructed. Through an end-to-end training process, the optimal feature transformation is automatically learned to generate multiple semantically representative feature descriptors for each category. (Multi-foreground local prototypes are then used.) The mathematical representation is as follows: ; Flatten the input features into , for , where in represents the input features H and W. for , LeakyReLU(x) is equivalent to max(0.01x,x). It is a regularization layer with a probability of 0.
3. for , for , for , for , out is the number of custom output prototypes.
4. The few-shot medical image segmentation method based on edge-aware multi-prototype learning according to claim 3, characterized in that: For the final randomly selected set of foreground features If the number of available foreground pixels is insufficient, an oversampling strategy is employed, which supplements feature points by repeatedly selecting a subset of samples; randomly sampled foreground features are filled into the original feature map, and replacements are made at positions where the background mask is zero; mathematically, this can be represented as: ; in This indicates that n original foreground feature stitching operations were performed, where n represents... , This indicates that the data is rounded down, and some foreground features are randomly selected to supplement the insufficient number of features.
5. The few-shot medical image segmentation method based on edge-aware multi-prototype learning according to claim 3, characterized in that, The method for generating the multi-scale prototype in step S4 is as follows: First, the input features undergo spatial-channel cross-attention. The binary mask and its corresponding binary representation are subjected to multi-scale transformation, processed in parallel across three typical scale spaces. For each scale space, a mask-guided feature aggregation strategy is used to generate discriminative prototype vectors, where fine-grained scales focus on local structural features and coarse-grained scales capture global contextual information. By adaptively fusing multi-scale prototype representations, the final feature representation system can achieve synergistic optimization of local detail preservation and global semantic understanding. Its mathematical representation is as follows: ; ; Where s represents the three scales , It is a prototype generated at three scales. For multi-scale prototypes, For binary masking, , Used to prevent division by zero errors. This represents a stacked prototype feature across three scales; This represents the height of the feature map at scale s. This represents the width of the feature map at scale s. Indicates the index in the height direction of the feature map. Indicates the index in the width direction of the feature map.
6. The few-shot medical image segmentation method based on edge-aware multi-prototype learning according to claim 5, characterized in that, The method for fusing the multiple prototype foreground prototype P in step S4 is as follows: Fusion Support Prototype Inner boundary prototype Multi-foreground local prototype and multi-scale prototypes It integrates diverse feature information, including foreground, boundaries, local details, and features at different scales, enhancing the model's focus on boundary regions; its mathematical representation is as follows: ; This indicates a serial operation.
7. The few-shot medical image segmentation method based on edge-aware multi-prototype learning according to claim 6, characterized in that, The specific steps of step S5 are as follows: Using a multi-prototype foreground prototype P as input, multiple prototypes under the same category are adaptively fused through learnable attention weights to generate a discriminative and robust category representation. A lightweight attention network maps dimensional features to scalar attention weights, and then Softmax normalization satisfies convex combination constraints on the weights. Finally, the weighted sum of all prototypes within the same category is output, i.e., the prototype set. The mathematical representation is as follows: ; in , representing the learnable weight matrix, maps multidimensional features to one dimension, and K is the number of supporting sample prototypes for the current class. , is the bias term. Transpose of the weight matrix It is a SoftMax operation. This represents element-wise multiplication; The prototype calibration module receives a set of prototypes from the pre-attention aggregation layer. First, the arithmetic mean of the prototypes of each category is calculated as the initial class center. Then, the cosine similarity is used to measure the directional consistency between each prototype and the class center, and the similarity is converted into normalized attention weights using the Softmax function. Based on attention weights, the prototypes are weighted and aggregated; its mathematical representation is as follows: ; ; in , Let K represent the prototype feature of the i-th supporting sample, and K represent the number of supporting sample prototypes for the current category. This is the formula for calculating cosine similarity. It is an exponential function; Represents the i-th prototype Cosine similarity to the arithmetic mean class center of all prototypes in the current category; Indicates the original prototype The prototype obtained after weighted aggregation; and Let represent the cosine similarity between the k-th and j-th prototypes and the class center, respectively; This represents the prototype feature vector extracted by the network for the k-th supporting sample.
8. The few-shot medical image segmentation method based on edge-aware multi-prototype learning according to claim 1, characterized in that, The prediction process of the prototype prediction module in step S6 is as follows: First, a hierarchical prototype matching mechanism is constructed: for each query image's query feature map, the negative cosine similarity between it and the foreground prototypes of each category is calculated, and a learnable scaling factor and a category-specific threshold are introduced to adjust the similarity metric space; the adjusted distance is converted into probability values through the sigmoid function to form the initial segmentation response map; A multi-scale feature fusion strategy is adopted. After prototype matching is performed on feature maps extracted from different network layers, they are uniformly upsampled to the original image size through bilinear interpolation. Learnable weight coefficients are used to adaptively weight the prediction results of different layers. The final prediction result is represented by a binary classification output in the form of softmax, and the foreground probability and background probability are concatenated to form a two-channel output tensor. Its mathematical representation is as follows: ; ; in It's the Sigmoid function, which compresses the output to [0, 1]. This is the formula for calculating cosine similarity. It is a tensor concatenation function. These are learnable weights. It is bilinear interpolation upsampling; This represents the predicted probability value of the nth query sample belonging to a certain category, which is compressed to the [0,1] interval by the Sigmoid function; This represents the depth features of the nth query sample; The prototype feature representing a certain category is generated from the feature mean of similar samples in the support set or other aggregation methods; This represents the prediction result for all N samples. The final segmentation probability map is generated by stitching together weighted aggregation and upsampling. This represents the original prediction result for the nth sample.
9. The few-shot medical image segmentation method based on edge-aware multi-prototype learning according to claim 8, characterized in that, The optimization process of the loss calculation module in step S6 is as follows: First, class prototype features are extracted using the query set prediction mask. Then, back-validation is performed by calculating the cross-entropy loss between the support set prediction mask and the true label. Normalization is performed using the number of classes and the number of samples to ensure gradient balance across different training episodes. The mathematical representation is as follows: ; ; ; in, It is a predictive mask. It is to obtain query features. It supports image features. It's a similarity calculation. It is bilinear interpolation upsampling. This is the category threshold, where K is the number of categories and S is the number of shots. It is an indicator function. It is cross-entropy loss. It supports set prediction masks; Represents the prototype of feature of category c; This represents the set of all pixels belonging to category c, typically defined by a predicted mask or a ground truth mask. The logical sigmoid function compresses real number inputs into the (0,1) interval, and its output can be interpreted as a probability. This represents a learnable weight parameter used to weight feature similarity at different levels or locations; Indicates the category label currently being processed; This represents the prediction mask for class c of the s-th sample in the support set, obtained through prototype... With supporting features Similarity calculation is used to generate the result; This represents a loss function that measures the difference between the predicted mask and the true mask in the support set, promoting the alignment of features between the prototype and the support set. Simultaneously, an edge-aware loss function is designed to enhance boundary localization capabilities in few-sample semantic segmentation by explicitly modeling geometric structural features. The edge-aware loss function employs a multi-stage collaborative optimization strategy: firstly, an edge detection module is constructed based on the classic Sobel operator, and then horizontal convolutional kernels are used to perform edge detection. and vertical convolution kernel Calculate the spatial gradient of the support set mask to generate an edge ground map with pixel-level accuracy. Subsequently, a lightweight edge prediction network was designed, which replicates the foreground probability map output by the model to form a three-channel input, and predicts the edge response map. Finally, end-to-end optimization is performed using weighted binary cross-entropy loss; its mathematical representation is as follows: ; ; ; ; ; in Limit the value to the range [0, 1]. The predicted mask is copied three times for a specific channel and then stitched together along the channel dimension. It is a lightweight edge detection network. It is a binary cross-entropy loss function. This represents a function that performs a convolution operation on the input feature map; This represents the depth feature representation extracted from the query image; This indicates the detection of vertical edges; Indicates the detection of horizontal edges; The edge probability map represented by the network is generated by a lightweight edge detection network; This represents the binary cross-entropy loss, which measures the difference between the predicted edge and the actual edge. Represents the true edge label at position (i,j); This represents the probability value that the model predicts the position (i,j) to be on the edge.
Citation Information
Patent Citations
Medical image segmentation method and device, computer equipment and storage medium
CN114419020A
Improved few-sample medical image target region segmentation method
CN119722604A