A Single-Domain Generalization Method for Medical Image Segmentation
The single-domain generalization strategy for medical image segmentation uses data augmentation and a dual-branch network to learn domain-invariant features, addressing privacy and annotation challenges, and enhances model performance and precision across diverse medical image datasets.
Patent Information
- Application Number
- CN202310129544.7
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2023-02-03
- Publication Date
- 2025-07-15
- Estimated Expiration
- 2043-02-03
AI Technical Summary
Existing deep learning models face the problem of insufficient cross-center and cross-modal adaptation and generalization capabilities in medical image segmentation tasks, especially in terms of data privacy and high data annotation costs.
A single domain generalization strategy is adopted to enhance the simulated domain offset through weak data, and a dual-branch consistent network is used to learn domain invariant representation, combining features to guide the whitening method, and improve the cross-center generalization ability and segmentation accuracy of the model.
It effectively improves the segmentation performance and robustness of the model in different data domains, and can show good generalization ability and segmentation accuracy on multi-center data sets.
Smart Images

Figure CN116596832B_ABST
Abstract
Description
Technical Field
[0001] The present invention belongs to the field of medical image analysis, and particularly relates to a single-domain generalization method for medical image segmentation. Background Art
[0002] In a computer-aided clinical diagnosis system, the segmentation algorithm of human tissues, organs, lesions, etc. is an important link. At present, deep learning systems have been widely used in medical image segmentation problems and achieved high accuracy. However, due to factors such as different devices, different patient groups, and different imaging parameters, deep learning models often face a significant decline in performance when processing images in a new data domain. To reach the clinically available level, it is an urgent problem to improve the cross-center and cross-modal adaptation and generalization capabilities of deep learning models.
[0003] Generally speaking, the existing methods for improving the domain generalization (DG) ability of deep learning networks can be divided into: data manipulation-based methods and alignment-based methods. The former aims to simulate unknown target domain data through data augmentation or data generation to improve the generalization ability of the model. For example, Chinese Patent CN114694150, with the publication date of July 1, 2022, discloses a method and system for improving the generalization ability of a digital image classification model, which uses a data augmentation method based on mixed samples and uses the gradient information of the classification network for data augmentation to improve the generalization ability of the model. Its disadvantages are: on the one hand, data augmentation is a simple expansion, transformation, or randomization of the input image, and it is difficult to accurately simulate the target domain data, resulting in limited generalization ability of the trained model; on the other hand, considering that medical images contain rich tissue or organ details, it is difficult for the data generation model to accurately generate realistic medical images; while the latter aims to align the data distributions of multiple source domains to learn domain-invariant representations, and then improve the generalization ability of the model on unseen target domains. For example, Chinese Patent CN115018874A, with the publication date of September 6, 2022, discloses a domain generalization method for fundus vascular segmentation based on frequency domain analysis, which uses a feature normalization algorithm based on frequency domain analysis to learn a more unified semantic frequency distribution space among multiple source domains, thereby enhancing the expression ability and generalization ability of the model. Its disadvantages are: on the one hand, in the field of medical images, data has privacy issues, and data sharing is prohibited before taking complex procedures including ethical approval and patient consent, and it is difficult to aggregate data from multiple domains. On the other hand, in medical image segmentation tasks, data annotation is difficult, costly, and time-consuming, and requires experts in related fields for annotation, and it is difficult to collect annotated data from multiple domains. Summary of the Invention
[0004] The object of the present invention is to provide a single-domain generalization method for medical image segmentation in view of the deficiencies existing in the prior art.
[0005] The present invention first adopts a single-domain generalization strategy, avoiding the problems of cross-center data privacy and high data annotation costs; based on the medical image segmentation problem, in view of the characteristics of medical images, the present invention innovatively uses weak data augmentation to simulate domain shift, and proposes a dual-branch consistency network to learn domain-invariant representations, improving the cross-center generalization ability of the model; at the same time, a feature-guided whitening method is proposed to further guide the model to focus on semantic information and ignore style information, so as to improve the model's expression ability and segmentation accuracy.
[0006] A single-domain generalization method for medical image segmentation, comprising:
[0007] (1) Unify the size of the labeled input instance image x in the training domain, and perform data augmentation operations to obtain augmented data x′ to simulate domain shift.
[0008] (2) Use two encoders Encoder with shared parameters to extract features from the original image x and the augmented image x′ respectively, obtain two feature maps M and M′ with dimensions of C×H×w, and perform instance normalization;
[0009] (3) Input the features M and M′ into two decoders Decoder with shared parameters respectively to obtain prediction probability maps f(x) and f(x′). The two prediction probabilities Figure 1 are respectively calculated with the true label y to obtain the segmentation loss to improve the model performance; on the other hand, a consistency loss is calculated between the two prediction maps to learn domain-invariant representations and improve the generalization ability of the model;
[0010] (4) Perform a difference operation on the feature maps M and M′ to obtain a highly sensitive style information feature map M s , and perform dimensionality reduction through global average pooling and normalization to obtain a sensitive vector s. The sensitive vector and its transpose s T are used for inner product to obtain a sensitivity map SM, and a whitening mask WM is obtained through threshold selection. Finally, the covariance matrices ∑ M and ∑′ M of the feature maps M and M′ are guided for whitening and the feature whitening loss is calculated
[0011] (5) Minimize the overall loss. When training the above segmentation network, when the total loss function does not reach the preset convergence condition, the parameters of the above segmentation network are iteratively updated until the total loss function reaches the preset convergence condition.
[0012] (6) Target image segmentation. For a given target domain image, the segmentation model outputs the probability of each pixel in the target image belonging to a category, and sets a threshold to label the pixel as a segmented target or background.
[0013] The image processing enhancement method in the above step (1) is as follows:
[0014] (i) Unify the image size. Use the bilinear interpolation algorithm to scale all images x to a fixed size, and use the nearest neighbor difference to scale the corresponding label y to the same size. This makes the input image conform to the input specifications of the segmentation network;
[0015] (ii) Image weak enhancement operation. Given that the domain shift of medical images mainly comes from changes such as image brightness and color. Therefore, perform weak enhancement operations such as illumination and color changes on the unified-size image x to obtain enhanced data x′ to simulate the domain shift.
[0016] The feature extraction and normalization method in step (2) described above is as follows:
[0017] (i) Use two encoders with shared parameters to extract features from the original image x and the enhanced image x′, as shown in the formula.
[0018]
[0019] Among them, the encoder Encoder can adopt various types of segmentation network feature extractors, such as UNet, 3DUNet, DeepLab, etc.
[0020] (ii) Use the instance normalization method to normalize the extracted features.
[0021]
[0022] Where H and W represent the sizes of the feature maps M and M′, c represents the channel index, and j, k, l, m represent the position indices, ∈→0 to prevent the denominator from being 0.
[0023] The domain-invariant representation learning method in step (3) described above is as follows:
[0024] (i) Input the features M and M′ into two decoders (Decoder) with shared parameters to obtain the predicted probability maps f(x) and f(x′);
[0025] f(x) = Decoder(M)
[0026] f(x′) = Decoder(M′)
[0027] (ii) Calculate the segmentation loss of the predicted probability maps f(x) and f(x′) with the true label y To improve the model performance. The segmentation loss is used to measure the difference between the predicted output and the ground truth label, and the pixel-level cross-entropy loss is adopted as the loss function;
[0028]
[0029] where H and W represent the size of the input image, i represents the index of the pixel, C represents the categories of pixel predictions in the segmentation, c is the index of the category, y ic ∈ {0, 1} is the one-hot encoding of the ground truth label, and p ic is the predicted class probability. Minimizing the segmentation loss can optimize the model performance and improve the model segmentation accuracy.
[0030] (iii) Calculate the consistency loss between the predicted probability maps f(x) and f(x′) to learn domain-invariant representations and improve the model generalization ability. The consistency loss is used to measure the difference between the predicted probability maps of the original image and the augmented image, and the mean squared error loss is adopted as the loss function;
[0031]
[0032] where f(·) represents the segmentation network composed of an encoder and a decoder, and f(x) and f(x′) represent the predicted probability maps of dimension C×H×W output by the network. Minimizing the consistency loss can align the prediction distributions of the network for the original image and the augmented image, enabling the network to focus on the semantic information shared by the domains and ignore the domain-specific style information, thereby learning domain-invariant representations and improving the out-of-domain generalization ability of the network.
[0033] The feature-guided whitening method in step (4) is as follows:
[0034] (i) Given that the augmented image and the original image have the same semantic information but different style information, perform a difference operation on their feature maps to obtain a style information feature map M of dimension C×H×W that only contains style information s ;
[0035] M s = M - M′
[0036] The difference operation can decouple the semantic information and the style information, facilitating the elimination of the influence brought by the style information while keeping the semantic information unchanged.
[0037] (ii) Perform global average pooling operation on the style information feature map M s to obtain a sensitivity vector s, to reduce the dimension and computational amount, and at the same time suppress overfitting; and normalize the sensitivity vector s through a normalization operation to improve the model convergence speed;
[0038]
[0039]
[0040] where H and W represent the dimensions of the style information feature map M s and c is the channel number index, i, j represent the position indices, and Max(·) and Min(·) are used to calculate the maximum and minimum values in the channel dimension.
[0041] (iii) Take the inner product of the sensitive vector s and its transpose s T to obtain the sensitive map SM, and perform threshold selection on the sensitive map to obtain the whitening mask WM, which is used to eliminate the style information in the feature map, further guiding the encoder to focus on semantic information and ignore style information, and improving the generalization ability;
[0042] SM = (s)(s) T ∈R C×C
[0043]
[0044] where i, j represent the position indices, the top(·) function represents returning the value of the top element in the sensitive map SM, ε is a fixed whitening range threshold, set to 0.1 - 0.9, and more preferably, set to 0.4.
[0045] (vi) Calculate the covariance matrices ∑ M and ∑′ M of the feature maps M and M′ respectively to match the dimension of the obtained whitening mask WM;
[0046]
[0047]
[0048] where H and W represent the dimensions of the feature maps M and M′.
[0049] (v) And use the whitening mask to whiten the style features of the covariance matrices ∑ M and ∑′ M to guide the model to focus on semantic information and ignore style information, so as to improve the model's expression ability and segmentation accuracy. The whitening loss is:
[0050]
[0051] where represents the arithmetic mean operation, and ⊙ represents element-wise multiplication.
[0052] The method for minimizing the total loss in step (5) described above is:
[0053] (i) The total calculated loss is the segmentation loss the consistency loss and the feature whitening loss as a linear combination;
[0054]
[0055] where λ1 and λ2 are hyperparameters, and the ranges of the λ1 and λ2 parameters are (0 to 1], which are used to balance the influence of the two losses on the total loss. Further preferably, λ1 and λ2 are respectively set to 0.3 and 0.6.
[0056] (ii) Iteratively update the parameters of the above segmentation network until the total loss function reaches the preset convergence condition;
[0057] The target image segmentation method in step (6) is as follows:
[0058] (i) Input the target image x T into the segmentation model to obtain the predicted probability map f(x T );
[0059] (ii) Set a threshold of 0.5 to 0.9 to label the prediction map. If the predicted probability is greater than the threshold, it is labeled as the segmentation target; if it is less than the threshold, it is labeled as the background to obtain the final segmentation mask.
[0060] Further preferably, set a threshold of 0.5 to label the prediction map. If the predicted probability is greater than 0.5, it is labeled as the segmentation target; if it is less than 0.5, it is labeled as the background to obtain the final segmentation mask;
[0061] Compared with the prior art, the beneficial effects of the present invention are:
[0062] (1) The present invention has a single-domain generalization strategy for domain-invariant representation learning. By weakly augmenting the source domain data to simulate domain shift and using a two-branch consistency loss to learn domain-invariant representations, the generalization ability and robustness of the model are improved;
[0063] (2) Decouple the features through feature difference operations to obtain highly sensitive style features, and use the style features to further perform style whitening on the features to guide the model to focus on semantic information and ignore style information, thereby improving the model's expression ability and segmentation accuracy. BRIEF DESCRIPTION OF THE DRAWINGS
[0064] Figure 1 is a schematic diagram of the overall process of the medical image single-domain generalization segmentation method of the present invention;
[0065] Figure 2 is a multi-center fundus segmentation result map;
[0066] Figure 3 It is a multi - center prostate segmentation result diagram;
[0067] Figure 4 It is a multi - center hepatocellular tumor segmentation result diagram.
[0068] Figure 5 It is a schematic flow diagram of the single - domain generalization method for medical image segmentation of the present invention. Specific implementation mode
[0069] As Figure 1 and Figure 5 shown, a single - domain generalization method for medical image segmentation is as follows:
[0070] Process and enhance the single - source data. Use the bilinear interpolation algorithm to scale all images x to a fixed size, and use the nearest - neighbor interpolation to scale the corresponding labels y to the same size; and perform weak enhancement operations on the single - source domain images to simulate domain shift.
[0071] Then use two encoders with shared parameters to extract features from the original image and the enhanced image. The specific steps are as follows:
[0072] (i) Use two encoders with shared parameters to extract features from the original image x and the enhanced image x′, as shown in the formula.
[0073]
[0074] Among them, the encoder Encoder can adopt various types of segmentation network feature extractors, such as UNet, DeepLab, etc.
[0075] (ii) Use the instance normalization method to normalize the extracted features.
[0076]
[0077] Among them, H and W represent the sizes of the feature maps M and M′, c represents the channel index, j, k, l, m represent the position index, ∈→0 to prevent the denominator from being 0.
[0078] Adopt a two - branch consistency network to learn domain - invariant representations and improve the segmentation accuracy and generalization ability of the model. The specific steps are as follows:
[0079] (i) Input the features M and M′ into two decoders (Decoder) with shared parameters respectively to obtain the predicted probability maps f(x) and f(x′);
[0080] f(x)=Decoder(M)
[0081] f(x′) = Decoder(M′)
[0082] (ii) Calculate the segmentation loss by comparing the predicted probability maps f(x) and f(x′) with the ground truth label y to improve the model performance. The segmentation loss is used to measure the difference between the predicted output and the ground truth label, and the pixel-wise cross-entropy loss is adopted as the loss function;
[0083]
[0084] where H and W represent the size of the input image, i represents the index of the pixel, C represents the class of the pixel prediction in the segmentation, c is the index of the class, y ic ∈ {0, 1} is the one-hot encoding of the ground truth label, and p ic is the predicted class probability. Minimizing the segmentation loss can optimize the model performance and improve the model segmentation accuracy.
[0085] (iii) Calculate the consistency loss between the predicted probability maps f(x) and f(x′) to learn domain-invariant representations and improve the model generalization ability. The consistency loss is used to measure the difference between the predicted probability maps of the original image and the augmented image, and the mean squared error loss is adopted as the loss function;
[0086]
[0087] where f(·) represents the segmentation network composed of an encoder and a decoder, and f(x) and f(x′) represent the predicted probability maps of size C×H×W output by the network. Minimizing the consistency loss can align the predicted distributions of the network for the original image and the augmented image, enabling the network to focus on the semantic information shared by the domains and ignore the domain-specific style information, thereby learning domain-invariant representations and enhancing the out-of-domain generalization ability of the network.
[0088] Secondly, a feature-guided whitening module is proposed, and the specific steps are as follows:
[0089] (i) Given that the augmented image and the original image have the same semantic information but different style information, perform a subtraction operation on their feature maps to obtain a style information feature map M of size C×H×W that only contains style information s ;
[0090] M s = M - M′
[0091] The subtraction operation can decouple the semantic information and the style information, facilitating the elimination of the influence brought by the style information while keeping the semantic information unchanged.
[0092] (ii) Perform global average pooling operation on the style information feature map M s to obtain the sensitivity vector s, so as to reduce the dimension and computational amount, and at the same time suppress overfitting; and normalize the sensitivity vector s through normalization operation to improve the model convergence speed;
[0093]
[0094]
[0095] where H and W represent the size of the style information feature map M s , c is the channel number index, i and j represent the position index, and Max(·) and Min(·) are used to calculate the maximum and minimum values in the channel dimension.
[0096] (iii) Take the inner product of the sensitivity vector s and its transpose s T to obtain the sensitivity map SM, and perform threshold selection on the sensitivity map to obtain the whitening mask WM, which is used to eliminate the style information in the feature map, further guide the encoder to focus on semantic information and ignore style information, and improve the generalization ability;
[0097] SM = (s)(s) T ∈R C×C
[0098]
[0099] where i and j represent the position index, the top(·) function represents returning the value of the top element in the sensitivity map SM, ε is a fixed selection whitening range threshold, set to 0.1 - 0.9, and further preferably, set to 0.4.
[0100] (vi) Calculate the covariance matrices ∑ M and ∑′ M of the feature maps M and M′ respectively to match the dimension of the whitening mask WM obtained in 4.3;
[0101]
[0102]
[0103] where H and W represent the sizes of the feature maps M and M′.
[0104] (v) And use the whitening mask to perform style feature whitening on the covariance matrices ∑ M and ∑′ M to guide the model to focus on semantic information and ignore style information, so as to improve the model expression ability and segmentation accuracy. The whitening loss is:
[0105]
[0106] Among them represents the arithmetic mean operation, and ⊙ represents element-wise multiplication.
[0107] Calculate the total loss And continuously iterate the above steps and update the segmentation network parameters until the total loss reaches the preset convergence.
[0108] Finally, perform the segmentation of the target image, obtain the tissue organs or lesion images to be segmented in other centers, input them into the trained single-domain generalization segmentation model, and automatically segment the masks of the corresponding organs or lesions.
[0109] The present invention uses two multi-center public datasets (prostate MRI segmentation dataset and fundus segmentation dataset) and one multi-center private dataset (hepatocellular tumor segmentation dataset) to evaluate the performance of the present invention. The prostate MRI segmentation dataset includes MRI images of the T2 phase of the prostate collected from six different clinical centers; the fundus segmentation dataset contains color retinal fundus images from four medical institutions; the hepatocellular tumor segmentation dataset is composed of MRI images of the T2 phase of hepatocellular tumors collected by the applicant from seven medical institutions. During the experiment, for each dataset, the data of one center (domain) is selected for model training, and the data of the remaining centers is used for model testing to test the generalization ability of the model across centers.
[0110] The segmentation effects of the present invention on the three datasets are respectively as Figure 2 , 3 , and shown in Figure 4. Among them, the baseline model is a segmentation model that does not adopt any optimization strategy, and the true label is annotated by domain experts. It can be found that on the three datasets, the generalization ability and segmentation performance of the model of the present invention have been greatly improved compared with the baseline model, and the present invention has strong universality and can achieve good generalization ability in various segmentation tasks.
Claims
1. A single-domain generalization method for medical image segmentation, characterized in that, Including the following steps: (1) Unify the size of the labeled medical images x in the training domain and perform data augmentation to obtain augmented data x ′ , to simulate domain shift; (2) Use two Encoders that share parameters to extract features from the original image x and the enhanced image x ′ respectively, obtaining two feature maps M and M', and performing instance normalization; (3) Input the feature maps M and M′ into two decoders Decoder with shared parameters respectively to obtain the predicted probability maps f(x) and f(x′). On the one hand, calculate the segmentation loss between the two predicted probability maps and the ground truth label y respectively On the other hand, calculate the consistency loss between the two predicted maps (4) Perform a difference operation on the feature maps M and M' to obtain a highly sensitive style information feature map M s , and perform dimensionality reduction through global average pooling and normalization to obtain a sensitive vector s. The sensitive vector and its transpose s T perform an inner product to obtain a sensitive map SM, and obtain a whitening mask WM through threshold selection. Finally, calculate the feature whitening loss for the covariance matrices ∑ M and ∑' M (5) Minimize the total loss: When training the encoder Encoder and the decoder Decoder, when the total loss function does not reach the preset convergence condition, iteratively update the parameters until the total loss function reaches the preset convergence condition to obtain the segmentation model; (6) Target image segmentation: For a given medical image to be tested, the segmentation model outputs the probability of the category to which each pixel point of the medical image to be tested belongs, and sets a threshold to mark the pixel point as a segmentation target or background.
2. The single-domain generalization method for medical image segmentation according to claim 1, wherein In step (1), the sizes of the labeled medical images x in the unified training domain are unified, and data augmentation is performed to obtain augmented data x ′ , which specifically includes: (i) Use the bilinear interpolation algorithm to scale all images x to a fixed size, and use the nearest neighbor difference to scale the corresponding label y to the same size so that the input image meets the input specifications of the segmentation network; (ii) Obtain the augmented data x by data augmentation on the uniformly sized image x ′ to simulate domain shift.
3. The single-domain generalization method for medical image segmentation according to claim 1, wherein In step (3), the segmentation loss is calculated specifically as follows: where H and W represent the input image size, i represents the index of the pixel, C represents the class of the pixel prediction in the segmentation, c is the index of the class, and y ic ∈ {0, 1} is the one-hot encoding of the ground truth label, and p ic is the predicted class probability.
4. The single-domain generalization method for medical image segmentation according to claim 1, wherein In step (3), the consistency loss is calculated specifically as follows: Among them, f(·) represents the segmentation network composed of an encoder and a decoder.
5. The single-domain generalization method for medical image segmentation according to claim 1, wherein In step (4), the difference operation is performed on the feature maps M and M' to obtain the highly sensitive style information feature map M s , which specifically includes: Perform a difference operation on the feature map M of the original image and the feature map M' of the enhanced image to obtain a highly sensitive style information feature map M s = M - M'.
6. The single-domain generalization method for medical image segmentation according to claim 1, characterized in that In step (4), dimensionality reduction is performed through global average pooling and normalization, specifically including: where H and W represent the size of the style information feature map M s , c is the channel number index, i, j represent the position index, and Max(·) and Min(·) are used to calculate the maximum and minimum values in the channel dimension.
7. The single-domain generalization method for medical image segmentation according to claim 1, wherein In step (4), the whitening mask WM is obtained through threshold selection, specifically including: Where i and j represent position indices, the top(·) function represents returning the value of the top element in the sensitivity map SM, and ε is a fixed selected whitening range threshold, set to 0.1 - 0.
9.
8. The single-domain generalization method for medical image segmentation according to claim 1, wherein In step (4), for the covariance matrices ∑ M and ∑′ M of the feature maps M and M′, calculate the feature whitening loss Specifically, it includes: 4.1) The covariance matrix calculation method is: Where H and W represent the sizes of the feature maps M and M′; 4.2) Whitening loss The calculation method is: Among them represents the arithmetic mean operation, and ⊙ represents element-wise multiplication.
9. The single-domain generalization method for medical image segmentation according to claim 1, characterized in that In step (5), when training the encoder Encoder and the decoder Decoder, when the total loss function does not reach the preset convergence condition, iteratively update the parameters until the total loss function reaches the preset convergence condition to obtain the segmentation model, specifically including: (i) Computed total loss is the segmentation loss consistency loss and the feature whitening loss as a linear combination; Among them, λ1 and λ2 are hyperparameters, and the ranges of the λ1 and λ2 parameters are (0 - 1], which are used to balance the influence of the two losses on the total loss; (ii) Iteratively update the parameters of the encoder Encoder and the decoder Decoder until the total loss function reaches the preset convergence condition to obtain the segmentation model.
10. The single-domain generalization method for medical image segmentation according to claim 1, wherein In step (6), for a given medical image to be tested, the segmentation model outputs the probability of the category to which each pixel point of the medical image to be tested belongs, and sets a threshold to mark the pixel point as a segmentation target or background, specifically including: (i) Input the medical image x to be tested T into the segmentation model to obtain the predicted probability map f(x T ); (ii) Set 0.5 - 0.9 as the threshold to mark the prediction map. If the prediction probability is greater than the threshold, it is marked as a segmentation target, and if it is less than the threshold, it is marked as the background to obtain the final segmentation mask.
Citation Information
Patent Citations
Fundus blood vessel segmentation domain generalization method based on frequency domain analysis
CN115018874A