3D all-core nuclear magnetic image segmentation method based on semi-supervision

By adopting improved teacher-student model and FMDD-UNet network structure in 3D full-heart nuclear magnetic image segmentation, combined with dual decoder and global category prototype training, the shortcomings of the existing semi-supervised methods in fine processing and global modeling capabilities are solved, and higher segmentation accuracy and better category balance are achieved.

CN120047455APending Publication Date: 2025-05-27HEBEI UNIV OF TECH
View PDF 0 Cites 0 Cited by

Patent Information

Application Number
CN202510114556.1
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2025-01-24
Publication Date
2025-05-27

AI Technical Summary

Technical Problem

When dealing with 3D full-heart nuclear magnetic image segmentation, the existing semi-supervised methods have problems such as insufficient fine processing, weak global modeling capabilities and imbalance in categories.

Method used

A semi-supervised 3D full-heart nuclear magnetic image segmentation method is designed, using an improved teacher-student model segmentation network model, combining the FMDD-UNet network structure and dual decoder, pre-training and global category prototype training is used with labeled data, and finally deep training is carried out to improve segmentation accuracy.

Benefits of technology

Through this method, the segmentation accuracy of 3D full-heart nuclear magnetic images can be improved while using a large amount of unlabeled data, relieve the problem of category imbalance, and show the effect comparable to that of the supervision segmentation method in clinical diagnosis.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN120047455A_ABST
    Figure CN120047455A_ABST
Patent Text Reader

Abstract

The invention discloses a semi-supervision-based 3D all-heart nuclear magnetic image segmentation method, and the method employs an improved teacher-student model segmentation network model, the student model employs a three-dimensional U-shaped network structure based on a Mama encoder to extract the global and local information of an image in a 3D medical image, and employs a dual decoder and a segmentation decoder to output a segmentation result. A reconstruction decoder reconstructs an output feature of the encoder. A teacher-student model is trained in a semi-supervised mode, firstly, the student model is pre-trained by using labeled data, then, further training is performed in combination with a global category prototype of the labeled data, and finally, deep training is performed on the teacher-student model by using the labeled data and unlabeled data. And network parameters of the teacher model are updated through the index moving average value of the student model. The performance of the semi-supervised segmentation method is equivalent to that of a supervised segmentation method, and the semi-supervised segmentation method has important significance in clinical medical diagnosis.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention relates to the technical field of image segmentation, and in particular to a semi-supervised 3D whole-heart nuclear magnetic resonance image segmentation method. Background Art

[0002] Cardiovascular disease has always been a difficult and hot research topic in the medical field, and people are paying more and more attention to it. With the advancement of modern medical technology and auxiliary technologies in other fields, the structure and function of the heart can be qualitatively and quantitatively evaluated through technologies such as magnetic resonance imaging (MRI), computed tomography (CT) and ultrasound imaging. Segmenting and determining several substructures of the heart is a basic step in disease diagnosis and surgical planning. Accurately segmenting the substructures of the heart can help medical workers track the target organs more easily, so as to make more accurate medical diagnoses. However, the current segmentation of medical images faces many problems and challenges, specifically the following two aspects:

[0003] 1) Difficulty in labeling: Since medical images are usually 3D images composed of several 2D slices, which requires a lot of professional knowledge, medical image labeling often requires doctors and other professionals to label. Therefore, labeling data is very time-consuming and labor-intensive. Therefore, how to use a large amount of unlabeled data for medical image processing is a direction with great research significance and application value.

[0004] 2) Medical images are more complex: Compared with natural images, medical images usually show blurred structures, low contrast between different regions of the image, and noisy and chaotic characteristics. Therefore, it is necessary to design a reasonable method to obtain detailed spatial structural information of the image. For example, the pixel values ​​of the myocardium and liver parts of the heart substructure are very similar and difficult to distinguish with the naked eye; at the same time, in MRI images, the chest, back, neck and other parts often introduce additional noise, which weakens the segmentation effect of the model.

[0005] With the development of deep learning, a large number of model structures based on CNN convolutional structures and Transformer attention mechanisms have emerged. CNN and its variants can effectively capture local and detailed spatial features because convolution uses a sliding mechanism to operate on images. At the same time, the sharing mechanism of convolution makes the model parameters small. However, due to its mechanism, the ability to process global features is insufficient. Transformer and its variants are generally based on the attention mechanism and can model global information. However, due to the sequence-based calculation of attention, the Transformer has a large number of parameters, high computational complexity, and consumes a lot of computing resources. Therefore, it is extremely challenging to achieve local and global modeling and a model with a reasonable number of parameters. At the same time, in order to utilize a large amount of unlabeled data, Huang et al. solved some of the labeling problems by generating pseudo labels for unlabeled anatomical structures based on a semi-supervised learning paradigm, and combined additional regularization such as anatomical structure size and consistency between models to stabilize training. However, without fine spatial structure information, the model will generate inaccurate pseudo labels, further damaging the performance of the model. Summary of the invention

[0006] In order to solve the problems of insufficient processing, weak global modeling ability and easy category imbalance in existing semi-supervised methods, a semi-supervised 3D whole heart nuclear magnetic resonance image segmentation method was provided to improve the accuracy of segmenting 3D whole heart nuclear magnetic resonance images.

[0007] The technical solution of the present invention to solve the technical problem is to design a semi-supervised 3D whole heart nuclear magnetic resonance image segmentation method, characterized in that the method comprises the following steps:

[0008] Step 1: Preprocess the 3D full heart image data to obtain the data set to be trained

[0009] The training samples in the dataset are obtained by randomly cropping the 3D full heart image data to a set size, adding random offset values ​​to increase the diversity of the data, and finally performing z-score standardization on the data; the dataset consists of L+U training samples, of which L are labeled data and U are unlabeled data;

[0010] Step 2: Build a segmentation network model

[0011] The overall network architecture of the segmentation network model adopts a teacher-student model. The student model and the teacher model have the same network structure. The student model is an FMDD-UNet network structure improved from 3D UNet, specifically including an FMmamba-Encoder encoder, a segmentation decoder Decider seg and reconstruct the decoderrec , wherein the FMmamba-Encoder encoder is composed of a first convolution block, a second convolution block, a third convolution block, a fourth convolution block, a fifth convolution block, a first Freq-Mamba Block, a second Freq-Mamba Block, and a third Freq-Mamba Block;

[0012] The input image X of the student model is first processed by the first convolution block, the second convolution block, the third convolution block, the fourth convolution block, and the fifth convolution block in sequence. The output features of the previous convolution block are used as the input features of the next convolution block. The output features of the first convolution block, the second convolution block, the third convolution block, the fourth convolution block, and the fifth convolution block are recorded as F 01 、F 02 、F 03 、F 04 、F 05 ;

[0013] After the input feature X is input to the first convolution block, a convolution operation with a convolution kernel of 1×1×1 and a stride of 1 is first performed to obtain y, y=F(X)+X, where F represents the above convolution operation; y is a 5D tensor of size b×c×h×w×d, where b represents the batch size, h represents the spatial height, w represents the spatial width, d is the spatial depth, and c is the number of channels; then two convolution operations are performed on y. The first convolution operation is a convolution operation with a convolution kernel of 3×3×3 and a stride of 2. The number of channels is doubled from c. c', and InstanceNorm normalization and ReLU activation are performed to obtain the intermediate layer features, and then a second convolution operation is performed; the second convolution operation is also a convolution operation with a convolution kernel of 3×3×3 and a step size of 2, the number of channels c' remains unchanged, and InstanceNorm normalization and ReLU activation are performed; the final result of the second convolution is subjected to a MaxPool maximum pooling of size 2 to obtain a 5D tensor of size h / 2×w / 2×d / c', which is F 01 ; The operations inside the second convolution block, the third convolution block, the fourth convolution block, and the fifth convolution block are exactly the same as those in the first convolution block. The difference is that the number of output channels of the latter convolution block is twice the number of output channels of the previous convolution block;

[0014] The output feature F of the fifth convolutional block 05As the input of the first Freq-Mamba Block, the output of the first Freq-MambaBlock is used as the input of the second Freq-Mamba Block, the output of the second Freq-Mamba Block is used as the input of the third Freq-Mamba Block, and the output feature of the third Freq-Mamba Block is recorded as F 06 ; The basic structure, number of feature channels, etc. of the first Freq-Mamba Block, the second Freq-Mamba Block, and the third Freq-Mamba Block are exactly the same, and the parameters are not shared; The first Freq-Mamba Block includes a wavelet transform operation, two Mamba modules, an inverse transform operation of a wavelet transform, a BN operation, and a ReLu operation; The feature F 05 After being input into the first Freq-Mamba Block, it is processed in two paths; for the first path, the feature F 05 First, after a wavelet transform operation, the result is input into the first Mamba module. The output of the first Mamba module undergoes an inverse wavelet transform operation to obtain the feature FM 1 ; For the second path, feature F 05 After being processed by the second Mamba module, the feature FM is obtained 2 ; Then the feature FM 1 With the characteristic FM 2 The concatenation is performed in the channel dimension, and the obtained result is then subjected to BN operation and ReLu operation in sequence to obtain the output of the first Freq-Mamba Block;

[0015] The first Mamba module has the same basic structure as the second Mamba module, but different parameters. The first Mamba module is used as an example to explain the principle. The first Mamba module includes a compression layer, a normalization layer, a first linear layer, a second linear layer, a first activation function layer, a one-dimensional convolution layer, a second activation function layer, a state space model, a third linear layer, and a reshape function layer. For the F input to the first Mamba module, a , the compression layer first changes its dimension from [b,c,h,w,d] to [b,c,h×w×d], and then after the normalization layer processing, we get F b ; Then F b The first path is processed through a linear layer, a one-dimensional convolutional layer, a SiLu activation function, and a state space model, and outputs a feature F of [b, c, h×w×d] dimensions. c; The second path is processed by a linear layer and SiLu activation function, outputting the feature F of [b,c,h×w×d] dimensions d ; Output F of the first path c and the output F of the second path d Multiply them together, and then process them through a linear layer to output the feature F of [b,c,h×w×d] dimensions e , F e After being processed by the reshape function layer, the output F of [b,c,h,w,d] dimensions is finally obtained f ;

[0016] Segmentation Decoder seg and reconstruct the decoder rec The basic structure of the two networks is the same, but the parameters are not shared. They all include the first convolution upsampling module, the second convolution upsampling module, the third convolution upsampling module, the fourth convolution upsampling module and the Final Conv module.

[0017] For the segmentation decoder seg , the feature F 04 With feature F 06 As the input of the first convolution upsampling module, the first convolution upsampling module outputs the feature F 11 ; The feature F 03 With feature F 11 As the input of the second convolution upsampling module, the second convolution upsampling module outputs the feature F 12 ; The feature F 02 With feature F 12 As the input of the third convolution upsampling module, the third convolution upsampling module outputs the feature F 13 ; The feature F 01 With feature F 13 As the input of the fourth convolution upsampling module, the fourth convolution upsampling module outputs the feature F 14 ; The feature F 14 As a segmentation decoder Decoder seg The Final Conv module first performs a convolution operation with a kernel of 3×3×3 and a step size of 1, converts the number of channels to proto_dim, and outputs the secondary segmentation result out sub_seg Then perform a second convolution operation with a kernel of 1 and a step size of 1, convert proto_dim into the number of segmented categories, and then perform the secondary segmentation result out sub_seg Perform BatchNorm normalization and ReLU activation processing to output the final segmentation result out seg ;

[0018] Segmentation Decoder seg The first convolution upsampling module receives the input feature F 04 With feature F 06 , for F 06 First, perform an upsampling operation using trilinear interpolation mode with a scaling factor of (2,2,2) to obtain the upsampled intermediate result mid_outputs01; then calculate mid_outputs1 and F 04 The size difference in the spatial dimension is used to determine the corresponding filling amount, and the F 04 Fill the intermediate result mid_outputs02; then concatenate mid_outputs01 and mid_outputs02 in the channel dimension; finally, input the concatenated result into the convolution module consisting of two layers of convolution operations. The first convolution operation is a convolution operation with a convolution kernel of 3×3×3 and a step size of 1. The number of channels is determined by F 04 The number of channels c is 0.5 times the original one, that is, c', and InstanceNorm normalization and ReLU activation are performed to obtain the intermediate layer features, and then the second convolution operation is performed; the second convolution operation is also a convolution operation with a convolution kernel of 3×3×3 and a step size of 1. The number of channels c' remains unchanged, and InstanceNorm normalization and ReLU activation are performed, and finally the output feature F of the first convolution upsampling module is output 11 ;

[0019] Segmentation Decoder seg The operations inside the second convolution upsampling module, the third convolution upsampling module, and the fourth convolution upsampling module are exactly the same as those in the first convolution upsampling module. The difference is that the number of output channels is different. The number of output channels of the latter convolution block is 0.5 times the number of output channels of the previous convolution block. The decoder is reconstructed. rec The first convolution upsampling module receives two input features F 04 With feature F 06 , for F 06 First, perform an upsampling operation using trilinear interpolation mode with a scaling factor of (2,2,2) to obtain the upsampled intermediate result mid_outputs11; then calculate mid_outputs11 and F 04 The size difference in the spatial dimension is used to determine the corresponding filling amount, and the F 04Pad to get the intermediate result mid_outputs12; then concatenate mid_outputs11 and mid_outputs12 in the channel dimension; finally, input the concatenated result into the convolution module consisting of two layers of convolution operations. The first convolution operation is a convolution operation with a convolution kernel of 3×3×3 and a step size of 1. The number of channels is determined by F 04 The number of channels c is 0.5 times the original one, that is, c', and the InstanceNorm normalization kernel ReLU activation is performed to obtain the intermediate layer features for the second convolution operation; the second convolution operation is also a convolution operation with a convolution kernel of 3×3×3 and a step size of 1. The number of channels c' remains unchanged, and InstanceNorm normalization and ReLU activation are performed to finally output the output features of the first convolution upsampling module;

[0020] Reconstruction decoder Decoder rec The operations inside the second convolution upsampling module, the third convolution upsampling module, and the fourth convolution upsampling module are exactly the same as those in the first convolution upsampling module. The difference is that the number of output channels is different. The number of output channels of the latter convolution block is 0.5 times the number of output channels of the previous convolution block. The output of the fourth convolution upsampling block is input to the reconstruction decoder Decoder rec The Final Conv module first performs the first convolution operation on the input with a convolution kernel of 3×3×3 and a step size of 1, and the number of channels is converted to proto_dim. Then, the second convolution operation is performed to convert proto_dim to 1. The output of the first convolution operation is then BatchNorm normalized and ReLU activated to obtain the final reconstruction result out rec ;

[0021] Step 3: Pre-training of student models

[0022] The student model is pre-trained using the labeled data in the dataset. The pre-training phase is terminated after 1000 iterations. The initial network parameters are initialized using the default Kaiming. Random sampling is performed to input a batch of labeled data in the dataset into the initialized student model. The segmentation results are used to seg Calculate the linear classification loss with the Ground-Truth segmentation label; reconstruct the decoder output reconstructed result out rec The original image input to the student model is fed into the perceptual loss network and the perceptual loss is calculated; the loss function to be optimized is as follows:

[0023] L cls_ce_3d =L ce (f ds (f e (xi )),y i ) (1)

[0024] L seg_dice_3d =L dice (softmax(f ds (f e (x i ))),y i ) (2)

[0025] L cls_3d =weight ce_w ×L cls_ce_3d +weight dice_w ×L seg_dice_3d (3)

[0026] L percen_3d =L 1 (vgg(f dr (f e (x i )),vgg(x i ))) (4)

[0027] L a =L cls_3d +L percen_3d (5)

[0028] Among them, x i represents input, f e (·) represents the output of FMmamba-Encoder, f ds (·) represents the segmentation decoder Decoder seg The output, f dr (·) represents the reconstruction decoder rec The output of , vgg(·) represents the output of the perceptual loss network; y i Represents x i Ground-Truth segmentation label; weight ce_w 、weight dice_w They represent the cross entropy loss weight and Dice loss weight, both of which are 0.5; L ce (·,·) and L dice (·,·) means using cross entropy loss and Dice loss, multiplying by weights and summing them up to get the linear classification loss L cls_3d ; Perceptual loss L percen_3d Use L 1 The loss is calculated and finally summed with the linear classification loss to obtain the training loss of a batch of data in the pre-training stage of the student model;

[0029] According to the training loss, the stochastic gradient descent optimizer is used to reversely update the network parameters of the student model once to complete the training of one batch of data; the network parameters that have completed the training of the previous batch of data are used as the initial parameters of the next batch of training, and the next batch of labeled data is input, and the process of training a batch of data is repeated continuously until the training of the last batch of labeled data in the data set is completed, completing one round of training; the network parameters that have completed the previous round of training are used as the initial parameters of the next round of training, and the training is continuously iterated until the iteration round reaches 1000 times, and the pre-trained student model is obtained;

[0030] Step 4: Global Category Prototype Training of Student Model

[0031] Step 4.1: First, the global category prototype P 0 Set to all zero parameters of [num_classes, subcluster, proto_dim] dimensions, where num_classes, subcluster, and proto_dim represent the number of segmented classes, the number of subcluster centers in a single class prototype, and the dimension of the global class prototype, respectively;

[0032] Random sampling, input a batch of labeled data in the dataset into the student model pre-trained in the third step, split decoder Decoder seg Output segmentation result out seg , secondary segmentation result out sub_seg , the reconstruction decoder outputs the reconstruction result out rec ;

[0033] For the secondary segmentation result out sub_seg Perform a dimension transformation from [b,c,h,w,d] to [b×h×w×d,c] to obtain out_feat; the ground truth is first transformed from [b,h,w,d] to [b×d,h,w] to obtain label_2d, and label_2d is transformed to [b×h×w×d] to obtain label_expand; then, from out_feat, according to its corresponding label_expand category label, filter out the feature subsets belonging to each category in turn; the specific operation is: according to the position of the pixel of the kth category of the label_expand category label, that is, the corresponding pixel value is k, filter and collect the features of the corresponding position in out_feat; for the collected features of the kth category, use K-means clustering to obtain the small category prototype corresponding to this category; then concatenate all the small category prototypes into a tensor to obtain the initialized global category prototype P 1, the dimension is [num_classes,subcluster,proto_dim];

[0034] Step 4.2: The secondary segmentation result out obtained in step 4.1 sub_seg First, we transform the dimension from [b,c,h,w,d] to [b×h×w×d,c] to get out_feat; the ground truth is first transformed from [b,h,w,d] to [b×d,h,w] to get label_2d, which is then transformed to [b×h×w×d] through label_expand; out_feat and the current global category prototype P are combined. 1 Perform matrix multiplication along the last dimension and output feat_proto_sim with the dimension [b×h×w×d,subcluste,num_classes]; then calculate the maximum value tmp of feat_proto_sim in the subcluster dimension, and the maximum value tmp is regularized by Layernorm and the dimension is changed to [b×d,num_classes,h,w], so as to obtain the secondary segmentation result out sub_seg The global category prototype P 1 The segmentation result proto_seg, that is, nearest_proto_distance;

[0035] Then use out_feat, feat_proto_sim, label_expand, and nearest_proto_distance to update the momentum of the global category prototype; first, search the nearest_proto_distance on the num_classes dimension to get the pred_seg with the maximum value index dimension of [b×d,h,w]; secondly, stretch label_expand and pred_seg to pred_seg of [b×d×h×w] s Perform query calculation to obtain the Boolean matrix mask of the positions where the elements of the two are equal;

[0036] Then, starting from category 0, we traverse num_classes categories and perform momentum updates on the small category prototypes of each category. Taking the update of the kth category as an example, the update method for each category is as follows;

[0037] First, find the index of feat_proto_sim according to the last dimension, the index value is category k, and get the feature tmp k_0; Then find the position of the pixel value equal to k in label_expand, record it as position index_k, and use this position in tmp k_0 Find the index of the first dimension and get the similarity measure q k ;

[0038] Then the similarity measure q k Perform a sinkhorn operation to obtain the final similarity measure q; then find the feature m_k of the Boolean matrix mask at the index_k position, copy m_k to expand it subcluster times in the column dimension, and obtain the mask matrix m k_0 ; Then find the index in the first dimension of the feature out_feat, find the feature m_w at the index_k position, copy m_w to expand it proto_dim times in the column dimension, and get the mask matrix m k_1 ;

[0039] Then calculate the updated value f of the small category prototype of the current class k , the formula is:

[0040] Where m q_k =qm k_0 , represents the effective soft assignment weight, w q_k =m_k×m k_1 ; Then, the momentum update formula is used to update the newly calculated update value f k Incorporating old category archetypes In the example above, we get the updated small category prototype

[0041]

[0042] where μ is the momentum coefficient; the update value f k Maintain consistency through l2-normalization, that is, the normalize operation; perform the above operation on each small category prototype to obtain each updated small category prototype, and then splice each updated small category prototype to obtain an updated global category prototype P 2 ;

[0043] Step 4.3: Decoder according to the segmentation in step 4.1 seg Output segmentation result out seg , secondary segmentation result out sub_seg , the reconstruction decoder outputs the reconstruction result out rec , calculate the training loss, the loss function is as follows:

[0044] L proto_ce =Lce (Proto(f ds_sub (f e (x i ))),toSlice(y i )) (6)

[0045] L proto_seg_dice =L dice (Softmax(Proto(f ds_sub (f e (x i )))),toSlice(y i )) (7)

[0046] L proto_3d =weight ce_w ×L proto_ce +weight dice_w ×L proto_se_dicd (8)

[0047] weight ce_w 、weight dice_w is the weight matrix parameter; according to formula (3), formula (4), and formula (8), the total training loss L of the batch of labeled data is obtained b , the specific calculation formula is:

[0048] L b =L proto_3d +L percen_3d +L cls_3d (9)

[0049] L percen_3d , L cls_3d According to formula (4) and formula (3), the perceptual loss and linear classification loss are calculated respectively. proto_3d is the current global category prototype P 2 loss;

[0050] For formula (6), Proto(f ds_sub (f e (x i ))) refers to the secondary segmentation result f output by the segmentation decoder ds_sub (f e (x i )))'s current global category prototype P 2 The segmentation result has the dimension [b×d,num_classed,h,w]; toSlice(y i ) means to convert the label y of shape [b,h,w,d] i Map to [b×d,h,w];

[0051] For formula (7), Proto(f ds_sub (f e (x i )))、toSlice(y i ) is consistent with the operation in formula (6), Softmax(Proto(f ds_sub (f e (x i )))) refers to the secondary segmentation result f output by the segmentation decoder ds_sub (f e (x i ))'s current global category prototype P 2 Perform Softmax operation on the segmentation result and adjust the dimension to [b×d,num_classes,h,w];

[0052] According to the total loss L of a batch b , using the stochastic gradient descent optimizer, the trainable parameters of the student model are updated once, completing the iterative training of a batch of labeled data;

[0053] Step 4.4: Input the next batch of labeled data in the data set into the pre-trained student model that has completed one iteration of training in step 4.3, and repeat the process of steps 4.2, 4.3, and 4.4 until the last batch of labeled data in the data set is trained and one round of training is completed; use the network parameters that have completed the previous round of training as the initial parameters for the next round of training, and repeat the training process for one round until the iterative round reaches 1000 times, complete the global category prototype training of the student model, and obtain the current global category prototype P o ;

[0054] Step 5: Deep training of segmentation network model

[0055] Step 5.1: The network parameters of the student model that has completed the global category prototype training in the fourth step are used as the initial network parameters of the deep training stage of the segmentation network model. The initial values ​​of the parameters of the teacher model are all set to 0. The parameter values ​​of the teacher model are initialized according to the initial network parameter values ​​of the student model. Specifically, in the mean teacher model, the student model weight θ S The exponential moving average of is used to update the weights θ of the teacher model T ;

[0056] Step 5.2: Randomly sample and select a batch of training data from the dataset. A batch of training data includes half labeled data and half unlabeled data. Input the labeled data in a batch into the student model that has completed the global category prototype training in step 4. The student model outputs the segmentation result of the labeled data out seg_l , secondary segmentation result out sub_seg_l And the reconstruction result out rec_l ;

[0057] According to the output of the student model for the labeled data, we first use the method in step 4.2 to generate the global category prototype P o Based on this, the global category prototype is updated to obtain the current global category prototype; then the segmentation result out seg_l , secondary segmentation result out sub_seg_l And the true label GroundTruth respectively calculates the linear classification loss L of the labeled data cls_3d_l and the current global category prototype loss L proto_3d_l , the reconstruction result out rec_l The labeled image of the input model is input to the perceptual loss network and the perceptual loss L is calculated. percen_3d_l , linear classification loss L cls_3d_l , the current global category prototype loss L proto_3d_l , perceptual loss L percen_3d_l The sum is the training loss L for labeled data l ;

[0058] Then the unlabeled data in a batch is input into the student model that has completed the global category prototype training and the teacher model that has completed the network parameter initialization, and the segmentation result output by the teacher model is output. seg_ul As the pseudo label of the corresponding input unlabeled data when the student model calculates the loss; after the unlabeled data is input to the student model, the segmentation result of the student model is used seg_ul , secondary segmentation result out sub_seg_ul And the pseudo label Pseudo Label generated by the teacher model is used to calculate the corresponding linear classification loss L cls_3d_ul and the current global category prototype loss L proto_3d_ul , the reconstruction result out rec_ul The image of the input teacher model is input to the perceptual loss network and the perceptual loss L is calculated percen_3d_ul , linear classification loss L cls_3d_ul , the current global category prototype loss L proto_3d_ul , perceptual loss L percen_3d_ul The sum is the loss of unlabeled data L ul ;

[0059] The above linear classification loss, perceptual loss and current global category prototype loss are calculated according to formula (3), formula (4) and formula (8) respectively; for a batch of training data, the training loss is L;

[0060] L=L l +consistency_weight×L ul

[0061] In the above formula, consistency_weight is the consistency weight of the teacher-student model, which is updated with the number of training rounds. According to the training loss of the training data of this batch, the network parameters of a student model are reversely updated. At the same time, the network parameters of the teacher model are updated once by using the exponential moving average EMA using the updated network parameters of the student model to complete the training of a batch of data samples.

[0062] Step 5.3: Select training set D train The next batch of training data in the dataset is taken. Similarly, the training data of this batch includes half of the labeled data and half of the unlabeled data. The network parameters of the student model and the teacher model when the previous batch of data samples are trained are used as the initial parameters for the next batch of training data. The process in step 5.2 and step 5.3 is repeated continuously until the last batch of training data in the dataset, including half of the labeled data and half of the unlabeled data, is trained to complete one round of training. The network parameters at the completion of the previous round of training are used as the initial network parameters at the beginning of the next round of training. The training round is repeated continuously until the preset 14,000 times are reached to obtain a trained segmentation network model.

[0063] Step 6: 3D whole heart MRI image segmentation

[0064] First, the patient's 3D whole-heart MRI image data is obtained, and then the 3D image data is center-cropped to obtain a picture of the same size as set in the first step, and the pixels of the picture are z-score standardized. After that, it is input into the student model in the segmentation network model trained in the fifth step to obtain the segmentation result, thus completing the segmentation of the 3D whole-heart MRI image data.

[0065] Compared with the prior art, the beneficial effects of the present invention are as follows: a 3D whole heart magnetic resonance image segmentation method based on semi-supervision adopts an improved teacher-student model segmentation network model, the teacher model and the student model adopt the same structure, the student model adopts a three-dimensional U-shaped mesh structure based on the Mamba encoder to extract the global and local information of the image in the 3D medical image, and adopts a dual decoder, a segmentation decoder outputs the segmentation result, and a reconstruction decoder is used to reconstruct the output features of the encoder. The teacher-student model is trained in a semi-supervised manner, firstly the student model is pre-trained with labeled data, then further trained in combination with the global category prototype of the labeled data, and finally the teacher-student model is deeply trained with labeled data and unlabeled data, the network parameters of the teacher model are updated by the exponential moving average of the student model, and the final trained segmentation network model is obtained. In the pre-training stage, the sum of linear classification loss and perceptual loss is used to update the network parameters. In the subsequent two stages, the sum of linear classification loss, perceptual loss, and current global category prototype loss is used to update the network parameters, instead of using a pixel-by-pixel loss function that only relies on low-level pixel information, so as to constrain the reconstructed image that deviates from the input image semantically and spatially, obtain a more refined spatial structure from the input image, and improve the segmentation accuracy of the whole heart nuclear magnetic resonance image. The method of the present invention can use a large amount of unlabeled data to refine and constrain the learning of the model, alleviate the problem of category imbalance while obtaining the fine spatial structure of the image, and can alleviate the lack of labeled data in clinical practice. The information of a large amount of unlabeled data is used to improve the accuracy of the model for whole heart segmentation, and the performance of the semi-supervised segmentation method is comparable to that of the supervised segmentation method. Accurate segmentation of cardiac substructures is an important prerequisite for quantitative description and surgical planning of cardiovascular diseases. The method of the present invention can play a role in assisting doctors in clinical diagnosis and treatment, and is of great significance in clinical medical diagnosis. BRIEF DESCRIPTION OF THE DRAWINGS

[0066] Figure 1 This is a network structure diagram of a student model of an embodiment of a semi-supervised 3D whole-heart magnetic resonance image segmentation method of the present invention.

[0067] Figure 2 The schematic diagram of the principle of the state space model of the Mamba module of the student model of an embodiment of the semi-supervised 3D whole heart magnetic resonance image segmentation method of the present invention.

[0068] Figure 3 A schematic diagram of the structure and principle of the Freq-Mamba Block of the student model of an embodiment of a semi-supervised 3D whole-heart magnetic resonance image segmentation method of the present invention.

[0069] Figure 4Schematic diagram of the structure of the 3D UNet network in the prior art.

[0070] Figure 5 The figure is a schematic diagram of the structure and principle of a perceptual loss network used in an embodiment of a semi-supervised 3D whole-heart magnetic resonance image segmentation method of the present invention. DETAILED DESCRIPTION

[0071] The specific embodiments of the present invention are given below. The specific embodiments are only used to further illustrate the present invention in detail and do not limit the protection scope of the claims of this application.

[0072] The present invention provides a 3D whole heart nuclear magnetic resonance image segmentation method based on semi-supervision (hereinafter referred to as the method), which comprises the following steps:

[0073] Step 1: Preprocess the 3D full heart image data to obtain the data set to be trained

[0074] The training samples in the dataset are obtained by randomly cropping the 3D full heart image data to a set size (80,80,80), adding random offset values ​​to increase the diversity of the data, and finally performing z-score standardization on the data. The dataset consists of L+U training samples, of which L are labeled data and U are unlabeled data. The labeled dataset is represented as The unlabeled dataset is represented as For 3D data, x i ∈R H×W×D , R H×W×D Indicates the dimension of the input data; y i ∈{0,1,2…5} H×W×D ,y i is the Ground-Truth segmentation label, and H, W, and D are the height, width, and thickness of the input image, respectively. As an embodiment, the present invention uses the MMWHS 3D whole heart image dataset for preprocessing to obtain a dataset to be trained. This dataset contains 7 substructures of the heart, including the myocardium (Myo), left atrium (LA), left ventricle (LV), right atrium (RA), right ventricle (RV), ascending aorta (AA), and pulmonary artery (PA). The present invention mainly targets the left and right four chambers and myocardium (i.e., LA, LV, RA, RV, Myo) in the cardiac substructures, totaling 5 categories. The preprocessing specifically refers to: randomly cropping the 3D dataset to a size of (80, 80, 80), adding random offset values ​​to increase the diversity of the data, and finally performing z-score standardization on the data.

[0075] Step 2: Build a segmentation network model

[0076] The overall network architecture of the segmentation network model adopts a teacher-student model, and the student model and the teacher model have the same network structure. The student model is an FMDD-UNet network structure improved from 3D UNet (see Appendix Figure 2 ), including FMmamba-Encoder encoder, segmentation decoder Decoder seg and reconstruct the decoder rec , wherein the FMmamba-Encoder encoder consists of a first convolution block, a second convolution block, a third convolution block, a fourth convolution block, a fifth convolution block, a first Freq-Mamba Block, a second Freq-Mamba Block, and a third Freq-Mamba Block.

[0077] The input image X of the student model is first processed by the first convolution block, the second convolution block, the third convolution block, the fourth convolution block, and the fifth convolution block in sequence. The output features of the previous convolution block are used as the input features of the next convolution block. The output features of the first convolution block, the second convolution block, the third convolution block, the fourth convolution block, and the fifth convolution block are recorded as F 01 、F 02 、F 03 、F 04 、F 05 .

[0078] After the input feature X (dimension is [b, 1, 80, 80, 80]) is input into the first convolution block, a convolution operation with a convolution kernel of 1×1×1 and a stride of 1 is first performed to obtain y, y=F(X)+X, where F represents the above convolution operation; y is a 5D tensor of size b×c×h×w×d, where b represents the batch size, h represents the spatial height, w represents the spatial width, d is the spatial depth, and c is the number of channels; then y is subjected to two convolution operations, the first convolution operation is a convolution operation with a convolution kernel of 3×3×3 and a stride of 2, the number of channels is changed from c to twice the original, which is c' (c'=2c), and InstanceNorm normalization and ReLU activation processing are performed to obtain the intermediate layer features, and then a second convolution operation is performed; the second convolution operation is also a convolution operation with a convolution kernel of 3×3×3 and a stride of 2, the number of channels c' remains unchanged, and InstanceNorm normalization and ReLU activation processing are performed. The final result of the second convolution is pooled with a MaxPool of size 2 to obtain a 5D tensor of size h / 2×w / 2×d / c', which is F 01The internal operations of the second convolution block, the third convolution block, the fourth convolution block, and the fifth convolution block are exactly the same as those of the first convolution block. The difference is that the number of output channels is different. The number of output channels of the latter convolution block is twice the number of output channels of the previous convolution block. As an embodiment, the number of output channels of the first convolution block, the second convolution block, the third convolution block, the fourth convolution block, and the fifth convolution block are 16, 32, 64, 128, and 256, respectively, and the image dimension transformations are 80, 40, 20, 10, and 5 respectively.

[0079] The output feature F of the fifth convolutional block 05 As the input of the first Freq-Mamba Block, the output of the first Freq-MambaBlock is used as the input of the second Freq-Mamba Block, the output of the second Freq-Mamba Block is used as the input of the third Freq-Mamba Block, and the output feature of the third Freq-Mamba Block is recorded as F 06 The basic structure, number of feature channels, etc. of the first Freq-Mamba Block, the second Freq-Mamba Block, and the third Freq-Mamba Block are exactly the same, and parameters are not shared. As an embodiment, the number of output channels of the first Freq-Mamba Block, the second Freq-Mamba Block, and the third Freq-Mamba Block are all 256, and the image dimensions are all 10.

[0080] The first Freq-Mamba Block includes a wavelet transform operation, two Mamba modules, an inverse transform operation of a wavelet transform, a BN (batch normalization operation) operation, and a ReLu operation (ReLu activation function processing operation); the feature F 05 After being input into the first Freq-Mamba Block, it is processed in two paths; for the first path, the feature F 05 First, a wavelet transform operation (FFT) is performed, and the result is input into the first Mamba module. The output of the first Mamba module is subjected to an inverse wavelet transform operation to obtain the feature FM 1 ; For the second path, feature F 05 After being processed by the second Mamba module, the feature FM is obtained 2 ; Then the feature FM 1 With the characteristic FM 2 The concatenation is performed in the channel dimension, and the result is then subjected to BN and ReLu operations in sequence to obtain the output of the first Freq-Mamba Block.

[0081] The first Mamba module has the same basic structure as the second Mamba module, but different parameters. The first Mamba module is used as an example to explain the principle. The first Mamba module includes a compression layer (Flatten), a normalization layer (LayerNorm), a first linear layer (Linear), a second linear layer (Linear), a first activation function layer (SiLu), a one-dimensional convolution layer (1D Conv), a second activation function layer (SiLu), a state space model (SSM), a third linear layer, and a reshape function layer. For the F input to the first Mamba module, a , the compression layer first changes its dimension from [b,c,h,w,d] to [b,c,h×w×d], and then after the normalization layer processing, we get F b Then, F b The first path is processed through a linear layer, a one-dimensional convolutional layer, a SiLu activation function, and a state space model, and outputs a feature F of [b, c, h×w×d] dimensions. c ; The second path is processed by a linear layer and SiLu activation function, outputting the feature F of [b,c,h×w×d] dimensions d . The output F of the first path c and the output F of the second path d Multiply them together, and then process them through a linear layer to output the feature F of [b,c,h×w×d] dimensions e , F e After being processed by the reshape function layer, the output F of [b,c,h,w,d] dimensions is finally obtained f .

[0082] State Space Model (SSM, such as Figure 2 The process of converting input X into output y is mainly realized by the state update equation and the output equation. The state update equation is as follows:

[0083] h(t+1)=Ah(t)+Bx t

[0084] Where h(t) represents the current state, x t is the data of a time step t of X. A is the state transfer matrix, which describes how the state of the system is transferred between time steps and defines the dynamic characteristics of the system. B is the control matrix, which describes the external input x t How to act on the system state h(t) and affect the dynamic evolution of the system. The output equation is as follows:

[0085] y(t)=Ch(t)+Dx t

[0086] C represents the observation matrix, which describes how the system state h(t) is mapped to the output y(t), and defines the contribution of the state to the observation value. D represents the direct transfer matrix, which describes how the input directly affects the output and defines the direct coupling relationship between the input and the output. A, B, C, and D are all parameter matrices that are updated during training. They are initialized randomly and are automatically updated during the training process.

[0087] In the state update equation, matrix B maps the input to the state and determines the direct impact of the input on the state, while matrix A describes the dynamic relationship between the states and is responsible for propagating the impact of the input to subsequent time steps through the state recursion. In the output equation, matrix C projects the state to the output space and determines the indirect contribution of the state to the output, while matrix D describes the direct impact of the input on the output. Ultimately, the input affects the state through the synergy of A and B, and then is mapped to the output through C and D, completing the input-output conversion of the dynamic system. A and B control the state update, and C and D map the input X.

[0088] Segmentation Decoder seg and reconstruct the decoder rec The basic structure of is the same, but the parameters are not shared. They all include the first convolution upsampling module, the second convolution upsampling module, the third convolution upsampling module, the fourth convolution upsampling module and the Final Conv module (including two convolution operations, the first convolution kernel size is 3, and the padding step is 1; the second convolution kernel size is 1, the step size is 1, and the padding is 0, which is only used to convert the number of channels, corresponding to the real label and the original image respectively).

[0089] For the segmentation decoder seg , the feature F 04 With feature F 06 As the input of the first convolution upsampling module, the first convolution upsampling module outputs the feature F 11 ; The feature F 03 With feature F 11 As the input of the second convolution upsampling module, the second convolution upsampling module outputs the feature F 12 ; The feature F 02 With feature F 12 As the input of the third convolution upsampling module, the third convolution upsampling module outputs the feature F 13 ; The feature F 01 With feature F 13 As the input of the fourth convolution upsampling module, the fourth convolution upsampling module outputs the feature F 14 ; The feature F 14 As a segmentation decoder Decoder segThe Final Conv module first performs a convolution operation with a kernel of 3×3×3 and a step size of 1, converts the number of channels to proto_dim (the dimension of the subsequent category prototype), and outputs the secondary segmentation result out sub_seg Then perform a second convolution operation with a kernel of 1 and a step size of 1, convert proto_dim into the number of segmented categories, and then perform the secondary segmentation result out sub_seg Perform BatchNorm normalization and ReLU activation processing to output the final segmentation result out seg .

[0090] Segmentation Decoder seg The first convolution upsampling module receives the input feature F 04 With feature F 06 , for F 06 First, perform an upsampling operation using trilinear interpolation mode with a scaling factor of (2,2,2) to obtain the upsampled intermediate result mid_outputs01; then calculate mid_outputs1 and F 04 The size difference in the spatial dimension is used to determine the corresponding filling amount, and the F 04 Fill the intermediate result mid_outputs02; then concatenate mid_outputs01 and mid_outputs02 in the channel dimension; finally, input the concatenated result into the convolution module consisting of two layers of convolution operations. The first convolution operation is a convolution operation with a convolution kernel of 3×3×3 and a step size of 1. The number of channels is determined by F 04 The number of channels c is 0.5 times the original one, i.e., c' (c'=c / 2), and InstanceNorm normalization and ReLU activation are performed to obtain the intermediate layer features, and then the second convolution operation is performed; the second convolution operation is also a convolution operation with a convolution kernel of 3×3×3 and a step size of 1. The number of channels c' remains unchanged, and InstanceNorm normalization and ReLU activation are performed, and finally the output feature F of the first convolution upsampling module is output 11 .

[0091] Segmentation Decoder segThe internal operations of the second convolution upsampling module, the third convolution upsampling module, and the fourth convolution upsampling module are exactly the same as those of the first convolution upsampling module, the difference is that the number of output channels is different, and the number of output channels of the latter convolution block is 0.5 times the number of output channels of the previous convolution block; as an embodiment, the numbers of output channels of the first convolution upsampling block, the second convolution upsampling block, the third convolution upsampling block, and the fourth convolution upsampling block are 128, 64, 32, and 16, respectively, and the image dimensions are 10, 20, 40, and 80, respectively.

[0092] Reconstruction decoder Decoder rec The first convolution upsampling module receives two input features F 04 With feature F 06 , for F 06 First, perform an upsampling operation using trilinear interpolation mode with a scaling factor of (2,2,2) to obtain the upsampled intermediate result mid_outputs11; then calculate mid_outputs11 and F 04 The size difference in the spatial dimension is used to determine the corresponding filling amount, and the F 04 Pad to get the intermediate result mid_outputs12; then concatenate mid_outputs11 and mid_outputs12 in the channel dimension; finally, input the concatenated result into the convolution module consisting of two layers of convolution operations. The first convolution operation is a convolution operation with a convolution kernel of 3×3×3 and a step size of 1. The number of channels is determined by F 04 The number of channels c is 0.5 times the original one, i.e. c' (c'=c / 2), and the InstanceNorm normalization kernel ReLU activation is performed to obtain the intermediate layer features for the second convolution operation; the second convolution operation is also a convolution operation with a convolution kernel of 3×3×3 and a step size of 1. The number of channels c' remains unchanged, and InstanceNorm normalization and ReLU activation are performed to finally output the output features of the first convolution upsampling module.

[0093] Reconstruction decoder Decoder rec The operations inside the second convolution upsampling module, the third convolution upsampling module, and the fourth convolution upsampling module are exactly the same as those in the first convolution upsampling module, except that the number of output channels is different. The number of output channels of the latter convolution block is 0.5 times the number of output channels of the previous convolution block. As an embodiment, the number of output channels of the first convolution upsampling block, the second convolution upsampling block, the third convolution upsampling block, and the fourth convolution upsampling block are 128, 64, 32, and 16, respectively, and the image dimensions are 10, 20, 40, and 80, respectively. The output of the fourth convolution upsampling block is input to the reconstruction decoder Decoder recThe Final Conv module first performs a convolution operation with a kernel of 3×3×3 and a step size of 1 on the input, and the number of channels is converted to proto_dim (i.e., the dimension of the subsequent category prototype). Then, a second convolution operation is performed to convert proto_dim to 1, and the output of the first convolution operation is subjected to BatchNorm normalization and ReLU activation to obtain the final reconstruction result out rec .

[0094] Step 3: Pre-training of student models

[0095] The student model is pre-trained using the labeled data in the dataset. The pre-training phase stops after 1000 iterations, and the initial network parameters are initialized using the default Kaiming initialization (i.e., normal distribution initialization).

[0096] Using random sampling, a batch of labeled data in the data set is input into the initialized student model, and the segmentation result out seg Calculate the linear classification loss with the Ground-Truth segmentation label. Reconstruct the decoder output out rec The original image input to the student model is sent to the perceptual loss network (the perceptual loss network uses the 3D VGG model as the backbone, with the number of channels being [1, 64, 256, 256, 512]. Each module contains a convolutional layer, a maximum pooling layer, and a ReLu activation function layer. For details, please refer to the attached Figure 5 ), and calculate the perceptual loss. The loss function to be optimized is as follows:

[0097] L cls_ce_3d =L ce (f ds (f e (x i )),y i ) (1)

[0098] L seg_dice_3d =L dice (softmax(f ds (f e (x i ))),y i ) (2)

[0099] L cls_3d =weight ce_w ×L cls_ce_3d +weight dice_w ×L seg_sice_3d (3)

[0100] L percen_3d =L1 (vgg(f dr (f e (x i )),vgg(x i ))) (4)

[0101] L a =L cls_3d +L percen_3d (5)

[0102] Among them, x i represents input, f e (·) represents the output of FMmamba-Encoder, f ds (·) represents the segmentation decoder Decoder seg The output, f dr (·) represents the reconstruction decoder rec The output of , vgg(·) represents the output of the perceptual loss network 3D VGG. i Represents x i Ground-Truth segmentation label. weight ce_w 、weight dice_w They represent the cross entropy loss weight and Dice loss weight, respectively, both are 0.5. ce (·,·) and L dice (·,·) means using cross entropy loss and Dice loss, multiplying by weights and summing them up to get the linear classification loss L cls_3d . Perceptual loss L percen_3d Use L 1 The loss is calculated and finally summed with the linear classification loss to obtain the training loss of a batch of data in the pre-training stage of the student model.

[0103] According to the training loss, the stochastic gradient descent (SGD) optimizer (with an initial learning rate of 0.01 and a momentum parameter of 0.9) is used to reversely update the student model (including the encoder FMmamba-Encoder, the segmentation decoder Decoder seg And reconstruct the decoder Decoder rec ) to complete the training of one batch of data; the network parameters that completed the previous batch of data training are used as the initial parameters for the next batch of training, and the next batch of labeled data is input. The process of training one batch of data is repeated continuously until the training of the last batch of labeled data in the data set is completed, completing one round of training; the network parameters that completed the previous round of training are used as the initial parameters for the next round of training, and the training is continuously iterated until the number of iterations reaches 1000, and the student model that has completed the pre-training is obtained.

[0104] Step 4: Global Category Prototype Training of Student Model

[0105] Step 4.1: First, the global category prototype P 0 Set to all zero parameters of [num_classes, subcluster, proto_dim] dimensions, where num_classes, subcluster, and proto_dim represent the number of segmented classes as 6, the number of subcluster centers in a single class prototype as 4, and the dimension of the global class prototype as 64, respectively.

[0106] Input a batch of labeled data from the dataset (selected by random sampling) into the student model pre-trained in the third step, and split the decoder seg Output segmentation result out seg , secondary segmentation result out sub_seg , the reconstruction decoder outputs the reconstruction result out rec .

[0107] For the secondary segmentation result out sub_seg The dimension is transformed from [b,c,h,w,d] to [b×h×w×d,c] to obtain out_feat. The ground truth is first transformed from [b,h,w,d] to [b×d,h,w] to obtain label_2d (this process is the toSlice(·) process), and label_2d is transformed to [b×h×w×d] to obtain label_expand.

[0108] Then, from out_feat, according to its corresponding label_expand category label, the feature subset belonging to each category is filtered out in turn. The specific operation is: according to the position of the pixel of the kth category of the label_expand category label (the corresponding pixel value is k, k = 0, 1, ..., num_classes, num_classes represents the total number of categories), the features of the corresponding position in out_feat are filtered and collected.

[0109] For the features of the kth category collected, K-means clustering is used to obtain the small category prototype corresponding to this category. The specific operation of K-Means is to randomly initialize subcluster subcenters according to the number of sub-prototype centers in each small category prototype, and then for each feature point in the features screened out of this category, calculate its distance from each cluster center (usually using Euclidean distance, etc.), and divide it into the class represented by the subcenter with the closest distance; after completing the classification of all feature points, recalculate the center position of the cluster according to the features of the sample points contained in each cluster (generally taking the mean of each feature dimension); repeat the above classification and recalculation of cluster centers until 300 iterations are reached. At this time, the subcluster cluster centers determined are the final clustering results of the data, and the feature points are divided into subcluster different subcategories. Finally, the output of the small category prototype p that can represent this class is c_k (k=0, 1…num_classes), that is, the eigenvalues ​​of several subcluster centers, and the small category prototype with dimension [subcluster,proto_dim] is obtained.

[0110] Finally, K-means clustering is performed on num_classes categories to obtain the small category prototype p of each category. c_k Then all the small category prototypes are concatenated into a tensor to obtain the initialized global category prototype P 1 , the dimension is [num_classes,subcluster,proto_dim].

[0111] Step 4.2: The secondary segmentation result out obtained in step 4.1 sub_seg First, we transform the dimension from [b,c,h,w,d] to [b×h×w×d,c] to get out_feat. The ground truth is first transformed from [b,h,w,d] to [b×d,h,w] to get label_2d, which is then transformed to [b×h×w×d] label_expand. We combine out_feat and the current global category prototype P 1Perform matrix multiplication along the last dimension and output feat_proto_sim with dimensions [b×h×w×d,subcluster,num_classes]. Then calculate the maximum value tmp of feat_proto_sim in the subcluster dimension (the maximum value has dimensions [b×h×w×d,num_classed]). The maximum value tmp is regularized by Layernorm and its dimensions are changed to [b×d,num_classes,h,w], and the secondary segmentation result out is obtained. sub_seg The global category prototype P 1 The segmentation result proto_seg, that is, nearest_proto_distance.

[0112] Then use out_feat, feat_proto_sim, label_expand, and nearest_proto_distance to update the momentum of the global category prototype. First, search the nearest_proto_distance on the num_classes dimension to get the pred_seg with the maximum index dimension [b×d,h,w]. Then query and calculate the pred_segs of label_expand and pred_seg after stretching to [b×d×h×w] to get the Boolean matrix mask of the positions where the elements of the two are equal (True if equal, False if not equal).

[0113] Then, starting from category 0, we traverse num_classes categories and perform momentum updates on the small category prototypes of each category. Taking the update of the kth category as an example, the update method for each category is as follows.

[0114] First, find the index of feat_proto_sim according to the last dimension, the index value is category k, and get the feature tmp k_0 Then find the position of the pixel value equal to k in label_expand, record it as position index_k, and use this position in temp k_0 Find the index of the first dimension and get the similarity measure q k .

[0115] Then the similarity measure q k Perform sinkhorn operation (the similarity measure q kThrough iterative adjustment, the default iteration is 3 times, so that the sum of its rows and columns is close to uniform distribution), the final similarity measure q (dimension is [b×h×w×d, sub_cluster]). Then find the feature m_k of the Boolean matrix mask at the index_k position, copy m_k to expand it subcluster times in the column dimension, and get the mask matrix m k_0 Then find the index in the first dimension of out_feat, find the feature m_w at index_k, copy m_w to expand it proto_dim times in the column dimension, and get the mask matrix m k_1 .

[0116] Then calculate the updated value f of the small category prototype of the current class k , the formula is:

[0117] Where m q_k =qm k_0 , represents the effective soft assignment weight, w q_k =m_k×m k_1 ; Then, the momentum update formula is used to update the newly calculated update value f k Incorporating old category archetypes In the example above, we get the updated small category prototype

[0118]

[0119] Where μ is the momentum coefficient. Update value f k The consistency is maintained through l2-normalization, i.e., the normalize operation. The above operation is performed on each small category prototype to obtain each updated small category prototype, and then each updated small category prototype is spliced ​​to obtain an updated global category prototype P 2 .

[0120] Step 4.3: Decoder according to the segmentation in step 4.1 seg Output segmentation result out seg , secondary segmentation result out sub_seg , the reconstruction decoder outputs the reconstruction result out rec , calculate the training loss, the loss function is as follows:

[0121] L proto_ce =L ce (Proto(f ds_sub (f e (x i ))),toSlice(y i )) (6)

[0122] L proto_seg_dice =L dice (Softmax(Proto(f ds_sub (f e (x i )))),toSlice(y i )) (7)

[0123] L proto_3d =weight ce_w ×L proto_ce +weight dice_w ×L proto_seg_dice (8)

[0124] weight ce_w 、weight dice_w is the weight matrix parameter; according to formula (3), formula (4), and formula (8), the total training loss L of the batch of labeled data is obtained b , the specific calculation formula is:

[0125] L b =L proto_3d +L percen_3d +L cls_3d (9)

[0126] L percen_3d , L cls_3d According to formula (4) and formula (3), the perceptual loss and linear classification loss are calculated respectively. proto_3d is the current global category prototype P 2 loss.

[0127] For formula (6), Proto(f ds_sub (f e (x i ))) refers to the secondary segmentation result f output by the segmentation decoder ds_sub (f e (x i ))'s current global category prototype P 2 The segmentation result has the dimension [b×d,num_classes,h,w]; toSlice(y i ) means to convert the label y of shape [b,h,w,d] i Map to [b×d,h,w].

[0128] For formula (7), Proto(f ds_sub (f e (x i )))、toSlice(y i) is consistent with the operation in formula (6), Softmax(Proto(f ds_sub (f e (x i )))) refers to the secondary segmentation result f output by the segmentation decoder ds_sub (f e (x i )))'s current global category prototype P 2 Perform Softmax operation on the segmentation result and adjust the dimension to [b×d,num_classes,h,w].

[0129] According to the total loss L of a batch b , using the stochastic gradient descent (SGD) optimizer (initial learning rate is 0.01, momentum parameter is 0.9), update the student model (including encoder FMmamba-Encoder, segmentation decoder Decoder seg And reconstruct the decoder Decoder rec ) to complete the iterative training of a batch of labeled data.

[0130] Step 4.4: Input the next batch of labeled data in the data set into the pre-trained student model that has completed one iteration of training in step 4.3, and repeat the process of steps 4.2, 4.3, and 4.4 until the last batch of labeled data in the data set is trained and one round of training is completed; use the network parameters that have completed the previous round of training as the initial parameters for the next round of training, and repeat the training process for one round until the iterative round reaches 1000 times, complete the global category prototype training of the student model, and obtain the current global category prototype P o .

[0131] Step 5: Deep training of segmentation network model

[0132] Step 5.1: Use the network parameters of the student model that has completed the global category prototype training in the fourth step as the initial network parameters of the segmentation network model in the deep training phase. The initial learning rate of the entire training iteration process is 0.01, and the momentum parameter is 0.9. The starting values ​​of the teacher model parameters are all set to 0, and the teacher model parameter values ​​are initialized according to the initial network parameter values ​​of the student model. The specific operation is that in the mean teacher model, the student model weight θ S The exponential moving average (EMA) of T , to integrate information from different training steps. Specifically, at training iteration step t, the weight θ of the teacher model is T Updated to: in Respectively represent the parameters of the current student model and the current teacher model parameters, Indicates based on and The parameters of the teacher model for the next step are obtained; α is used to control the rate of EMA decay, which is dynamically updated with the number of iterations (that is, updated once for one batch training).

[0133] Step 5.2: Use random sampling to select a batch of training data in the dataset. A batch of training data includes half labeled data and half unlabeled data. Input the labeled data in a batch into the student model that has completed the global category prototype training in the fourth step. The student model outputs the segmentation result of the labeled data out seg_l , secondary segmentation result out sub_seg_l And the reconstruction result out rec_l ;

[0134] According to the output of the student model for the labeled data, we first use the method in step 4.2 to generate the global category prototype P o Based on this, the global category prototype is updated to obtain the current global category prototype; then the segmentation result out seg_l , secondary segmentation result out sub_seg_l And the true label GroundTruth respectively calculates the linear classification loss L of the labeled data cls_3d_l and the current global category prototype loss L proto_3d_l , the reconstruction result out rec_l The labeled image of the input model is input to the perceptual loss network and the perceptual loss L is calculated. percen_3d_l , linear classification loss L cls_3d_l , the current global category prototype loss L proto_3d_l , perceptual loss L percen_3d_l The sum is the training loss L for labeled data l .

[0135] Then the unlabeled data in a batch is input into the student model that has completed the global category prototype training and the teacher model that has completed the network parameter initialization, and the segmentation result output by the teacher model is output. seg_ul As the pseudo label of the corresponding input unlabeled data when the student model calculates the loss; after the unlabeled data is input to the student model, the segmentation result of the student model is used seg_ul , secondary segmentation result out sub_seg_ul And the pseudo label Pseudo Label generated by the teacher model is used to calculate the corresponding linear classification loss L cls_3d_ul and the current global category prototype loss Lproto_3d_ul , the reconstruction result out rec_ul The image of the input teacher model is input to the perceptual loss network and the perceptual loss L is calculated percen_3d_ul , linear classification loss L cls_3d_ul , the current global category prototype loss L proto_3d_ul , perceptual loss L percen_3d_ul The sum is the loss of unlabeled data L ul .

[0136] The above-mentioned linear classification loss, perceptual loss and current global category prototype loss are calculated according to formula (3), formula (4) and formula (8) respectively.

[0137] For a batch of training data, the training loss is L:

[0138] L=L l +consistency_weight×L ul

[0139] In the above formula, consistency_weight is the consistency weight of the teacher-student model, which is updated with the number of training rounds (that is, updated once per training round). The specific update operation is: the current number of iterations iter_num is divided by 100, then smoothed with Sigmoid, and then multiplied by the preset basic consistency weight parameter value (initial 0.5), so as to dynamically change the value of consistency_weight as the rounds progress during the training process.

[0140] According to the training loss of the training data of this batch, the network parameters of a student model are reversely updated. At the same time, the network parameters of the teacher model are updated once through the exponential moving average EMA using the updated network parameters of the student model to complete the training of a batch of data samples;

[0141] Step 5.3: Randomly sample the training set D train The next batch of training data in the dataset is taken. Similarly, the training data in this batch includes half labeled data and half unlabeled data. The network parameters of the student model and the teacher model when the previous batch of data samples are trained are used as the initial parameters for the next batch of training data. The process in step 5.2 and step 5.3 is repeated until the last batch of training data in the dataset, including half labeled data and half unlabeled data, is trained to complete a round of training. The network parameters at the completion of the previous round of training are used as the initial network parameters at the beginning of the next round of training. This process is repeated until the training round reaches the preset 14,000 times, and a trained segmentation network model is obtained.

[0142] Step 6: 3D whole heart MRI image segmentation

[0143] In a clinical setting, we first obtain the patient's 3D whole-heart magnetic resonance imaging (MRI) image data, then perform center cropping on the 3D image data, obtain an image of the same size (80,80,80) as in the first step, and perform z-score normalization on the pixels of the image. Then, we input the image into the student model of the trained segmentation network model in the fifth step, and then use its encoder FMmamba-Encoder and segmentation decoder Decoder to obtain the image. seg To perform segmentation prediction on the input image, obtain the segmentation result, and complete the segmentation of the image data of 3D whole-heart magnetic resonance imaging.

[0144] The essence of the present invention is to design a single encoder-dual decoder network FMDD-UNet based on 3DUNet and Mamba, and then combine it with prototype learning and semi-supervised learning methods to solve the semi-supervised learning task; taking the MMWHS publicly available dataset as an example, the method of the present invention and the existing models are used to perform 3D whole-heart magnetic resonance imaging segmentation prediction, and the segmentation effects of various models are shown in Table 1.

[0145] Table 1 Comparison of segmentation results of various methods

[0146]

[0147]

[0148] As can be seen from the above table, the average dice similarity coefficient of the Mean-Teacher segmentation model is 79.32%, the average dice similarity coefficient of the Generative Adversarial Semi-Supervised Segmentation Model (AdvSemiSeg) is 80.99%, and the average dice similarity coefficient of the Class Imbalanced Semi-Supervised Segmentation Model (SSCI) is 84.39%. The average dice similarity coefficient of the method of the present invention is 85.67%, which is 1.28 percentage points higher than 84.39% of SSCI. The method of the present invention demonstrates its superior performance in the overall cardiac structure segmentation task, and the improvement of the segmentation effect will bring great help to the diagnosis of diseases.

[0149] Any matters not described in the present invention are applicable to the prior art.

Claims

1. A semi-supervised 3D whole-heart magnetic resonance image segmentation method, characterized in that: The method comprises the following steps: Step 1: Preprocess the 3D full heart image data to obtain the data set to be trained The training samples in the dataset are obtained by randomly cropping the 3D full heart image data to a set size, adding random offset values ​​to increase the diversity of the data, and finally performing z-score standardization on the data; the dataset consists of L+U training samples, of which L are labeled data and U are unlabeled data; Step 2: Build a segmentation network model The overall network architecture of the segmentation network model adopts a teacher-student model. The student model and the teacher model have the same network structure. The student model is an FMDD-UNet network structure improved from 3D UNet, specifically including an FMmamba-Encoder encoder, a segmentation decoder Decoder seg and reconstruct the decoder rec , wherein the FMmamba-Encoder encoder is composed of a first convolution block, a second convolution block, a third convolution block, a fourth convolution block, a fifth convolution block, a first Freq-Mamba Block, a second Freq-Mamba Block, and a third Freq-Mamba Block; The input image X of the student model is first processed by the first convolution block, the second convolution block, the third convolution block, the fourth convolution block, and the fifth convolution block in sequence. The output features of the previous convolution block are used as the input features of the next convolution block. The output features of the first convolution block, the second convolution block, the third convolution block, the fourth convolution block, and the fifth convolution block are recorded as F 01 、F 02 、F 03 、F 04 、F 05 ; After the input feature X is input to the first convolution block, a convolution operation with a convolution kernel of 1×1×1 and a stride of 1 is first performed to obtain y, y=F(X)+X, where F represents the above convolution operation; y is a 5D tensor of size b×c×h×w×d, where b represents the batch size, h represents the spatial height, w represents the spatial width, d is the spatial depth, and c is the number of channels; then two convolution operations are performed on y. The first convolution operation is a convolution operation with a convolution kernel of 3×3×3 and a stride of 2. The number of channels is doubled from c. c', and InstanceNorm normalization and ReLU activation are performed to obtain the intermediate layer features, and then a second convolution operation is performed; the second convolution operation is also a convolution operation with a convolution kernel of 3×3×3 and a step size of 2, the number of channels c' remains unchanged, and InstanceNorm normalization and ReLU activation are performed; the final result of the second convolution is subjected to a MaxPool maximum pooling of size 2 to obtain a 5D tensor of size h / 2×w / 2×d / c', which is F 01 ; The operations inside the second convolution block, the third convolution block, the fourth convolution block, and the fifth convolution block are exactly the same as those in the first convolution block. The difference is that the number of output channels of the latter convolution block is twice the number of output channels of the previous convolution block; The output feature F of the fifth convolutional block 05 As the input of the first Freq-Mamba Block, the output of the first Freq-MambaBlock is used as the input of the second Freq-Mamba Block, the output of the second Freq-Mamba Block is used as the input of the third Freq-Mamba Block, and the output feature of the third Freq-Mamba Block is recorded as F 06 ; The basic structure, number of feature channels, etc. of the first Freq-Mamba Block, the second Freq-Mamba Block, and the third Freq-Mamba Block are exactly the same, and the parameters are not shared; The first Freq-Mamba Block includes a wavelet transform operation, two Mamba modules, an inverse transform operation of a wavelet transform, a BN operation, and a ReLu operation; The feature F 05 After being input into the first Freq-Mamba Block, it is processed in two paths; for the first path, the feature F 05 First, after a wavelet transform operation, the result is input into the first Mamba module. The output of the first Mamba module undergoes an inverse wavelet transform operation to obtain feature FM1. For the second path, feature F 05 After being processed by the second Mamba module, feature FM2 is obtained; then feature FM1 and feature FM2 are concatenated in the channel dimension, and the obtained result is successively subjected to BN operation and ReLu operation to obtain the output of the first Freq-Mamba Block; The first Mamba module has the same basic structure as the second Mamba module, but different parameters. The first Mamba module is used as an example to explain the principle. The first Mamba module includes a compression layer, a normalization layer, a first linear layer, a second linear layer, a first activation function layer, a one-dimensional convolution layer, a second activation function layer, a state space model, a third linear layer, and a reshape function layer. For the F input to the first Mamba module, a , the compression layer first changes its dimension from [b,c,h,w,d] to [b,c,h×w×d], and then after the normalization layer processing, we get F b ; then F b The first path is processed through a linear layer, a one-dimensional convolutional layer, a SiLu activation function, and a state space model, and outputs a feature F of [b, c, h×w×d] dimensions. c ; The second path is processed by a linear layer and SiLu activation function, outputting the feature F of [b,c,h×w×d] dimensions d ; Output F of the first path c and the output F of the second path d Multiply them together, and then process them through a linear layer to output the feature F of [b,c,h×w×d] dimensions e , F e After being processed by the reshape function layer, the output F of [b,c,h,w,d] dimensions is finally obtained f ; Segmentation Decoder seg and reconstruct the decoder rec The basic structure of the two networks is the same, but the parameters are not shared. They all include the first convolution upsampling module, the second convolution upsampling module, the third convolution upsampling module, the fourth convolution upsampling module and the Final Conv module. For the segmentation decoder seg , the feature F 04 With feature F 06 As the input of the first convolution upsampling module, the first convolution upsampling module outputs the feature F 11 ; The feature F 03 With feature F 11 As the input of the second convolution upsampling module, the second convolution upsampling module outputs the feature F 12 ; The feature F 02 With feature F 12 As the input of the third convolution upsampling module, the third convolution upsampling module outputs the feature F 13 ; The feature F 01 With feature F 13 As the input of the fourth convolution upsampling module, the fourth convolution upsampling module outputs the feature F 14 ; The feature F 14 As a segmentation decoder Decoder seg The Final Conv module first performs a convolution operation with a kernel of 3×3×3 and a step size of 1, converts the number of channels to proto_dim, and outputs the secondary segmentation result out sub_seg Then perform a second convolution operation with a kernel of 1 and a step size of 1, convert proto_dim into the number of segmented categories, and then perform the secondary segmentation result out sub_seg Perform BatchNorm normalization and ReLU activation processing to output the final segmentation result out seg ; Segmentation Decoder seg The first convolution upsampling module receives the input feature F 04 With feature F 06 , for F 06 First, perform an upsampling operation using trilinear interpolation mode with a scaling factor of (2,2,2) to obtain the upsampled intermediate result mid_outputs01; then calculate mid_outputs1 and F 04 The size difference in the spatial dimension is used to determine the corresponding filling amount, and the F 04 Fill the intermediate result mid_outputs02; then concatenate mid_outputs01 and mid_outputs02 in the channel dimension; finally, input the concatenated result into the convolution module consisting of two layers of convolution operations. The first convolution operation is a convolution operation with a convolution kernel of 3×3×3 and a step size of 1. The number of channels is determined by F 04 The number of channels c is 0.5 times the original one, that is, c', and InstanceNorm normalization and ReLU activation are performed to obtain the intermediate layer features, and then the second convolution operation is performed; the second convolution operation is also a convolution operation with a convolution kernel of 3×3×3 and a step size of 1. The number of channels c' remains unchanged, and InstanceNorm normalization and ReLU activation are performed, and finally the output feature F of the first convolution upsampling module is output 11 ; Segmentation Decoder seg The operations inside the second convolution upsampling module, the third convolution upsampling module, and the fourth convolution upsampling module are exactly the same as those in the first convolution upsampling module. The difference is that the number of output channels is different. The number of output channels of the latter convolution block is 0.5 times the number of output channels of the previous convolution block. The decoder is reconstructed. rec The first convolution upsampling module receives two input features F 04 With feature F 06 , for F 06 First, perform an upsampling operation using trilinear interpolation mode with a scaling factor of (2,2,2) to obtain the upsampled intermediate result mid_outputs11; then calculate mid_outputs11 and F 04 The size difference in the spatial dimension is used to determine the corresponding filling amount, and the F 04 Pad to get the intermediate result mid_outputs12; then concatenate mid_outputs11 and mid_outputs12 in the channel dimension; finally, input the concatenated result into the convolution module consisting of two layers of convolution operations. The first convolution operation is a convolution operation with a convolution kernel of 3×3×3 and a step size of 1. The number of channels is determined by F 04 The number of channels c is 0.5 times the original one, that is, c', and the InstanceNorm normalization kernel ReLU activation is performed to obtain the intermediate layer features for the second convolution operation; the second convolution operation is also a convolution operation with a convolution kernel of 3×3×3 and a step size of 1. The number of channels c' remains unchanged, and InstanceNorm normalization and ReLU activation are performed to finally output the output features of the first convolution upsampling module; Reconstruction decoder Decoder rec The operations inside the second convolution upsampling module, the third convolution upsampling module, and the fourth convolution upsampling module are exactly the same as those in the first convolution upsampling module. The difference is that the number of output channels is different. The number of output channels of the latter convolution block is 0.5 times the number of output channels of the previous convolution block. The output of the fourth convolution upsampling block is input to the reconstruction decoder Decoder rec The Final Conv module first performs the first convolution operation on the input with a convolution kernel of 3×3×3 and a step size of 1, and the number of channels is converted to proto_dim. Then, the second convolution operation is performed to convert proto_dim to 1. The output of the first convolution operation is then BatchNorm normalized and ReLU activated to obtain the final reconstruction result out rec ; Step 3: Pre-training of student models The student model is pre-trained using the labeled data in the dataset. The pre-training phase is terminated after 1000 iterations. The initial network parameters are initialized using the default Kaiming. Random sampling is performed to input a batch of labeled data in the dataset into the initialized student model. The segmentation results are used to seg Calculate the linear classification loss with the Ground-Truth segmentation label; reconstruct the decoder output reconstructed result out rec The original image input to the student model is fed into the perceptual loss network and the perceptual loss is calculated; the loss function to be optimized is as follows: L cls_ce_3d =L ce (f ds (f e (x i )),y i ) (1) L seg_dice_3d =L dice (softmax(f ds (f e (x i ))), y i ) (2) L cls_3d =weight ce_w ×L cls_ce_3d +weight dice_D ×L seg_dice_3d (3) L percen_3d =L1(vgg(f dr (f e (x i )),vgg(x i ))) (4) L a =L cls_3d +L percen_3d (5) Among them, x i represents input, f e (·) represents the output of FMmamba-Encoder, f ds (·) represents the segmentation decoder Decoder seg The output of f dr (·) represents the reconstruction decoder rec The output of , vgg(·) represents the output of the perceptual loss network; y i Represents x i Ground-Truth segmentation label; weight ce_w 、weight dice_w They represent the cross entropy loss weight and Dice loss weight, both of which are 0.5; L ce (·,·) and L dice (·,·) means using cross entropy loss and Dice loss, multiplying by weights and summing them up to get the linear classification loss L cls_3d ; Perceptual loss L percen_3d Use L1 loss to calculate, and finally sum it with linear classification loss to get the training loss of a batch of data in the pre-training stage of the student model; According to the training loss, the stochastic gradient descent optimizer is used to reversely update the network parameters of the student model once to complete the training of one batch of data; the network parameters that have completed the training of the previous batch of data are used as the initial parameters of the next batch of training, and the next batch of labeled data is input, and the process of training a batch of data is repeated continuously until the training of the last batch of labeled data in the data set is completed, completing one round of training; the network parameters that have completed the previous round of training are used as the initial parameters of the next round of training, and the training is continuously iterated until the iteration round reaches 1000 times, and the pre-trained student model is obtained; Step 4: Global Category Prototype Training of Student Model Step 4.1: First, set the global category prototype P0 to all zero parameters of [num_classes, subcluster, proto_dim] dimensions, where num_classes, subcluster, and proto_dim represent the number of segmented categories, the number of subcluster centers in a single category prototype, and the dimension of the global category prototype, respectively; Random sampling, input a batch of labeled data in the dataset into the student model pre-trained in the third step, split decoder Decoder seg Output segmentation result out seg , secondary segmentation result out sub_seg , the reconstruction decoder outputs the reconstruction result out rec ; For the secondary segmentation result out sub_seg Perform dimension transformation from [b,c,h,w,d] to [b×h×w×d,c] to obtain out_feat; transform the ground truth from [b,h,w,d] to [b×d,h,w] to obtain label_2d, and then transform label_2d to [b×h×w×d] to obtain label_expand; then filter out the feature subsets belonging to each category from out_feat according to the corresponding label_expand category label; the specific operation is: according to the position of the pixel of the kth category of the label_expand category label, that is, the corresponding pixel value is k, filter and collect the features of the corresponding position in out_feat; for the collected features of the kth category, use K-means clustering to obtain the small category prototype corresponding to this category; then concatenate all the small category prototypes into a tensor to obtain the initialized global category prototype P1, with the dimension of [num_classes,subcluster,proto_dim]; Step 4.2: The secondary segmentation result out obtained in step 4.1 sub_seg First, the dimension is transformed from [b,c,h,w,d] to [b×h×w×d,c] to obtain out_feat; the real label Ground Truth is first transformed from [b,h,w,d] to [b×d,h,w] to obtain label_2d, and label_2d is transformed into [b×h×w×d] label_expand through the dimension transformation; out_feat and the current global category prototype P1 are matrix multiplied along the last dimension, and the output dimension is [b×h×w×d,subcluster,num_classes] feat_proto_sim; then the maximum value tmp of feat_proto_sim on the subcluster dimension is calculated, and the maximum value tmp is regularized by Layernorm and the dimension is changed to [b×d,num_classes,h,w], that is, the secondary segmentation result out sub_seg The segmentation result proto_seg of the global category prototype P1, that is, nearest_proto_distance; Then use out_feat, feat_proto_sim, label_expand, and nearest_proto_distance to update the momentum of the global category prototype; first, search the nearest_proto_distance on the num_classes dimension to get the pred_seg with the maximum value index dimension of [b×d,h,w]; secondly, stretch label_expand and pred_seg to pred_seg of [b×d×h×w] s Perform query calculation to obtain the Boolean matrix mask of the positions where the elements of the two are equal; Then, starting from category 0, we traverse num_classes categories and perform momentum updates on the small category prototypes of each category. Taking the update of the kth category as an example, the update method for each category is as follows; First, find the index of feat_proto_sim according to the last dimension, the index value is category k, and get the feature tmp k_0 ; Then find the position of the pixel value equal to k in label_expand, record it as position index_k, and use this position in tmpk _0 Find the index of the first dimension and get the similarity measure q k ; Then the similarity measure q k Perform a sinkhorn operation to obtain the final similarity measure q; then find the feature m_k of the Boolean matrix mask at the index_k position, copy m_k to expand it subcluster times in the column dimension, and obtain the mask matrix m k_0 ; Then find the index in the first dimension of the feature out_feat, find the feature m_w at the index_k position, copy m_w to expand it proto_dim times in the column dimension, and get the mask matrix m k_1 ; Then calculate the updated value f of the small category prototype of the current class k , the formula is: Where m q_k =qm k_0 , represents the effective soft assignment weight, w q_k =m_k×m k_1 ; Then, the momentum update formula is used to update the newly calculated update value f k Incorporating old category archetypes In the example above, we get the updated small category prototype where μ is the momentum coefficient; the update value f k Maintain consistency through l2-normalization, that is, the normalize operation; perform the above operation on each small category prototype to obtain each updated small category prototype, and then concatenate each updated small category prototype to obtain an updated global category prototype P2; Step 4.3: Decoder according to the segmentation in step 4.1 seg Output segmentation result out seg , secondary segmentation result out sub_seg , the reconstruction decoder outputs the reconstruction result out rec , calculate the training loss, the loss function is as follows: L proto_ce =L ce (Proto(f ds_ sub(f e (x i ))),toSlice(y i )) (6) L proto_seg_dice =L dice (Softmax(Proto(f ds_sub (f e (x i )))),toSlice(y i )) (7) L proto_3d =weight ce_w ×L proto_ce +weight dice_w ×L proto_seg_dice (8) weight ce_w 、weight dice_w is the weight matrix parameter; according to formula (3), formula (4), and formula (8), the total training loss L of the batch of labeled data is obtained b , the specific calculation formula is: L b =L proto_3d +L percen_3d +L cls_3d (9) L percen_3d , L cls_3d According to formula (4) and formula (3), the perceptual loss and linear classification loss are calculated respectively. proto_3d is the current global category prototype P2 loss; For formula (6), Proto(f ds_sub (f e (x i ))) refers to the secondary segmentation result f output by the segmentation decoder ds_sub (f e (x i )) is the segmentation result of the current global category prototype P2, with a dimension of [b×d,num_classes,h,w]; toSlice(yi) means to slice the label y of the shape [b,h,w,d] i Map to [b×d,h,w]; For formula (7), Proto(f ds_sub (f e (x i )))、toSlice(y i ) is consistent with the operation in formula (6), Softmax(Proto(f ds_sub (f e (x i )))) refers to the secondary segmentation result f output by the segmentation decoder ds_sub (f e (x i )) Perform a Softmax operation on the segmentation result of the current global category prototype P2 and adjust the dimension to [b×d,num_classes,h,w]; According to the total loss L of a batch b , using the stochastic gradient descent optimizer, update the trainable parameters of the student model once and complete the iterative training of a batch of labeled data; Step 4.4: Input the next batch of labeled data in the data set into the pre-trained student model that has completed one iteration of training in step 4.3, and repeat the process of steps 4.2, 4.3, and 4.4 until the last batch of labeled data in the data set is trained and one round of training is completed; use the network parameters that have completed the previous round of training as the initial parameters for the next round of training, and repeat the training process for one round until the iterative round reaches 1000 times, complete the global category prototype training of the student model, and obtain the current global category prototype P o ; Step 5: Deep training of segmentation network model Step 5.1: The network parameters of the student model that has completed the global category prototype training in the fourth step are used as the initial network parameters of the deep training stage of the segmentation network model. The starting values ​​of the parameters of the teacher model are all set to 0. The parameter values ​​of the teacher model are initialized according to the initial network parameter values ​​of the student model; the student model weight θ S The exponential moving average of is used to update the weights θ of the teacher model T ; Step 5.2: Randomly sample and select a batch of training data from the dataset. A batch of training data includes half labeled data and half unlabeled data. Input the labeled data in a batch into the student model that has completed the global category prototype training in step 4. The student model outputs the segmentation result of the labeled data out seg_l , secondary segmentation result out sub_seg_l And the reconstruction result out rec_l ; According to the output of the student model for the labeled data, we first use the method in step 4.2 to generate the global category prototype P o Based on this, the global category prototype is updated to obtain the current global category prototype; then the segmentation result out seg_l , secondary segmentation result out sub_seg_l And the true label Ground Truth respectively calculates the linear classification loss L of the labeled data cls_3d_l and the current global category prototype loss L proto_3d_l , the reconstruction result out rec_l The labeled image of the input model is input to the perceptual loss network and the perceptual loss L is calculated. percen_3d_l , linear classification loss L cls_3d_l , the current global category prototype loss L proto_3d_l , perceptual loss L percen_3d_l The sum is the training loss L for labeled data l ; Then the unlabeled data in a batch is input into the student model that has completed the global category prototype training and the teacher model that has completed the network parameter initialization, and the segmentation result output by the teacher model is output. seg_ul As the pseudo label of the corresponding input unlabeled data when the student model calculates the loss; after the unlabeled data is input to the student model, the segmentation result of the student model is used seg_ul , secondary segmentation result out sub_seg_ul And the pseudo label PseudoLabel generated by the teacher model are used to calculate the corresponding linear classification loss L cls_3d_ul and the current global category prototype loss L proto_3d_ul , the reconstruction result out rec_ul The image of the input teacher model is input to the perceptual loss network and the perceptual loss L is calculated percen_3d_ul , linear classification loss L cls_3d_ul , the current global category prototype loss L proto_3d_ul , perceptual loss L percen_3d_ul The sum is the loss of unlabeled data L ul ; The above linear classification loss, perceptual loss and current global category prototype loss are calculated according to formula (3), formula (4) and formula (8) respectively; for a batch of training data, the training loss is L; L=L l +consistency_weight×L ul In the above formula, consistency_weight is the consistency weight of the teacher-student model, which is updated with the number of training rounds. According to the training loss of the training data of this batch, the network parameters of a student model are reversely updated. At the same time, the network parameters of the teacher model are updated once by using the exponential moving average EMA using the updated network parameters of the student model to complete the training of a batch of data samples. Step 5.3: Select training set D train The next batch of training data in the dataset is taken. Similarly, the training data of this batch includes half of the labeled data and half of the unlabeled data. The network parameters of the student model and the teacher model when the previous batch of data samples are trained are used as the initial parameters for the next batch of training data. The process in step 5.2 and step 5.3 is repeated continuously until the last batch of training data in the dataset, including half of the labeled data and half of the unlabeled data, is trained to complete one round of training. The network parameters at the completion of the previous round of training are used as the initial network parameters at the beginning of the next round of training. The training round is repeated continuously until the preset 14,000 times are reached to obtain a trained segmentation network model. Step 6: 3D whole heart MRI image segmentation First, the patient's 3D whole-heart MRI image data is obtained, and then the 3D image data is center-cropped to obtain a picture of the same size as set in the first step, and the pixels of the picture are z-score standardized. After that, it is input into the student model in the trained segmentation network model in the fifth step to obtain the segmentation result, thus completing the segmentation of the 3D whole-heart MRI image data.

2. The semi-supervised 3D whole heart magnetic resonance image segmentation method according to claim 1, characterized in that: In the second step, the number of output channels of the first convolution block, the second convolution block, the third convolution block, the fourth convolution block, and the fifth convolution block are 16, 32, 64, 128, and 256, respectively, and the image dimension transformation is 80, 40, 20, 10, and 5 correspondingly.

3. The semi-supervised 3D whole heart magnetic resonance image segmentation method according to claim 1, characterized in that: In the second step, the number of output channels of the first convolution upsampling block, the second convolution upsampling block, the third convolution upsampling block, and the fourth convolution upsampling block are 128, 64, 32, and 16, respectively, and the image dimensions are 10, 20, 40, and 80, respectively.

4. The semi-supervised 3D whole heart magnetic resonance image segmentation method according to claim 1, characterized in that: In the second step, the number of output channels of the first Freq-Mamba Block, the second Freq-Mamba Block, and the third Freq-Mamba Block are all 256, and the image dimension is 10.

5. The semi-supervised 3D whole heart magnetic resonance image segmentation method according to claim 1, characterized in that: The perceptual loss network uses the 3D VGG model as the backbone with the number of channels being [1, 64, 256, 256, 512]. Each module includes a convolutional layer, a maximum pooling layer, and a ReLu activation function layer.

6. The semi-supervised 3D whole heart magnetic resonance image segmentation method according to claim 1, characterized in that: In the third, fourth, and fifth steps, when the student model starts training, the initial learning rate is 0.01 and the momentum parameter is 0.

9.

7. The semi-supervised 3D whole heart magnetic resonance image segmentation method according to claim 1, characterized in that: In the fourth step, for the features of the kth category collected, K-means clustering is used to obtain the small category prototype corresponding to this category. The specific operation is as follows: according to the number of sub-prototype centers in each small category prototype, subcluster sub-centers are randomly initialized, and then for each feature point in the features screened out by this category, its distance from each cluster center is calculated, and it is divided into the class represented by the sub-center with the closest distance; after completing the classification of all feature points, the center position of the cluster is recalculated according to the features of the sample points contained in each cluster; the above classification and recalculation of cluster centers are repeated until 300 iterations are reached. At this time, the subcluster cluster centers determined are the final clustering results of the data, and the feature points are divided into subcluster different sub-categories. Finally, the output of the small category prototype p that can represent this class is c_k , that is, the eigenvalues ​​of several subcluster centers, and the small category prototype with dimension [subcluster,proto_dim] is obtained.

8. The semi-supervised 3D whole heart magnetic resonance image segmentation method according to claim 1, characterized in that: In the fifth step, the student model weight θ S The exponential moving average of is used to update the weights θ of the teacher model T The specific operations are: At training iteration step t, the weight θ of the teacher model is T Updated to: in Respectively represent the parameters of the current student model and the current teacher model parameters, Indicates based on and The parameters of the teacher model for the next step are obtained; α is used to control the rate of EMA decay, which is dynamically updated with the number of iterations.

9. The semi-supervised 3D whole heart magnetic resonance image segmentation method according to claim 1, characterized in that: In the fifth step, consistency_weight is updated as follows: the current number of iterations iter_num is divided by 100, then smoothed using Sigmoid, and then multiplied by the preset basic consistency weight parameter value, thereby dynamically changing the value of consistency_weight as the rounds progress during training.

10. The semi-supervised 3D whole heart magnetic resonance image segmentation method according to claim 1, characterized in that: In the fourth step, the similarity measure q k The sinkhorn operation is iterated 3 times by default, so that the sum of its rows and columns is close to a uniform distribution.