A medical image segmentation method based on feature rearrangement and gated axial attention
By adopting feature rearrangement and gating axial attention mechanism in the medical image segmentation model, the global and local branching are coordinatedly trained, and the difficulties of existing models in learning global and remote semantic information interaction are solved, and higher medical image segmentation accuracy is achieved.
Patent Information
- Application Number
- CN202111262731.X
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2021-10-28
- Publication Date
- 2025-05-13
- Estimated Expiration
- 2041-10-28
AI Technical Summary
The existing medical image segmentation model has difficulties in learning global and remote semantic information interactions, resulting in the segmentation accuracy not being able to fully meet the needs of medical applications.
Feature rearrangement and gated axial attention mechanism are adopted to coordinate the training of global branches and local branches through an end-to-end method, extract the global information interaction and local information interaction characteristics of the input original image, and control the circulation of information in the network through the gated mechanism.
This method can segment medical images more accurately, especially on small sample data sets, which can also learn good position deviation information, improving the accuracy of medical image segmentation.
Smart Images

Figure CN114049314B_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to the field of medical image segmentation, and in particular to an end-to-end medical image segmentation method. Background Art
[0002] With the popularity of deep convolutional neural networks in computer vision, deep convolutional neural networks have been used for medical image segmentation tasks. Networks such as U-Net, Res-UNet, and U-Net++ have been specifically proposed to perform image segmentation of various medical imaging modalities. These methods have also achieved good performance on many difficult datasets, demonstrating the effectiveness of CNN in learning discriminative features to segment organs or lesions from medical scans. Although CNN-based methods have achieved excellent performance in the field of medical image segmentation, due to the inherent locality of convolution operations, it is difficult for CNN-based methods to learn the exact global and long-range semantic information interactions. Its segmentation accuracy cannot fully meet medical applications. Medical image segmentation remains a challenging task in medical image analysis.
[0003] Recently, Transformer-based methods have made great progress in the field of computer vision. The main reason for the success of Transformers is that they can learn long-term dependencies between input tokens and learn global and long-range semantic information interactions. The Transformer-based deformable axial attention block decomposes 2D self-attention into two 1D self-attentions and introduces position-sensitive axial attention, which has been used for panoramic segmentation. The Transformer-based deformable gated axial attention can use a gating mechanism to control the flow of attention information in the network.
[0004] Most existing medical image segmentation models use convolutional layers or pooling layers to downsample the original image. For upsampling, convolution plus bilinear interpolation is used for upsampling. In the process of downsampling using convolutional layers or pooling methods, part of the original image information will be lost. The present invention uses feature rearrangement to save the original information on the C channel. Since bilinear interpolation depends on the interpolation formula and has poor flexibility for different data sets, the present invention uses reverse feature rearrangement for upsampling. At the same time, the convolution block is replaced with an axial attention block with global and long-range semantic information interaction gate to achieve better segmentation accuracy for medical images. Summary of the invention
[0005] The present invention provides a medical image segmentation method based on feature rearrangement and gated axial attention. The method adopts feature rearrangement and gated axial attention mechanism, and trains global branches and local branches in an end-to-end manner, which can effectively extract the global information interaction and local information interaction features of the input original image, and control the flow of information in the network through the gating mechanism, so that the model can also learn good position deviation information on a small sample data set. Experimental results show that this method can segment medical images more accurately.
[0006] A medical image segmentation method based on feature rearrangement and gated axial attention, the steps of which are as follows:
[0007] Step 1. Dataset acquisition: Select three datasets from existing public medical image segmentation datasets;
[0008] The three data sets in the data acquisition are the gland segmentation data set Glas, which contains 85 training images and 80 test images, the cell nucleus segmentation data set MoNuSeg, which contains 30 training images and 14 test images, and the cell nucleus segmentation data set TNBC, which contains 35 training images and 15 test images.
[0009] Step 2. Data processing: on the medical image segmentation dataset obtained in step 1, adjust the images in the dataset to the same size; then randomly flip the adjusted training sample images horizontally / vertically to increase the diversity of training samples;
[0010] Step 3. Define a medical image segmentation model based on feature rearrangement and gated axial attention, which includes a global branch and a local branch; take the training image processed in step 2 and the real segmentation map of the training image as input;
[0011] Step 4. Loss function: The loss function is used to measure the error between the predicted value and the true sample label. Here, the cross entropy loss function is used.
[0012] Step 5. Define the Adam optimizer and set a reasonable learning rate for the model. The initial learning rate is set to 0.001. During the model training process, the learning rate slows down as the number of batches increases. The learning rate is adjusted to the original 0.8 every 50 batches, thereby effectively suppressing oscillation and finding better network parameters. At the same time, L2 regularization is used to effectively reduce overfitting.
[0013] The learning rate decay formula is as follows (3):
[0014] l p =l0×0.8 p / / 50 (3)
[0015] In the above formula, p is the number of training batches (epoch). The hyperparameter used to define the L2 regularization term is 0.0005.
[0016] Step 6. Network training and testing. Co-train the global branch and local branch in step 3. During training, evaluate the network on the test set provided by each dataset. The evaluation method uses the average IoU and the average F1 score.
[0017] Furthermore, the data processing described in step 2 is specifically implemented as follows:
[0018] First, the original images and the real segmentation maps in the dataset are resized to 128 × 128 size;
[0019] Finally, the resized training images and the corresponding segmented images are randomly flipped horizontally / vertically with a probability of 50%.
[0020] Furthermore, the global branch and local branch of the model described in step 3 are specifically implemented as follows:
[0021] Global branch encoding (enconder) part:
[0022] 3-1. Use a 7×7 convolution kernel with a step size of 1 and a padding of 3 for the input training sample image, retain the input height H and width W, and then map it through the BatchNorm layer and the ReLU activation function to obtain the feature block;
[0023] 3-2. Rearrange the features of the feature blocks, divide the H and W patches into 2×2 patches, and downsample (B, C, H, W) to (B, 4C, H / 2, W / 2); where B is the number of input images at a time, C is the number of channels of the feature block, and H and W are the height and width of the feature block respectively; feature rearrangement can retain the features of adjacent elements on the C channel, which can better retain information than using the pooling layer;
[0024] 3-3. The feature block after feature rearrangement is subjected to two convolutions without changing the height and width, plus the BatchNorm layer and the ReLU activation function mapping to obtain the feature block x, thereby enhancing the information flow of the local patch block;
[0025] 3-4. Input the feature block x into the gated axial attention block; a gated axial attention block first passes through a 1×1 convolution plus BatchNorm layer and ReLU activation function mapping; then, gated axial attention is applied along the width axis of the tensor, as shown in formula (1):
[0026]
[0027] Among them, N represents the width after the last downsampling, q ij =W q′ x,k ij =W k′ x,v ij =W v′ x are query, key and value respectively; W is a linear transformation, x is input and y is output;
[0028] q ij =W q′ x, represents the jth element in the i-th row, W q′ Represents the parameters of the linear transformation of x
[0029] k ij =W k′ x, represents the jth element in the i-th row, W k′ Represents the parameters of the linear transformation of x
[0030] v ij =W v′ x, represents the jth element in the i-th row, W v′ Represents the parameters of the linear transformation of x
[0031] Among them, qe,ke,e are all learnable position deviation terms, usually called relative position encoding;
[0032] Where G is a learnable gating parameter used to control the influence of the learned relative position encoding on the encoded non-local context; then, the gated axial attention is applied along the height axis of the tensor; the method is the same as the gated axial attention applied along the width axis; finally, a 1 × 1 convolution plus BatchNorm layer and ReLU activation function mapping are passed, and the feature block establishes a residual connection before and after each gated axial attention block;
[0033] 3-5. The feature block obtained after one gated axial attention is downsampled by feature rearrangement, and finally passes through two gated axial attention blocks to obtain a compact feature value block x1 with long-range dependency and global information;
[0034] Global branch decoding part:
[0035] First, the feature block x1 is passed through two gated axial attention blocks, and then the C channel is expanded using a 1×1 convolution so that the reverse feature rearrangement can be used for upsampling; then the reverse feature rearrangement is performed for upsampling; compared with the upsampling using bilinear interpolation, feature rearrangement is more flexible, and its parameters can be obtained by network learning, while bilinear interpolation is obtained by the formula; secondly, it is passed through a gated axial attention block, and then the C channel is expanded using a 1×1 convolution, and then the reverse feature rearrangement is performed for upsampling to obtain a feature block x2 with long-range dependency and global information and the same size as the input original image. All corresponding blocks of the decoding part and the encoding part use skip connections, and the connection method is addition;
[0036] Local branch part:
[0037] 301. The input training sample image is divided into 32×32 patches on the H and W patches, and then the following operations are performed on each patch:
[0038] 7×7 convolution kernel, stride 1, Padding set to 3, retain the input height and width, and then pass through the BatchNorm layer and ReLU activation function mapping to obtain the feature block;
[0039] 302. The feature block obtained in step 301 is subjected to feature rearrangement and downsampling, and then subjected to two convolutions without changing the height and width, plus a BatchNorm layer and a ReLU activation function mapping to obtain a feature block x';
[0040] 303. Downsample x' by three gated axial attention blocks and feature reordering, and the gated axial attention blocks used each time are 3, 4, and 1 respectively; obtain feature block x'1; then upsample x'1 by three gated axial attention blocks and inverse feature reordering, and the gated axial attention blocks used each time are 1, 4, and 3 respectively;
[0041] 304. Each patch block is spliced in the original order to obtain a feature block x'2 with local information and the same size as the input original image.
[0042] Furthermore, the loss function described in step 4 is as follows (2):
[0043]
[0044] Where H and W are the dimensions of the image, p(x,y) corresponds to a pixel in the image, and p'(x,y) represents the output prediction for a specific location (x,y).
[0045] Furthermore, the collaborative training and evaluation indicators described in step 6 are as follows:
[0046] The collaborative training of the global branch and the local branch is to add the feature blocks x2 and x'2 obtained in step 3, and then pass a 1×1 convolution to reduce the C channel to the number of categories required for segmentation;
[0047] The average IoU and average F1 score refer to first calculating the IoU and F1 score of each predicted segmentation map and the true segmentation map of different categories, and then taking their average; then adding the mean IoU and mean F1 score of all test sets respectively, and then dividing by the number of test images to get the average IoU and average F1 score; these two indicators can effectively evaluate the accuracy of model segmentation.
[0048] The beneficial effects of the present invention are as follows:
[0049] The present invention segments medical images based on feature rearrangement and gated axial attention. This method uses a downsampling method of feature rearrangement, which can retain more original image information compared to convolution and pooling. It uses an upsampling method of inverse feature rearrangement, which is more flexible than bilinear interpolation. The corresponding parameters can be learned by the network itself and are not limited to being determined by formulas. The gated axial attention model with global branches and local branches not only considers local information interaction and global information interaction, but also uses a gating mechanism to control the flow of information in the network. Appropriately adopting some training techniques, selecting ideal network parameters, optimization algorithms, and learning rate settings can improve the accuracy of medical image segmentation. BRIEF DESCRIPTION OF THE DRAWINGS
[0050] Figure 1 It is a flow chart of the present invention.
[0051] Figure 2 It is a schematic diagram of the network framework of the present invention.
[0052] Figure 3 This is a comparison chart of the segmentation effects of the present invention and other methods. DETAILED DESCRIPTION
[0053] The present invention will be further described below in conjunction with the accompanying drawings and embodiments.
[0054] like Figure 1 As shown, a medical image segmentation method based on feature rearrangement and gated axial attention specifically includes the following steps:
[0055] Step 1. Dataset acquisition: select three datasets from existing public medical image segmentation datasets. The datasets used in the present invention are respectively the gland segmentation dataset Glas, which contains 85 training images and 80 test images; the cell nucleus segmentation dataset MoNuSeg, which contains 30 training images and 14 test images; and the cell nucleus segmentation dataset TNBC, which contains 35 training images and 15 test images.
[0056] Step 2. Data processing, first, the original images and the real segmentation images in the dataset are unified into 128 × 128 sizes. Finally, the unified training images and the corresponding segmentation images are randomly flipped horizontally / vertically with a probability of 50%, which can effectively avoid overfitting on the one hand, and improve the model performance to a certain extent on the other hand.
[0057] Step 3. Figure 2 The figure shows the network framework of the medical image segmentation model based on feature rearrangement and gated axial attention, which consists of two parts, namely the global branch and the local branch. The training image processed in step 2 and the real segmentation map of the training image are used as input.
[0058] Global branches and local branches are implemented as follows:
[0059] Global branch encoding (enconder) part:
[0060] First, the input original image uses a 7×7 convolution kernel with a step size of 1 and a padding of 3, retaining the input height and width, and then passes through the BatchNorm layer and ReLU activation function mapping.
[0061] Next, we rearrange the features and divide the H and W patches into 2×2 patches, that is, we downsample (B, C, H, W) to (B, 4C, H / 2, W / 2). B is the number of input images at a time, C is the number of channels of the feature block, and H and W are the height and width of the feature block. This keeps the features of adjacent elements on the C channel, which can better preserve information than using the pooling layer.
[0062] Then, the feature blocks obtained above are convolved twice without changing the height and width.
[0063] The BatchNorm layer and the ReLU activation function map the feature block x to enhance the information flow of the local patch block.
[0064] Then, the feature block x is input into the gated axial attention block. A gated axial attention block first passes through a 1×1 convolution plus BatchNorm layer and ReLU activation function mapping. Secondly, the gated axial attention is applied along the width axis of the tensor, and the formula is as follows:
[0065]
[0066] Among them, N represents the width after the last downsampling, q ij =W q′ x,k ij =W k′ x,vij =W v′ x are query, key and value respectively; W is a linear transformation, x is input and y is output;
[0067] q ij =W q′ x, represents the jth element in the i-th row, W q′ Represents the parameters of the linear transformation of x
[0068] k ij =W k′ x, represents the jth element in the i-th row, W k′ Represents the parameters of the linear transformation of x
[0069] v ij =W v′ x, represents the jth element in the i-th row, W v′ Represents the parameters of the linear transformation of x
[0070] Among them, qe, ke, e are all learnable position deviation terms, usually called relative position encoding. Among them, G is a learnable gating parameter used to control the influence of the learned relative position encoding on the encoded non-local context. Then, the gated axial attention is applied along the height axis of the tensor. The method is consistent with the gated axial attention applied on the width axis. Finally, it passes through a 1×1 convolution plus BatchNorm layer and ReLU activation function mapping, and the feature block establishes a residual connection before and after passing through each gated axial attention block.
[0071] The feature block x first passes through a gated axial attention block, then performs feature rearrangement to downsample, and finally passes through two gated axial attention blocks to obtain a compact feature value block x1 with long-range dependencies and global information.
[0072] Global branch decoding part:
[0073] First, the feature block x1 passes through two gated axial attention blocks, and then uses 1×1 convolution to expand the C channel so that reverse feature rearrangement can be used for upsampling. Then reverse feature rearrangement is performed for upsampling. Compared with bilinear interpolation upsampling, feature rearrangement is more flexible and its parameters can be obtained by network learning, while bilinear interpolation is derived from the formula. Secondly, it passes through another gated axial attention block, and then uses 1×1 convolution to expand the C channel, and then reverse feature rearrangement is performed for upsampling to obtain a feature block x2 with long-range dependency and global information and the same size as the input original image. All corresponding blocks of the decoding part and the encoding part use skip connections, and the connection method is addition.
[0074] Local branch part:
[0075] First, divide the input original image into 32×32 patches on the H and W patches. Then, perform the following operations on each patch:
[0076] Use a 7×7 convolution kernel with a stride of 1 and a padding of 3. Keep the height and width of the input, and then map it through the BatchNorm layer and the ReLU activation function.
[0077] Secondly, the feature block obtained in the previous step is re-arranged and down-sampled, and then the feature block x' is obtained by adding a BatchNorm layer and a ReLU activation function mapping without changing the height and width twice.
[0078] Then, x' is downsampled by three gated axial attention blocks and feature reordering, and the gated axial attention blocks used each time are 3, 4, and 1 respectively. The feature block x'1 is obtained. Then x'1 is upsampled by three gated axial attention blocks and inverse feature reordering, and the gated axial attention blocks used each time are 1, 4, and 3 respectively.
[0079] Finally, each patch is concatenated in the original order to obtain a feature block x′2 with local information and the same size as the input original image.
[0080] Step 4. Define the loss function. The loss function is used to measure the error between the predicted value and the true sample label. The cross entropy loss function commonly used in segmentation tasks is used here as follows (2):
[0081]
[0082] Where w, h are the dimensions of the image, p(x, y) corresponds to a pixel in the image, and p'(x, y) represents the output prediction for a specific location (x, y).
[0083] Step 5. Define the Adam optimizer and set a reasonable learning rate for the model. The initial learning rate is set to 0.001. During the model training process, the learning rate slows down as the number of batches increases. The learning rate is adjusted to the original 0.8 every 50 batches, thereby effectively suppressing oscillation and finding better network parameters. At the same time, L2 regularization is used to effectively reduce overfitting, and the hyperparameter of the regularization term is 0.0005. The learning rate attenuation formula is defined as follows (3):
[0084] l p =l0×0.8 p / / 50 (3)
[0085] In the above formula, p is the number of training batches (epoch).
[0086] Step 6. Network training and testing, co-training global branches and local branches, adding the feature blocks x2 and x′2 obtained in step 3, and then performing a 1×1 convolution to reduce the C channel to the number of categories required for segmentation. Then, the corresponding prediction block is mapped to the C channel by the sigmoid function, and the cross entropy loss between it and the true segmentation map is calculated. Finally, the Adam optimizer defined in step 5 is used for gradient update. A total of 400 training batches are performed.
[0087] After every 50 training batches, a test is performed. First, a sigmoid function mapping is performed on the C channel, and then the corresponding pixels are classified into the category with the highest probability. The evaluation indicators used are average IoU and average F1 score. Average IoU and average F1 score refer to first calculating the IoU and F1 scores of different categories of each predicted segmentation map and the true segmentation map, and then taking their average. Then, the mean IoU and mean F1 score of all test sets are added together, and then divided by the number of test images to obtain the average IoU and average F1 score. These two indicators can effectively evaluate the accuracy of model segmentation.
[0088] The comparison model used in the experiment is the MedT model which has the best performance on the Glas and MoNuSeg datasets recently. The comparison of experimental indicators is shown in Table 1 below, and the comparison of segmentation effects is shown in the attached figure. Figure 3 .
[0089]
[0090] Table 1 Comparison of indicators between the present invention and the MedT model.
Claims
1. A medical image segmentation method based on feature rearrangement and gated axial attention, characterized in that The steps include: Step 1. Dataset acquisition: Select three datasets from existing public medical image segmentation datasets; Step 2. Data processing: on the medical image segmentation dataset obtained in step 1, adjust the images in the dataset to the same size; then randomly flip the adjusted training sample images horizontally / vertically to increase the diversity of training samples; Step 3. Define a medical image segmentation model based on feature rearrangement and gated axial attention, which includes a global branch and a local branch; take the training image processed in step 2 and the real segmentation map of the training image as input; Step 4. Loss function: The loss function is used to measure the error between the predicted value and the true sample label. Here, the cross entropy loss function is used. Step 5. Define the Adam optimizer and set a reasonable learning rate for the model. The initial learning rate is set to 0.
001. During the model training process, the learning rate slows down as the number of batches increases. The learning rate is adjusted to the original 0.8 every 50 batches, thereby effectively suppressing oscillation and finding better network parameters. At the same time, L2 regularization is used to effectively reduce overfitting. Step 6. Network training and testing: Co-train the global branch and local branch in step 3. During training, evaluate the network on the test set provided by each dataset. The evaluation method is the average IoU and the average F1 score. The data processing described in step 2 is specifically implemented as follows: First, the original images and the real segmentation maps in the dataset are resized to 128×128 size; Finally, the resized training image and the corresponding segmented image are randomly flipped horizontally / vertically with a probability of 50%; The global branch and local branch of the model described in step 3 are specifically implemented as follows: Global branch encoding part: 3-1. Use a 7×7 convolution kernel with a step size of 1 and a padding of 3 for the input training sample image, retain the input height H and width W, and then map it through the BatchNorm layer and the ReLU activation function to obtain the feature block; 3-2. Rearrange the features of the feature blocks, divide the H and W patches into 2×2 patches, and downsample (B, C, H, W) to (B, 4C, H / 2, W / 2); where B is the number of input images at a time, C is the number of channels of the feature block, and H and W are the height and width of the feature block respectively; feature rearrangement can retain the features of adjacent elements on the C channel, which can better retain information than using the pooling layer; 3-3. The feature block after feature rearrangement is subjected to two convolutions without changing the height and width, plus the BatchNorm layer and the ReLU activation function mapping to obtain the feature block x, thereby enhancing the information flow of the local patch block; 3-4. Input the feature block x into the gated axial attention block; a gated axial attention block first passes through a 1×1 convolution plus BatchNorm layer and ReLU activation function mapping; then, gated axial attention is applied along the width axis of the tensor, as shown in formula (1): Among them, N represents the width after the last downsampling, q ij =W q′ x,k ij =W k′ x,v ij =W v′ x is the query, key and value respectively; W is a linear transformation, x is the input and y is the output; q ij =W q′ x, represents the jth element in the i-th row, W q′ represents the parameter of the linear transformation of x; k ij =W k′ x, represents the jth element in the i-th row, W k′ represents the parameters of the linear transformation of x; v ij =W v′ x, represents the jth element in the i-th row, W v′ Represents the parameters of the linear transformation of x; in are all learnable position bias terms, usually called relative position encodings; G is a learnable gating parameter used to control the influence of the learned relative position encoding on the encoded non-local context; then, the gated axial attention is applied along the height axis of the tensor; the method is the same as the gated axial attention applied along the width axis; finally, a 1×1 convolution plus BatchNorm layer and ReLU activation function mapping are performed, and the feature block establishes a residual connection before and after passing through each gated axial attention block; 3-5. The feature block obtained after one gated axial attention is downsampled by feature rearrangement, and finally passes through two gated axial attention blocks; a compact feature value block x1 with long-range dependency and global information is obtained; Global branch decoding part: First, the feature block x1 is passed through two gated axial attention blocks, and then the C channel is expanded using a 1×1 convolution so that the reverse feature rearrangement can be used for upsampling; then the reverse feature rearrangement is performed for upsampling; compared with the upsampling using bilinear interpolation, feature rearrangement is more flexible, and its parameters can be obtained by network learning, while bilinear interpolation is obtained by the formula; secondly, it is passed through a gated axial attention block, and then the C channel is expanded using a 1×1 convolution, and then the reverse feature rearrangement is performed for upsampling to obtain a feature block x2 with long-range dependency and global information and the same size as the input original image. All corresponding blocks of the decoding part and the encoding part use skip connections, and the connection method is addition; Local branch part:
301. The input training sample image is divided into 32×32 patches on the H and W patches, and then the following operations are performed on each patch: 7×7 convolution kernel, stride 1, Padding set to 3, retain the input height and width, and then pass through the BatchNorm layer and ReLU activation function mapping to obtain the feature block; 302. The feature block obtained in step 301 is subjected to feature rearrangement and downsampling, and then subjected to two convolutions without changing the height and width, plus a BatchNorm layer and a ReLU activation function mapping to obtain a feature block x'; 303. Subtract x' from the original image by three gated axial attention blocks and feature reordering, and the gated axial attention blocks used each time are 3, 4, and 1 respectively; obtain feature block x1'; then subtract x1' from the original image by three gated axial attention blocks and reverse feature reordering, and the gated axial attention blocks used each time are 1, 4, and 3 respectively; 304. Each patch block is spliced in the original order to obtain a feature block x'2 with local information and the same size as the input original image.
2. The medical image segmentation method based on feature rearrangement and gated axial attention according to claim 1, characterized in that The loss function described in step 4 is as follows (2): Where H and W are the dimensions of the image, p(x,y) corresponds to a pixel in the image, and p'(x,y) represents the output prediction for a specific location (x,y).
3. The medical image segmentation method based on feature rearrangement and gated axial attention according to claim 2 is characterized in that The collaborative training and evaluation metrics described in step 6 are as follows: The collaborative training of the global branch and the local branch is to add the feature blocks x2 and x'2 obtained in step 3, and then pass a 1×1 convolution to reduce the C channel to the number of categories required for segmentation; The average IoU and average F1 score refer to first calculating the IoU and F1 scores of different categories of each predicted segmentation map and the true segmentation map, and then taking their average; Then add the mean IoU and mean F1 score of all test sets respectively, and divide by the number of test images to get the average IoU and average F1 score; these two indicators can effectively evaluate the accuracy of model segmentation.