Medical image segmentation method based on uncertainty estimation and multistage distillation
Through uncertainty estimation and multi-level knowledge distillation methods, the uncertainty of the teacher model is quantified, the high confidence pseudo-label is screened, and the hybrid attention mechanism is embedded in the 3D Unet network, which solves the problem of scarcity of labeled data in medical image segmentation and improves segmentation accuracy and boundary capture capabilities.
Patent Information
- Application Number
- CN202510728521.7
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-06-03
- Publication Date
- 2025-07-04
- Estimated Expiration
- 2045-06-03
AI Technical Summary
The scarcity of high-quality labeled data in medical image segmentation leads to limited improvement in performance of deep learning algorithms, especially in complex morphological structures and boundary region segmentation.
The medical image segmentation method based on uncertainty estimation and multi-stage distillation is adopted to quantify the predictive uncertainty of the teacher model through the uncertainty estimation module, screen high confidence pseudo-labels, and transfer the knowledge of the teacher model to the student model in stages through the multi-stage knowledge distillation module. At the same time, a hybrid attention module is embedded in the 3D Unet network, combining the channel and spatial attention mechanism, and using a composite loss function to optimize the training process of the student model.
Under the condition of limited labeled data, the accuracy and boundary capture capabilities of medical image segmentation are significantly improved, and more precise segmentation of lesion structure and boundaries are achieved, solving the problem of scarcity of labeled data.
Smart Images

Figure CN120259284A_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to medical image segmentation, and specifically to a medical image segmentation method based on uncertainty estimation and multi-level distillation. Background Art
[0002] Medical image segmentation has been widely used in clinical treatment, which can assist doctors in giving high-quality diagnostic opinions by detecting lesions or organs. Among them, deep learning has developed rapidly in recent years, and through deep learning technology, the performance of medical image segmentation has been significantly improved. Currently, the mainstream deep learning segmentation networks mainly include convolutional neural networks, Transformer models based on attention mechanisms, and hybrid models that combine the advantages of multiple architectures. CNN extracts local features of medical images through hierarchical convolutional operations, with strong spatial invariance and feature sharing capabilities, making it perform excellently in medical image segmentation tasks, especially suitable for structural analysis sensitive to local details. In contrast, Transformer relies on self-attention mechanisms to model long-range dependencies and can effectively capture global information, thus showing stronger modeling capabilities in segmentation tasks of complex morphological structures and improving the segmentation accuracy of boundary regions and small targets. However, regardless of the network model used, its performance improvement usually depends on large-scale labeled data, which is a major challenge for medical image segmentation. The annotation process of medical images not only requires highly professional medical knowledge but also takes a lot of time to ensure the accuracy of segmentation labels, making the acquisition cost of high-quality labeled data extremely expensive. To solve this problem, many studies in recent years have been devoted to exploring how to improve segmentation performance under the condition of limited labeled data.
[0003] Semi-supervised learning methods can improve the learning performance of the model by effectively using unlabeled data in the case of limited labeled data and a large amount of unlabeled data. Specifically, semi-supervised learning combines a small number of labeled samples and a large number of unlabeled samples, and through methods such as generating pseudo-labels, contrast learning, or consistency training, enables the model to extract valuable features from unlabeled data. Currently, semi-supervised learning (SSL) methods can be roughly divided into two categories. The first category is the pseudo-label-based method, which generates pseudo-labels through the model's prediction of unlabeled data and uses these pseudo-labels together with the labeled data for training. The other category of methods is the consistency regularization-based method, which learns the consistency of the inference results of the same unlabeled image under different perturbation conditions by using two deep convolutional neural networks, thereby better utilizing the information of unlabeled data.
[0004] However, due to the interference of imaging devices, individual differences among patients, and errors in manual annotation, etc., uncertainties widely exist in the medical datasets used for deep learning algorithm training. The existence of uncertainties severely limits the improvement of the performance of image segmentation algorithms. Summary of the Invention
[0005] Aiming at the above deficiencies in the prior art, a medical image segmentation method based on uncertainty estimation and multi-level distillation provided by the present invention solves the problem of scarce high-quality labeled data in medical image segmentation.
[0006] In order to achieve the above invention purpose, the technical solution adopted by the present invention is: A medical image segmentation method based on uncertainty estimation and multi-level distillation, comprising the following steps:
[0007] S1: Quantify the prediction uncertainty of the teacher model for the input medical image through the uncertainty estimation module to generate uncertainty weights;
[0008] S2: Screen high-confidence pseudo-labels based on the uncertainty weights, and transfer the knowledge of the teacher model to the student model in stages through the multi-level knowledge distillation module;
[0009] S3: Embed a hybrid attention module in the teacher model and the student model, and combine the channel attention mechanism and the spatial attention mechanism to enhance the feature capture of the lesion structure and boundary;
[0010] S4: Adopt a composite loss function that combines cross-entropy loss, Dice loss, and uncertainty-based distillation loss to optimize the training process of the student model;
[0011] S5: Apply the trained network to the medical image segmentation scenario to complete the medical image segmentation based on uncertainty estimation and multi-level distillation.
[0012] Further, the uncertainty estimation module in S1 is implemented by the Monte Carlo Dropout method, specifically including:
[0013] A1: Enable the Dropout layer in the inference stage of the teacher model to perform multiple forward predictions on the same input;
[0014] A2: Calculate the prediction entropy based on the probability distribution of the multiple prediction results as a quantization index of the segmentation uncertainty:
[0015]
[0016]
[0017] Wherein, is the prediction entropy, is the number of classes in the segmentation task, is the average softmax probability of T samples from the teacher model, is the prediction value of the class in the t-th forward pass of the teacher model;
[0018] A3: Dynamically adjust the screening threshold of pseudo-labels based on prediction entropy to filter the prediction results in high-uncertainty regions.
[0019] Furthermore, the implementation of the multi-level knowledge distillation module in S2 includes:
[0020] B1: Update the weights between the teacher model and the student model through exponential moving average to generate an intermediate model;
[0021] B2: Use the intermediate model as a new teacher model for secondary distillation to gradually transmit multi-level semantic features;
[0022] B3: During the distillation process, use the uncertainty map based on prediction entropy to optimize the mean square error loss and Dice loss between the predictions of the student model and the teacher model. At the same time, calculate the cross-entropy loss and Dice loss between the student model and the binarized labels to constrain the prediction consistency between the student model and the teacher model.
[0023] Furthermore, the network architectures of both the teacher model and the student model in S3 are 3D Unet, and a hybrid attention module is embedded in the bridging part of the 3D Unet. The hybrid attention module includes the following operations:
[0024] C1: Perform global average pooling and global max pooling on the input feature map respectively to generate two channel descriptors:
[0025]
[0026]
[0027] where, and are two channel descriptors, is the number of channels, and are the height and width of the feature map respectively, and respectively represent the th row and the th column elements of the feature map;
[0028] C2: Interact with the two channel descriptors through a multi-layer perceptron to generate channel attention weights:
[0029]
[0030] where, is the channel attention weight, is a multi-layer perceptron;
[0031] C3: Multiply the channel attention weight element-wise with the input feature map to obtain a channel-enhanced feature map, and perform spatial dimensional pooling on the channel-enhanced feature map, and generate a spatial attention weight through a convolution operation. The channel-enhanced feature map is expressed as:
[0032]
[0033] where, is the channel-enhanced feature map.
[0034] Furthermore, the composite loss function in S4 is:
[0035]
[0036] where, is the cross-entropy loss, is the Dice loss, is the distillation loss of uncertainty, and is the balance coefficient.
[0037] Furthermore, the cross-entropy loss is:
[0038]
[0039] where, is the total number of samples, is the prediction of the student model for the current sample and is the actual label of the current sample
[0040] Furthermore, the Dice loss is:
[0041]
[0042] where, is the intersection symbol.
[0043] Furthermore, the distillation loss of uncertainty is:
[0044]
[0045] where, is a Gaussian warm-up function related to time, is the current iteration number reached during the training process, is the prediction of the teacher model for the current sample The segmentation uncertainty, is the threshold of the most certain samples of the teacher model, is a hyperparameter, is the prediction of the student model for the current sample and is the prediction of the teacher model for the current sample and the mean squared error loss between them, is the prediction of the student model for the current sample and is the prediction of the teacher model for the current sample and the Dice loss between them, represents a selection operation. If holds, the value is 1; otherwise, the value is 0.
[0046]
[0047]
[0048] Among them, is a single voxel, is the total number of voxels in the current sample, is the prediction value of the student model for the th sample and the th voxel, is the prediction value of the teacher model for the th sample and the th voxel.
[0049] The beneficial effects of the present invention are as follows:
[0050] Aiming at the problem of scarce high-quality labeled data in medical image segmentation, the present invention proposes a segmentation network (UGMLDS-Net) based on uncertainty estimation and multi-level knowledge distillation. First, by introducing an uncertainty estimation module, the method quantifies the prediction uncertainty of the teacher model using the Monte Carlo Dropout method, thereby dynamically screening high-confidence pseudo-labels and reducing the negative impact of noisy pseudo-labels on model training. In addition, UGMLDS-Net adopts a multi-level knowledge distillation mechanism. Through staged knowledge transfer, the student model can gradually learn multi-level and multi-stage information in the teacher model, so as to more accurately capture the details and boundaries of lesions. At the same time, in order to enhance the model's ability to capture complex lesion structures and fuzzy boundaries, the present invention inserts a hybrid attention module into the 3D U-Net network. This module combines channel attention and spatial attention mechanisms, significantly improving the segmentation accuracy.
[0051] The proposed UGMLDS-Net in the present invention realizes more fine-grained knowledge transfer by combining uncertainty estimation and multi-level knowledge distillation mechanisms, enhances the learning ability of the student model, and efficiently utilizes unlabeled data under the condition of limited labeled data, providing an effective solution to the problem of scarce labeled data in medical image segmentation.
[0052] The present invention proposes an uncertainty estimation module to guide the student model to perform more refined learning in high-uncertainty regions, optimize the distillation process, and enhance the model's discrimination ability for difficult-to-segment regions.
[0053] The present invention proposes an uncertainty-based loss function to dynamically filter the knowledge of the teacher model during the multi-level distillation process, adjust the knowledge transfer through uncertainty weights, enable the student model to preferentially learn the information in high-confidence regions, and perform progressive knowledge optimization under the multi-level distillation framework, thereby improving the network performance. Brief Description of the Drawings
[0054] Figure 1 It is a flowchart of a medical image segmentation method based on uncertainty estimation and multi-level distillation according to the present invention.
[0055] Figure 2 It is a structural diagram of a medical image segmentation network based on uncertainty estimation and multi-level distillation according to the present invention.
[0056] Figure 3 It is a visual comparison diagram of segmentation using 10% labeled data on the Brats 19 dataset.
[0057] Figure 4 It is a visual comparison diagram of segmentation using 20% labeled data on the LA dataset. Detailed Embodiments
[0058] The following further describes the present invention with reference to the accompanying drawings and specific embodiments.
[0059] As Figure 1 shown, a medical image segmentation method based on uncertainty estimation and multi-level distillation includes the following steps:
[0060] S1: Quantify the prediction uncertainty of the teacher model for the input medical image through the uncertainty estimation module to generate uncertainty weights;
[0061] S2: Screen high-confidence pseudo-labels based on the uncertainty weights, and transfer the knowledge of the teacher model to the student model in stages through the multi-level knowledge distillation module;
[0062] S3: Embed a hybrid attention module in the teacher model and the student model, and combine the channel attention mechanism and the spatial attention mechanism to enhance the feature capture of the lesion structure and boundary;
[0063] S4: Adopt a composite loss function that combines cross-entropy loss, Dice loss, and uncertainty-based distillation loss to optimize the training process of the student model;
[0064] S5: Apply the trained network to the medical image segmentation scenario to complete medical image segmentation based on uncertainty estimation and multi-level distillation.
[0065] Currently, most deep learning-based medical image segmentation algorithms are data-driven, and the performance of the algorithms highly depends on the quality of the dataset. Tumors usually exhibit significant morphological and histological heterogeneity. There may be large differences in different structures within the tumor (such as necrotic regions, edematous regions, and enhancing cores), as well as the shape, size, and location of the tumor in different patients. Moreover, affected by differences in the resolution, contrast, and noise level of different imaging devices, problems such as modality loss or image artifacts may occur in the collected MRI images. In addition, when constructing a dataset for deep learning model training, medical images need to be annotated by professional doctors. However, due to the subjectivity of experts' judgment on tumor boundaries, the annotation results may be inconsistent due to differences among experts. In summary, the uncertainty introduced into the tumor segmentation dataset due to patient individual differences, tumor imaging noise, and experts' annotation differences poses a huge challenge to accurate medical image segmentation. Therefore, the present invention proposes an uncertainty quantification method, which quantifies the uncertainty in MRI imaging by using a teacher model and guides the student model to perform more refined segmentation of the tumor region based on this uncertainty information.
[0066] The Bayesian method quantifies uncertainty by specifying the prior distribution of the neural network parameters and calculating the posterior distribution of the parameters given the training data, giving the dataset and representing the random output of the Bayesian network as with the model likelihood represented as and Bayesian inference used to calculate the posterior For the weights w, given the prior capturing a set of reasonable model parameters:
[0067]
[0068] Given a new test sample the predictive distribution can be obtained by integrating over the model parameters according to the posterior distribution of the parameters:
[0069]
[0070] For MRI images, through the inference of the above formula, the posterior weighted average value of the model can be generated for each pixel in the image, which is also called the Bayesian model average (BMA). However, the difficulty of the above prediction distribution is that the posterior of the parameters is difficult to calculate analytically:
[0071]
[0072] The above formula needs to be integrated over all possible weight spaces to marginalize, but deep neural networks usually have millions of parameters, and it is impossible to integrate over all weight parameters. This has led to approximate Bayesian methods for approximating the posterior.
[0073] Since the MCDropout method only needs to add a dropout layer to the existing network model without changing its architecture, it has become a commonly used method for uncertainty quantification. Therefore, in this embodiment, the MCDropout method is used to quantify the uncertainty information in remote sensing images, and some neural network units are randomly discarded by the dropout layer during the training process to avoid overfitting.
[0074] Most existing image segmentation algorithms give a certain probability score. Using the MCDropout method, when training the model, we make the dropout layer perform multiple forward predictions on the same remote sensing image during testing, which is equivalent to obtaining samples from the posterior. In the experiment, we do this 50 times to obtain 50 Monte Carlo samples that approximate the predicted probability distribution.
[0075] Based on the samples drawn using the Monte Carlo method, the prediction entropy can be quantified to approximate the segmentation uncertainty of the teacher model.
[0076] The uncertainty estimation module in S1 is implemented by the Monte Carlo Dropout method, which specifically includes:
[0077] A1: Enable the Dropout layer during the inference stage of the teacher model and perform multiple forward predictions on the same input;
[0078] A2: Calculate the prediction entropy based on the probability distribution of the multiple prediction results as a quantization index for segmentation uncertainty:
[0079]
[0080]
[0081] Among them, is the prediction entropy, is the number of classes in the segmentation task, is the average softmax probability of T samples from the teacher model, is the prediction value of the class in the t-th forward pass of the teacher model;
[0082] A3: Dynamically adjust the screening threshold of pseudo-labels based on prediction entropy to filter the prediction results in high-uncertainty regions.
[0083] For a given input, the overall segmentation uncertainty can be estimated from the samples obtained through multiple forward predictions of the teacher model . Guided by the uncertainty of the teacher model, filter out the relatively unreliable predictions considered by the teacher model with high uncertainty, and only retain the relatively certain predictions of the teacher model for the student model to learn. This can, to a certain extent, avoid the transmission of incorrect information from the teacher model during the process of the student model's refined segmentation of the lesion area, thereby improving the segmentation accuracy of the tumor area.
[0084] Knowledge distillation was first proposed as a network compression method. By having a lightweight student model learn the output of the teacher model, the student network can achieve better results than learning directly by itself. Through continuous research and development, various methods for learning the knowledge of the teacher model have emerged, such as multi-teacher distillation, self-distillation, and online distillation. Although the existing knowledge distillation mechanisms have achieved good results in medical image segmentation, there are still certain deficiencies in the semi-supervised scenario. The existing knowledge distillation methods usually rely only on one-time knowledge transfer, that is, the knowledge transfer from the teacher model to the student model is completed through a single distillation process. However, in complex medical image segmentation, it is difficult to make full use of the multi-level and multi-stage information contained in the teacher model in this way. It is difficult for the student model to fully master the multi-scale structure and fine-grained structure information of the lesion after a single distillation. Therefore, the present invention proposes a multi-level knowledge distillation mechanism based on the above uncertainty quantification module, as Figure 2 shown. First, train the upper yellow teacher model and the middle blue intermediate model, which is the first knowledge distillation. After training is completed, use the intermediate model as the teacher model and train it together with the bottom green student model. This is the second knowledge distillation. All three models use 3D UNet as the basic architecture, and the hybrid attention module described in C1 is embedded in the 3D UNet. The entire processing logic can be viewed in two parts. The first distillation is the simultaneous training of the upper and middle two-layer networks, and the second distillation is the simultaneous training of the middle and bottom two-layer networks. Figure 2 The black-and-white prediction of the output of the middle blue intermediate model in is during the first distillation, and its output during the first distillation serves as the output of the student model. During the second distillation, it serves as the teacher model. Therefore, the output at this time.
[0085] In multi-level knowledge distillation, the teacher model (TeacherModel) first undergoes knowledge distillation once to obtain an intermediate model (MediateModel), and then this intermediate model (MedicateModel) is used as the teacher model for another round of distillation to obtain the final student model (StudentModel). Specifically, the teacher model (TeacherModel), the intermediate model (MediateModel), and the student model (StudentModel) all adopt the same network architecture. During the training process of each knowledge distillation, we update the network model weights of the corresponding teacher model to the exponential moving average of the student model weights , that is . Where γ is the hyperparameter for weight decay. During the knowledge distillation process, the mean squared error loss and Dice loss between the predictions of the student model and the teacher model are optimized using the uncertainty map based on prediction entropy. At the same time, the cross-entropy loss and Dice loss between the student model and the binarized labels are calculated, enabling the student model to not only learn the information in the labels but also learn the more certain knowledge of the teacher model through the guidance of uncertainty. In addition, through multi-level knowledge distillation, the multi-level and multi-stage information contained in the teacher model can be fully utilized to better grasp the fine-grained information of the lesion and refine the segmentation accuracy of the lesion boundary.
[0086] The implementation of the multi-level knowledge distillation module in S2 includes:
[0087] B1: Update the weights between the teacher model and the student model through exponential moving average to generate an intermediate model;
[0088] B2: Use the intermediate model as the new teacher model for secondary distillation to gradually transmit multi-level semantic features;
[0089] B3: During the distillation process, use the uncertainty map based on prediction entropy to optimize the mean squared error loss and Dice loss between the predictions of the student model and the teacher model, and at the same time calculate the cross-entropy loss and Dice loss between the student model and the binarized labels to constrain the prediction consistency between the student model and the teacher model.
[0090] The network architectures of the teacher model and the student model in S3 are both 3D Unet, and a hybrid attention module is embedded in the bridging part of the 3D Unet. The hybrid attention module includes the following operations:
[0091] C1: Perform global average pooling and global max pooling on the input feature map respectively to generate two channel descriptors:
[0092]
[0093]
[0094] Among them, and are two channel descriptors, is the number of channels, and are the height and width of the feature map respectively, and respectively represent the th row (height) and the th column (width) elements of the feature map;
[0095] C2: Interact with the two channel descriptors through a multi-layer perceptron to generate channel attention weights:
[0096]
[0097] Among them, is the channel attention weight, is the multi-layer perceptron;
[0098] C3: Multiply the channel attention weight with the input feature map element by element to obtain the channel-enhanced feature map, and perform spatial dimensional pooling on the channel-enhanced feature map to generate spatial attention weights through a convolution operation. The channel-enhanced feature map is expressed as:
[0099]
[0100] Among them, is the channel-enhanced feature map.
[0101] In this embodiment, both the teacher-student network structures use the 3D Unet structure. We inserted a Convolutional Block Attention Module (CBAM) into the central bridging part of the 3D Unet to enhance its expressive ability. The CBAM module combines the channel attention mechanism and the spatial attention mechanism, which can better capture the key features of the image and improve the performance of the model. The spatial attention mechanism focuses on the spatial dimension in the image. By performing max-pooling and average-pooling operations on the input feature map, an importance map of the spatial region is generated. Then, through convolutional operations, the pooling results are fused to generate the spatial attention weights for each position. In this way, the model can focus on the key regions in the spatial dimension, especially the boundary regions. By inserting the CBAM module, it helps with multi-scale feature extraction. The channel attention mechanism can selectively strengthen important channels, while the spatial attention mechanism can focus on key regions, providing a more detailed feature representation. CBAM can be easily integrated into the existing network architecture without adding excessive computational overhead, while significantly improving the performance of the model. For complex structures and subtle differences in the image, CBAM can help the model effectively capture the key information, especially in the segmentation performance of the boundary and detail parts, achieving better results
[0102] The composite loss function in S4 is:
[0103]
[0104] where is the cross-entropy loss, is the Dice loss, is the distillation loss of uncertainty, and are balance coefficients.
[0105] During the training process of knowledge distillation, by calculating the cross-entropy loss and Dice loss between the predictions of the student model and the labels, and the mean squared error loss and Dice loss between the predictions of the student model and the teacher model, the student model can not only learn the lesion information in the labels but also learn the relatively correct and reliable knowledge in the teacher model. Through multiple knowledge distillations, the segmentation accuracy of the student model for the lesion region and the lesion boundary can be effectively improved. In this embodiment and are taken as 0.5 and 1 respectively.
[0106] The cross-entropy loss calculates the cross-entropy loss between the prediction of the student model and the label The cross-entropy loss is:
[0107]
[0108]
[0109] Among them, is the total number of samples, is the prediction of the student model for the current sample , is the current sample 's actual label.
[0110] In medical image segmentation, the cross-entropy loss mainly optimizes the model by comparing the prediction and label pixel by pixel, and is very suitable for processing scenarios with balanced class distributions. However, in medical image segmentation, there are often serious class imbalance problems, that is, the lesion or organ area only accounts for a very small proportion in the image. This inter-class imbalance will cause the cross-entropy loss to mainly optimize the background area, thus affecting the accurate segmentation of the foreground area. Therefore, the present invention uses the Dice loss function to assist in optimizing the model.
[0111] The Dice loss function evaluates the model performance by calculating the Dice coefficient between the student model prediction and the true label. The Dice loss function can avoid the dominant role of excessive background pixels in loss calculation and is suitable for the segmentation task of tiny lesions. By focusing on the global optimization goal of the overall matching degree between the model prediction and the true label, rather than the independent calculation pixel by pixel, the Dice loss function can improve the overall consistency of the segmentation result and make the model prediction more accurate and smooth. By introducing the Dice loss function on the basis of the cross-entropy loss function, the segmentation performance of the model for lesions can be effectively improved in scenarios where the proportion of lesions is small, and the quality of the overall lesion segmentation result can be improved.
[0112] The Dice loss is:
[0113]
[0114] Among them, is the intersection symbol.
[0115] In addition, the present invention also introduces an uncertainty loss to filter the knowledge of the teacher model during the multi-level knowledge distillation process, guiding the student model to gradually learn the correct knowledge in the teacher model during multiple learning processes. The proposed uncertainty loss is as follows:
[0116] The distilled loss of the uncertainty is:
[0117]
[0118] Among them, is the Gaussian warm-up function related to time, is the current iteration number reached during the training process, is the segmentation uncertainty of the teacher model for the current sample , is the threshold for the most certain samples of the teacher model, is a hyperparameter, is the prediction of the student model for the current sample and is the prediction of the teacher model for the current sample and the mean squared error loss between them, is the prediction of the student model for the current sample and is the prediction of the teacher model for the current sample and the Dice loss between them, represents a selection operation. If holds, the value is 1; otherwise, the value is 0.
[0119] In the above uncertainty loss function, , represents the set maximum number of iterations. In this embodiment, the same Gaussian ramp-up function is used during the experiment to increase the threshold H from to through the formula H=(0.75 + 0.25×λ(t))× . represents the maximum uncertainty value, which is set to ln(2) in the experiment. is set to 0.75 in the experiment;
[0120]
[0121]
[0122] Among them, is a single voxel, is the total number of voxels of the current sample, is the predicted value of the student model for the th sample and the th voxel, is the predicted value of the teacher model for the th sample and the th voxel.
[0123] The mean squared error loss can capture the fine-grained differences between the predictions of the teacher model and the student model, prompting the student model to be closer to the output of the teacher model. In the case of a severe imbalance between the lesion area and the background area, the Dice loss helps the student model better learn and optimize the teacher model's handling and segmentation of the lesion area boundary. During the training of the student model, the segmentation threshold H of the teacher model for the sample increases from to , enabling the student model to learn from relatively certain examples to relatively unreliable samples. Additionally, by using hyperparameter tuning to adjust the learning weights of the mean squared error loss and the Dice loss, the student model can stabilize numerical regression while optimizing the teacher model's prediction of the shape and boundary of the lesion area, achieving an overall improvement in segmentation accuracy.
[0124] In one embodiment of the present invention, the experimental dataset herein uses the BraTS2019
[47] and LA
[48] datasets. BraTS, as an important public dataset for multimodal brain tumor segmentation, is widely used in the research of this topic. The BraTS2019 dataset contains 335 MRI samples of glioma patients with expert annotations. The MRI scans of each patient include four different modalities, namely T1, T1ce, T2, and Flair. We use the FLAIR modality for tumor segmentation because this modality can better display malignant tumors. In the experiment, all MRI scans are resampled to the same resolution size and the intensity is normalized to zero mean and unit variance. In the experiment, the BraTS2019 dataset is divided. 250 scans are used for training, 25 scans are used as the validation set, and the remaining 60 scans are used as the test set. Among the 250 training scans, 10% of 25 scans and 20% of 50 scans are used as labeled data according to different settings, and the remaining scans are used as unlabeled data. The LA dataset is a public dataset focused on the task of left atrium segmentation, containing multiple left atrium images with expert annotations. In this study, the LA dataset is divided into a training set of 80 scans and 20 scans as the validation set. The network is trained using 16 cases, that is, 20% of the labeled data.
[0125] To ensure a fair comparison of experimental results, the V-Net
[26] is used as the backbone structure in this experiment. To control the balance between the supervised segmentation loss and the unsupervised consistency loss, the Gaussian ramp function is used in the experiment 。The Stochastic Gradient Descent (SGD) optimizer is used to update the network parameters with an initial learning rate of 1e-2, which decays by 0.1 every 2500 iterations. The maximum number of training iterations is set to 6000. The batch size is set to 4, with each mini-batch containing 2 labeled images and 2 unlabeled images. In this experiment, 112×112×80 sub-volumes are randomly cropped as the network input, and a sliding window strategy is used to obtain the final segmentation result. Standard data augmentation techniques are used during training to avoid overfitting, including random flipping and rotation by 90 degrees, 180 degrees, and 270 degrees along the axial plane. To quantitatively evaluate the segmentation results, four complementary evaluation metrics are used. Two metrics, the region-based Dice similarity coefficient (Dice) and Jaccard index (Jaccard), are used to measure region mismatch. The Average Surface Distance (ASD) and 95% Hausdorff Distance (95HD), two boundary-based metrics, are used to evaluate the boundary error between the segmentation result and the ground truth.
[0126] Tables 1 and 2 show the quantitative comparison results of the method of the present invention and current mainstream semi-supervised medical image segmentation methods on the BraTS2019 dataset, including baseline models such as MT (MeanTeachers), UA-MT (uncertainty-aware self-ensembling mean teacher framework), URPC (Uncertainty Rectified Pyramid Consistency), DTC (Dual-task Consistency), and UG-MCL (Uncertainty-guided mutual consistency learning).
[0127] Table 1 compares the performance of each method on the BraTS2019 dataset using 10% labeled data
[0128]
[0129] Table 2 compares the performance of each method on the BraTS 2019 dataset using 20% labeled data
[0130]
[0131] As can be seen from the results in Table 1, the UGMLDS-Net (Uncertainty-Guided Multi-Level Distillation Segemntation Network) proposed in the present invention has achieved the best results in all four metrics. Specifically, in terms of the core metric of Dice similarity coefficient, UGMLDS-Net reached 85.21%, leading the second-ranked method UG-MCL by approximately 2%. This improvement indicates that UGMLDS-Net has a significant advantage in the overall segmentation accuracy of the lesion area and can more accurately identify and segment complex lesion structures. In terms of the key metric of 95HD, which measures the segmentation boundary error, the performance of UGMLDS-Net is even more prominent, with a result of 8.41, leading the 11.44 pixels of UG-MCL by approximately 35%. This shows that UGMLDS-Net not only performs excellently in overall segmentation accuracy but also has a significant advantage in boundary capture ability and can more precisely locate the edge area of the lesion. In addition, in the other two metrics, UGMLDS-Net also demonstrated comprehensive performance advantages. These results indicate that UGMLDS-Net is superior to other methods in terms of segmentation accuracy, boundary capture ability, and robustness.
[0132] In addition, in the other two metrics, UGMLDS-Net also demonstrated comprehensive performance advantages. These results indicate that UGMLDS-Net is superior to other methods in terms of segmentation accuracy, boundary capture ability, and robustness. Further, Table 3 shows the performance of UGMLDS-Net on the left atrial segmentation dataset. The experimental results show that UGMLDS-Net can still achieve the best segmentation results when only 20% of the labeled data is used. This further verifies the generality and effectiveness of UGMLDS-Net in different medical image segmentation tasks.
[0133] Table 3 Performance comparison of each method on the LA dataset using 20% labeled data
[0134]
[0135] To more intuitively demonstrate the performance advantages of UGMLDS-Net, Figure 3 and 4Visualization results in two representative brain regions and the left atrial segmentation task are presented respectively. By comparing the segmentation effects of different methods in these key regions, the excellent ability of UGMLDS-Net in dealing with complex lesion structures and fuzzy boundary regions can be clearly observed. Specifically, in the core region of the lesion, UGMLDSNet shows higher detail capture ability, being able to accurately identify and segment the fine structures of the lesion; while in the edge region, its boundary segmentation effect is smoother and more accurate, effectively avoiding the problems of over-segmentation or under-segmentation. This advantage is consistent with the significant improvement of the Dice coefficient and 95HD index in the experiment, further verifying the superiority of UGMLDS-Net in semi-supervised medical image segmentation tasks. These results indicate that UGMLDS-Net not only performs excellently in quantitative indicators, but also has significant advantages in the actual segmentation effect, being able to provide more reliable and accurate segmentation results for clinical diagnosis. In summary, by combining uncertainty estimation and multi-level knowledge distillation mechanisms, UGMLDS-Net realizes the efficient utilization of labeled and unlabeled data in semi-supervised medical image segmentation tasks, thus achieving the optimal performance in all evaluation metrics. This method provides an effective solution to the problem of scarce labeled data in medical image segmentation.
[0136] Those of ordinary skill in the art will realize that the embodiments described herein are for helping the reader understand the principles of the present invention, and it should be understood that the protection scope of the present invention is not limited to such specific statements and embodiments. Those of ordinary skill in the art can make various other specific deformations and combinations that do not depart from the essence of the present invention according to these technical revelations disclosed in the present invention, and these deformations and combinations are still within the protection scope of the invention.
Claims
1. A medical image segmentation method based on uncertainty estimation and multi-level distillation, characterized in that It includes the following steps: S1: Quantify the prediction uncertainty of the teacher model for the input medical image through the uncertainty estimation module to generate uncertainty weights; S2: Screen high-confidence pseudo-labels based on the uncertainty weights, and transfer the knowledge of the teacher model to the student model in stages through the multi-level knowledge distillation module; S3: Embed a hybrid attention module in the teacher model and the student model, and combine the channel attention mechanism and the spatial attention mechanism to enhance the feature capture of the lesion structure and boundary; S4: Adopt a composite loss function that combines cross-entropy loss, Dice loss, and uncertainty-based distillation loss to optimize the training process of the student model; S5: Apply the trained network to the medical image segmentation scenario to complete medical image segmentation based on uncertainty estimation and multi-level distillation.
2. The medical image segmentation method based on uncertainty estimation and multi-level distillation according to claim 1, wherein In the S1, the uncertainty estimation module is implemented by the Monte Carlo Dropout method, which specifically includes: A1: Enable the Dropout layer during the inference stage of the teacher model to perform multiple forward predictions on the same input; A2: Calculate the prediction entropy based on the probability distribution of the multiple prediction results as a quantization index of the segmentation uncertainty: ; ; wherein, is the prediction entropy, is the number of classes in the segmentation task, is the average softmax probability of T samples from the teacher model, is the prediction value of the class in the t-th forward pass of the teacher model; A3: Dynamically adjust the screening threshold of the pseudo-labels based on the prediction entropy to filter the prediction results in the high-uncertainty regions.
3. The medical image segmentation method based on uncertainty estimation and multi-level distillation according to claim 1, wherein, The implementation of the multi-level knowledge distillation module in the S2 includes: B1: Update the weights between the teacher model and the student model through exponential moving average to generate an intermediate model; B2: Use the intermediate model as a new teacher model for secondary distillation to gradually transfer multi-level semantic features; B3: During the distillation process, use the uncertainty map based on the prediction entropy to optimize the mean square error loss and Dice loss between the prediction of the student model and the prediction of the teacher model, and at the same time calculate the cross-entropy loss and Dice loss between the student model and the binary label to constrain the prediction consistency between the student model and the teacher model.
4. The medical image segmentation method based on uncertainty estimation and multi-level distillation according to claim 1, characterized in that In the S3, the network architectures of both the teacher model and the student model are 3D Unet, and a hybrid attention module is embedded in the bridging part of the 3D Unet. The hybrid attention module includes the following operations: C1: Perform global average pooling and global max pooling on the input feature map respectively to generate two channel descriptors: ; ; Among them, and are two channel descriptors, is the number of channels, and are the height and width of the feature map respectively, and represent the th row and the th column elements of the feature map respectively; C2: Interact with the two channel descriptors through a multi-layer perceptron to generate channel attention weights: ; Among them, is the channel attention weight, is a multi-layer perceptron; C3: Multiply the channel attention weights with the input feature map element by element to obtain the channel-enhanced feature map, and perform spatial dimension pooling on the channel-enhanced feature map, and generate spatial attention weights through convolution operations. The channel-enhanced feature map is expressed as: ; Among them, is the feature map after channel enhancement.
5. The medical image segmentation method based on uncertainty estimation and multi-level distillation according to claim 1, wherein The composite loss function in S4 is as follows: ; Among them, is the cross-entropy loss, is the Dice loss, is the distillation loss of uncertainty, and is the balance coefficient.
6. The medical image segmentation method based on uncertainty estimation and multi-level distillation according to claim 5, characterized in that The cross-entropy loss is: ; Among them, is the total number of samples, is the prediction of the student model for the current sample , is the current sample 's actual label.
7. The medical image segmentation method based on uncertainty estimation and multi-level distillation according to claim 6, wherein, The Dice loss is: ; Among them, is the intersection symbol.
8. The medical image segmentation method based on uncertainty estimation and multi-level distillation according to claim 7, wherein, The distillation loss of the uncertainty is: ; wherein, is a Gaussian warm-up function related to time, is the current iteration number reached during the training process, is the segmentation uncertainty of the teacher model for the current sample ; is the threshold of the most certain sample of the teacher model, is a hyperparameter, is the prediction of the student model for the current sample ; and the prediction of the teacher model for the current sample ; is the mean square error loss between them, is the prediction of the student model for the current sample ; and the prediction of the teacher model for the current sample ; is the Dice loss between them, represents a selection operation. If holds, the value is 1; otherwise, the value is 0. ; ; Among them, is a single voxel, is the total number of voxels in the current sample, is the prediction value of the student model for the th sample and the th voxel, is the prediction value of the teacher model for the th sample and the th voxel.
Citation Information
Patent Citations
Multi-stage unsupervised domain adaptive causal relationship identification method
CN114090770A
Mutual learning semi-supervised 3D medical image segmentation method based on heterogeneous perception
CN118762175A
Semi-supervised medical image segmentation method for eye movement guided hybrid data enhancement
CN119205802A
Cited By
Heterogeneous feature knowledge distillation method oriented to semantic segmentation task of smart home image
CN120726633A
Landslide identification method based on multilevel heterogeneous knowledge distillation
CN121074679A
Intelligent grouting simulation method and system based on large model
CN121145747A
Grouting intelligent simulation method and system based on large model
CN121145747B
Distillation learning method and device for skin problem segmentation model, equipment and medium
CN121415215A