A semi-supervised multi-organ segmentation method based on Rubik's Cube segmentation and restoration
Through the Rubik's Cube segmentation recovery idea and the dual-branch data enhancement method, the problem of mismatch between labeled data and unlabeled data distribution in multi-organ segmentation is solved, and more accurate multi-organ segmentation effect is achieved, which improves the multi-organ segmentation effect of semi-supervised learning.
Patent Information
- Application Number
- CN202211630590.7
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2022-12-16
- Publication Date
- 2025-07-08
- Estimated Expiration
- 2042-12-16
AI Technical Summary
The existing semi-supervised medical image segmentation method is difficult to effectively deal with the problem of mismatch between labeled data and unlabeled data in multi-organ segmentation task, resulting in the anatomy being ignored and the sub-optimization results are frequently reported.
Using the idea of slicing and restoration of Rubik's Cube, treat 3D images as Rubik's Cube, cut into small pieces and mix across images, design a dual-branch data enhancement method, including branches between and within images, use the anatomical prior of multiple organs, feature extraction and loss calculations are performed through deep neural networks, and model parameters are updated in combination with exponential moving average.
It effectively solves the problem of mismatch between labeled data and unlabeled data distribution, improves the accuracy and consistency of multi-organ segmentation, and improves the multi-organ segmentation effect of semi-supervised learning.
Smart Images

Figure CN115841494B_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to the technical field of computer vision and digital image processing, and in particular to a semi-supervised multi-organ segmentation method based on Rubik's Cube segmentation and restoration. Background Art
[0002] Abdominal multi-organ segmentation in medical images is an important task in many clinical applications, such as computer-aided diagnosis, organ volume measurement, etc. However, training an accurate multi-organ segmentation model usually requires a large amount of labeled data, and the acquisition process is both time-consuming and expensive. Semi-supervised learning has shown great potential in dealing with the scarcity of data annotation, and it attempts to transfer a large amount of prior knowledge learned from labeled images to unlabeled images. In recent years, semi-supervised learning has attracted more and more attention in the field of medical image analysis.
[0003] Current semi-supervised medical image segmentation methods mainly focus on segmenting single objects or local regions, such as segmenting the pancreas or the left atrium. Multi-organ segmentation is more challenging than single-organ segmentation because of the complex anatomical structure of organs. For example, the relatively fixed organ positions (the duodenum is always located at the head of the pancreas), the different shapes of organ appearances, and the different sizes of organs. Therefore, migrating current semi-supervised medical segmentation methods to multi-organ segmentation tasks will encounter serious problems. Compared with a single organ, the distribution differences introduced by multiple organs are much larger. Although labeled images and unlabeled images are always extracted from the same distribution, it is difficult to estimate an accurate distribution from them due to the limited number of labeled images. Therefore, there is always a mismatch problem between the estimated distributions of labeled and unlabeled images, which is even further amplified in multi-organ segmentation tasks. Current semi-supervised medical image segmentation methods lack the ability to handle such a large distribution gap, inevitably ignoring the internal anatomical structure of multiple organs and resulting in sub-optimal results. Summary of the Invention
[0004] In view of the above-mentioned defects of the prior art, the present invention is inspired by the idea of scrambling small pieces of a Rubik's Cube game and then restoring them to the starting position. We apply this idea of "segmentation-restoration" to semi-supervised learning. We regard 3D images as Rubik's Cubes, cut them into small pieces and perform cross-image mixing as the starting inputs within and between images respectively, and restore the small pieces to their original positions during the prediction stage. The purpose of the present invention is to utilize the anatomical prior of multi-organs themselves to solve the problem of distribution mismatch between labeled data and unlabeled data during the training process of semi-supervised learning. A dual-branch data augmentation method is designed for the relatively fixed relative positions and different organ sizes of multi-organs themselves, including an inter-image branch and an intra-image branch.
[0005] To achieve the above object, the present invention provides a semi-supervised multi-organ segmentation method based on Rubik's Cube segmentation. The method comprises the following steps:
[0006] Cut all images into small images for input to the intra-image branch; randomly mix the small images of the labeled images and unlabeled images across images to form mixed images for input to the inter-image branch; the inputs of the two branches respectively pass through a deep neural network to obtain features and predictions at two data levels of the mixed images and small images.
[0007] For the intra-image branch, input the features of the small images into a classifier, infer the relative positions of the small images within the image, and calculate the cross-entropy loss function of the classifier prediction and the corresponding relative positions.
[0008] Restore the prediction results of the two branches into segmentation predictions corresponding one by one to the original images; for the labeled images, calculate the loss function between the predictions of the two branches and the ground truth mask; for the unlabeled images, perform a weighted average of the predictions of the intra-image branch and the predictions of the teacher network based on the class distribution to obtain a pseudo-mask, and calculate the loss function between the inter-image predictions and the pseudo-mask.
[0009] Use the loss to perform gradient backpropagation, update the parameters of the student model and the classifier, and update the parameters of the teacher model using the exponential moving average method; when the training reaches convergence or the maximum number of times, obtain the final student network parameters.
[0010] Preferably, the deep neural network in the method is trained with a convolutional neural network "encoder-decoder" architecture as the backbone network, and the classifier is composed of two fully connected layers.
[0011] Preferably, the loss function is:
[0012]
[0013] Among them, the calculation methods of the loss functions for the labeled images and unlabeled images are respectively:
[0014]
[0015]
[0016] Among them, represents the current batch of labeled / unlabeled images, Θ s represents the student network parameters, represents the parameters of the encoder part of the student network, Θ cls represents the classifier parameters, and α, β represent the balance factors of the loss function.
[0017] Preferably, the cross-entropy loss function is as follows:
[0018]
[0019] Where X represents the original input image, represents cutting the original image into N 3 small pieces, represents the relative position of the small pieces in the original image, represents the cross-entropy loss function, σ represents the softmax layer, represents the classification head, represents the encoder of the student network. For labeled images, the segmentation loss functions of the inter-image and intra-image branches are respectively expressed as:
[0020]
[0021] Where, represents the dice loss function, respectively represent the segmentation predictions after restoration of the inter-image and intra-image branches, Y l represents the ground truth label. For unlabeled images, the present invention adopts a mixed supervision method, specifically as follows:
[0022]
[0023] Where, represents the dice loss function, represents the segmentation prediction after restoration of the inter-image branch, represents the final pseudo-label obtained after the small-piece level feature mixing module processes the teacher model prediction and the intra-image prediction.
[0024] Preferably, for unlabeled images, the specific method of obtaining a pseudo-mask by performing weighted averaging of the intra-image branch prediction and the teacher network prediction based on the class distribution is:
[0025] Initialize a class distribution repository D, and update the repository D every T rounds during training;
[0026] Let the current iteration number be t. If t % T ≠ 0, store the teacher prediction of the unlabeled image in this training round in D; if t % T = 0, count the number of pixels of each organ class in the pseudo-labels stored in D and normalize it between 0 and 1 to obtain the class distribution dictionary vector v:
[0027] v = {v0,..., v C-1}
[0028] Where C represents the number of organ classes, and each element in v represents the number of pixels of each class after normalization;
[0029] Subsequently, empty D;
[0030] For any pixel m in the teacher prediction, assume that the prediction of the teacher network for this pixel is Query the class distribution value in v corresponding to , that is, for any pixel, the corresponding class distribution value can be found in the dictionary vector v, so as to generate a pixel-level weight map Take Ω as the weight of the in-image branch prediction, take (1 - Ω) as the weight of the teacher network, and perform a weighted sum of the teacher prediction and the in-image branch prediction to obtain the final pseudo-mask.
[0031] Starting from the anatomical prior of abdominal multi-organ scanning itself, the present invention designs a data augmentation method specifically for semi-supervised multi-organ segmentation to effectively address the problem of the distribution mismatch between labeled data and unlabeled data in the context of semi-supervised learning. In view of the fixed relative positions of multi-organs themselves and the differences in organ sizes, a dual-branch data augmentation method is designed, including an inter-image branch and an in-image branch.
[0032] The following will further illustrate the concept, specific structure and technical effects of the present invention in conjunction with the drawings to fully understand the purpose, features and effects of the present invention. BRIEF DESCRIPTION OF THE DRAWINGS
[0033] Figure 1 is a flowchart of a preferred embodiment of the present invention. DETAILED DESCRIPTION OF THE INVENTION
[0034] The following introduces multiple preferred embodiments of the present invention with reference to the accompanying drawings of the specification to make its technical content clearer and easier to understand. The present invention can be embodied in many different forms of embodiments, and the protection scope of the present invention is not limited to the embodiments mentioned in the text.
[0035] In the drawings, components with the same structure are denoted by the same numerical reference signs, and components with similar structures or functions are denoted by similar numerical reference signs. The size and thickness of each component shown in the drawings are arbitrarily shown, and the present invention does not limit the size and thickness of each component. In order to make the illustration clearer, the thickness of some components in the drawings is appropriately exaggerated.
[0036] Refer to Figure 1 , the object of the present invention is to utilize the anatomical prior of CT multi-organs itself to solve the problem of the distribution mismatch between labeled data and unlabeled data during the training process in semi-supervised learning. In view of the fixed relative positions of CT multi-organs themselves and the differences in organ sizes, a dual-branch data augmentation method is designed, including an inter-image branch and an in-image branch. The specific steps are as follows:
[0037] Step 1: First, randomly select a labeled 3D CT image and an unlabeled 3D CT image; split the two CT image data into N 3 A small block image is used as the input of the intra-image branch; and under the premise of keeping the relative positions of each small block image in the original image, cross-image mixing is performed to obtain a mixed image for the input of the inter-image branch; the input images of the two branches are respectively passed through the deep neural network to obtain the prediction of the inter-image branch, as well as the features and predictions of the intra-image branch.
[0038] Step 2: For the in-image branch, set N 3 The features of small patches of images are fed into the classifier, and the inference is N 3 The relative position of each small patch of images in the 3D image is calculated, and the cross entropy loss between the classifier prediction and the corresponding relative position is calculated.
[0039] Step 3: Restore the prediction results of the two branches to the original Figure 1 For each corresponding segmentation prediction, for labeled images, the DICE loss function between the predictions of the two branches and the true value mask is calculated; for unlabeled images, the predictions of the intra-image branches and the predictions of the teacher network are weighted averaged based on the class distribution to obtain a pseudo mask, and the DICE loss function between the inter-image prediction and the pseudo mask is calculated;
[0040] Step 4: The final overall loss function is the weighted average of the losses in steps 1, 2, and 3, and the overall loss is used for gradient backpropagation to update the parameters of the student model and classifier; the teacher model parameters are updated using the exponential moving average method. The goal of this method is to obtain the final student network parameters when the training reaches convergence or the maximum number of times.
[0041] The deep neural network in the method is trained with the convolutional neural network "encoder-decoder" architecture as the backbone network, and specifically V-Net and 3D U-Net can be selected; the classifier is composed of two fully connected layers.
[0042] V-Net provides a 3D image segmentation method that uses an end-to-end training method and uses a new objective function based on the Dice coefficient to optimize the training. It can handle situations where there is a severe imbalance between the number of foreground and background voxels. In order to deal with situations where there is limited data available for training, it uses random nonlinear transformations and histogram matching to enhance the data.
[0043] 3D U-Net is a simple extension of UNet that replaces all 2D operations with 3D operations. For volumetric images, instead of separately inputting each slice for training, the entire image is input into the model. 3D U-Net is applicable to three-dimensional image segmentation problems.
[0044] The small patch image-level feature mixing module in the described method has the following specific steps:
[0045] By performing a weighted average of the teacher predictions and the in-image branch predictions for unlabeled images, a final pseudo-mask is generated to supervise the inter-image branch predictions. The weights used in this module are class distribution-based weight maps. The specific approach is as follows: First, initialize a class distribution repository D, and update the repository D every T rounds during training. Let the current iteration number be t. If t % T ≠ 0, store the teacher predictions of the unlabeled images in this training round in D; if t % T = 0, count the number of pixels of each organ class in the pseudo-labels stored in D and normalize it between 0 and 1 to obtain the class distribution dictionary vector v:
[0046] v = {v0,..., v C-1}
[0047] where C represents the number of organ classes, and each element in v represents the number of pixels of each class after normalization. Then empty D. For any pixel m in the teacher prediction, let the prediction of the teacher network for this pixel be Query the class distribution value in v corresponding to That is, for any pixel, the corresponding class distribution value can be found in the dictionary vector v, thereby generating a pixel-level weight map Use Ω as the weight for the in-image branch prediction, use (1 - Ω) as the weight for the teacher network, and perform a weighted sum of the teacher prediction and the in-image branch prediction to obtain the final pseudo-mask to supervise the prediction of the inter-image branch.
[0048] The final overall loss function of the described method is:
[0049]
[0050] where the calculation methods of the loss functions for labeled and unlabeled images are respectively:
[0051]
[0052]
[0053] where, represents the current batch of labeled / unlabeled images, Θ s represents the student network parameters, Denote the parameters of the student network encoder as Θ cls Denote the classifier parameters as α, and β represents the balancing factor of the loss function. The classifier loss function is calculated as follows:
[0054]
[0055] where X represents the input original image, denotes slicing the original image into N 3 patches, represents the relative position of the patch in the original image, denotes the cross-entropy loss function. For labeled images, the segmentation loss functions for the inter-image and intra-image branches are respectively expressed as:
[0056]
[0057] where, denotes the dice loss function, respectively represent the segmentation predictions after recovery for the inter-image and intra-image branches, and Y l represents the ground truth label. For unlabeled images, the present invention adopts a semi-supervised manner, specifically as follows:
[0058]
[0059] where, denotes the dice loss function, represents the segmentation prediction after recovery for the inter-image branch, represents the final pseudo-mask obtained after the small patch-level feature mixing module for the teacher model prediction and the intra-image prediction.
[0060] Based on the anatomical prior of the CT abdominal multi-organ scan itself, the present invention designs a data augmentation method specifically for semi-supervised CT multi-organ segmentation to effectively address the problem of the distribution mismatch between labeled data and unlabeled data in the context of semi-supervised learning.
[0061] The above has described in detail the preferred specific embodiments of the present invention. It should be understood that those of ordinary skill in the art can make many modifications and variations according to the concept of the present invention without creative labor. Therefore, all technical solutions that can be obtained by those skilled in the art in the technical field based on the concept of the present invention through logical analysis, reasoning, or limited experiments on the basis of the prior art should fall within the protection scope determined by the claims.
Claims
1. A semi-supervised multi-organ segmentation method based on Rubik's Cube segmentation, characterized in that, Regarding the 3D image as a Rubik's Cube and using multi-organ anatomical priors to design a data augmentation method, the method includes the following steps: Cut all images into small patch images for the input of the intra-image branch; randomly mix the small patch images of the labeled images and unlabeled images across images to form mixed images for the input of the inter-image branch; the inputs of the two branches respectively pass through a deep neural network to obtain features and predictions at two data levels of the mixed images and small patch images; For the intra-image branch, input the features of the small patch images into a classifier, infer the relative position of the small patch images within the image, and calculate the cross-entropy loss function of the classifier prediction and the corresponding relative position; Restore the prediction results of the two branches into segmentation predictions corresponding one-to-one with the original images; for the labeled images, calculate the loss function between the predictions of the two branches and the ground truth mask; for the unlabeled images, perform a weighted average of the predictions of the intra-image branch and the predictions of the teacher network based on the class distribution to obtain a pseudo-mask, and calculate the loss function between the inter-image prediction and the pseudo-mask; For the unlabeled images, the specific method of performing a weighted average of the predictions of the intra-image branch and the predictions of the teacher network based on the class distribution to obtain a pseudo-mask is as follows: Initialize a class distribution repository D, and update the repository D every T rounds during the training process; Let the current iteration number be t. If t % T ≠ 0, store the teacher predictions of the unlabeled images in this training round in D; if t % T = 0, count the number of pixels of each organ class in the pseudo-labels stored in D and normalize it between 0 and 1 to obtain a class distribution dictionary vector v: v = {v0,..., v C-1} Where C represents the number of organ classes, and each element in v represents the number of pixels of each class after normalization; Subsequently, empty D; For any pixel m in the teacher prediction, assume that the prediction of the teacher network for this pixel is Query the class distribution value in v corresponding to That is, for any pixel, the corresponding class distribution value can be found in the dictionary vector v, so as to generate a pixel-level weight map Take Ω as the weight of the in-image branch prediction, take (1 - Ω) as the weight of the teacher network, and perform a weighted sum of the teacher prediction and the in-image branch prediction to obtain the final pseudo-mask; Use the loss to perform gradient backpropagation, update the parameters of the student model and the classifier, and update the parameters of the teacher model using the exponential moving average method; when the training reaches convergence or the maximum number of times, obtain the final parameters of the student network.
2. The method according to claim 1, wherein The deep neural network in the method is trained with a convolutional neural network "encoder-decoder" architecture as the backbone network, and the classifier is composed of two fully connected layers.
3. The method according to claim 1, characterized in that The loss function is: Among them, the calculation methods of the loss functions for the labeled images and unlabeled images are respectively: Among them, represents the current batch of labeled / unlabeled images, Θ s represents the student network parameters, represents the parameters of the encoder part of the student network, Θ cls represents the classifier parameters, and α, β represent the balancing factors of the loss function.
4. The method according to claim 1, characterized in that, The cross-entropy loss function is as follows: Where X represents the original input image, represents cutting the original image into N 3 small pieces, represents the relative position of the small pieces in the original image, represents the cross-entropy loss function, σ represents the softmax layer, represents the classification head, represents the encoder of the student network. For labeled images, the segmentation loss functions for the inter-image and intra-image branches are respectively expressed as: Among them, represents the dice loss function, respectively represent the segmentation predictions after the restoration of the inter-image and intra-image branches, Y l represents the ground truth label. For unlabeled images, a mixed supervision method is adopted as follows: Among them, represents the dice loss function, represents the segmentation prediction after the recovery of the image - to - image branch, represents the final pseudo - label obtained after the teacher model prediction and the in - image prediction pass through the patch - level feature mixing module.