Adaptive channel pruning method for three-dimensional medical image segmentation network

Through adaptive channel pruning method and offline knowledge distillation technology, the problem of high complexity of medical image segmentation models is solved, and the model is lightweight and efficiently deployed, while maintaining or improving the segmentation accuracy.

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

Patent Information

Application Number
CN202510082789.8
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2025-01-20
Publication Date
2025-05-13

AI Technical Summary

Technical Problem

Existing medical image segmentation models are complex and cannot be deployed in resource-constrained edge systems or mobile devices. The existing lightweight methods are low in flexibility, high in design costs, and may lead to degraded segmentation performance.

Method used

An adaptive channel pruning method for three-dimensional medical image segmentation network is proposed. By automatically analyzing the network structure, the channels with dependencies are grouped, the optimal pruning rate of each module is adaptively calculated, and the lost segmentation performance is restored through offline knowledge distillation technology.

Benefits of technology

The model is lightweighted, the network complexity is reduced, while maintaining or exceeding the segmentation accuracy of the original model, and improving the model's learning efficiency and resource utilization efficiency.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN119990235A_ABST
    Figure CN119990235A_ABST
Patent Text Reader

Abstract

The invention relates to a self-adaptive channel pruning method for a three-dimensional medical image segmentation network, and the method comprises the following steps: obtaining a medical image, and carrying out the image preprocessing; constructing a segmentation network model, automatically analyzing a dependency relationship between adjacent layers in the network, dividing the dependency relationship into inter-layer dependency and intra-layer dependency, completing grouping of channels through a matrix D, and selecting a corresponding pruning scheme according to the dependency relationship; evaluating the importance of each channel according to the attention map difference; the optimal pruning rate of each module is adaptively calculated according to different sensitivities of each module in the network to complexity and segmentation performance; a channel importance threshold value is obtained through the pruning rate, and channels lower than the threshold value are pruned; and finally, restoring the lost segmentation performance by combining an offline knowledge distillation technology to obtain a pruned model for three-dimensional medical image segmentation. The model segmentation performance can be ensured, and the network complexity can be greatly reduced.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention belongs to the technical field of three-dimensional medical image segmentation, and in particular relates to an adaptive channel pruning method for a three-dimensional medical image segmentation network. Background Art

[0002] Medical image segmentation is a key task in medical image analysis, which aims to extract specific regions or structures from medical images for subsequent analysis, diagnosis and treatment. Early medical image segmentation methods usually rely on edge detection, template matching, statistical shape models, active contours and traditional machine learning techniques. Although these methods have achieved good results to a certain extent, due to the diversity of medical images, existing technologies still face many challenges in feature extraction. With the gradual application of deep learning in the medical field and the realization of end-to-end learning methods of neural networks for medical image segmentation tasks, the feature extraction effect has been significantly improved, which has promoted the further improvement of medical image segmentation accuracy.

[0003] However, deploying the trained model to the actual medical image analysis system has brought new difficulties. This is because when medical images involve three dimensions, the extractable features will be richer, and a larger model is required to analyze the features of the image, resulting in a significant increase in model complexity. The number of parameters is usually over millions, and the amount of calculation is at the billion level. This high-precision but low-efficiency situation has seriously hindered the promotion of segmentation technology in practical applications and is not conducive to the conservation and utilization of resources. Although the use of cloud computing solutions can alleviate the problem of resource shortage to a certain extent, this method has the risk of leaking data privacy. Therefore, how to lightweight the segmentation model so that it can be better put into practical applications has become a problem that needs to be solved urgently.

[0004] The existing lightweight methods for medical image segmentation networks are mainly to directly design a lightweight network for a specific segmentation task. For example, a lighter layered decoupled convolution is used instead of a standard convolution to reduce the number of network parameters, and then the attention mechanism is combined to improve the segmentation performance. The artificially designed lightweight network has also achieved good results in different segmentation tasks. However, this method requires researchers to have deep knowledge of the relevant field and spend a lot of time to verify, so the design cost is relatively high and the flexibility is not high. The lightweight methods widely used in image classification, target detection and other fields have higher flexibility, such as pruning and knowledge distillation. Since the field of medical image segmentation usually has high requirements for segmentation accuracy, general lightweight methods such as pruning or knowledge distillation may cause a significant reduction in segmentation performance. Therefore, it is necessary to design a more flexible lightweight method for medical image segmentation networks, which can not only ensure the model segmentation performance, but also greatly reduce the network complexity. Summary of the invention

[0005] The problem to be solved by the present invention is that the medical image segmentation model is highly complex and cannot be deployed in resource-constrained edge systems or mobile devices, and an adaptive channel pruning method for a three-dimensional medical image segmentation network is proposed. First, after loading the trained model, the present invention can automatically analyze the network structure and group the channels with dependencies, which is a key step in preparing for subsequent correct pruning. Secondly, the importance of each channel is evaluated. Then, the optimal pruning rate of each module is adaptively calculated. Finally, the channel importance threshold is calculated by the pruning rate, and the channels with importance lower than the threshold are removed. After pruning, the offline knowledge distillation technology is used to fine-tune and restore the decreased accuracy. The large model before pruning is used to guide the small model after pruning to perform a small number of iterations of training, and the knowledge learned by the large model is passed to the small model, so that the small model can reach or even exceed the segmentation accuracy of the large model.

[0006] In order to achieve the above object, the technical solution adopted by the present invention is:

[0007] An adaptive channel pruning method for a three-dimensional medical image segmentation network, the segmentation method comprising the following contents:

[0008] Step 1: Obtain medical images and perform image preprocessing;

[0009] Step 2: Build a segmentation network model and automatically analyze the dependencies between adjacent layers in the network, which are divided into inter-layer dependencies and intra-layer dependencies: The input channel in the i-th layer is represented as The output channel is represented as Inter-layer dependency is the relationship between layers directly connected, that is, the output of the upper layer i With the input of the lower layer j With the same feature graph, this relationship is described as at this time, and will be grouped together, that is, if you want to Pruning, then the channel will also be pruned in the same way; the intra-layer dependency depends on the properties of the layer itself. The input and output of this type of layer are not independent of each other, but affect each other. At this time, the input and output channels of the layer will be grouped together and share the same pruning scheme. This dependency relationship is described as

[0010] Then, grouping is performed according to the dependency relationship, ensuring that there is a dependency relationship between channels in the same group and that each channel is not grouped repeatedly. The channel grouping is completed through a matrix D.

[0011]

[0012] Among them, ∨ and ∧ represent "OR" and "AND" operations respectively. represents the indicator function, and returns "True" to indicate that the condition is met; sch represents the scheme; the channels are grouped through the matrix D, and the corresponding pruning scheme is selected according to the dependency relationship;

[0013] Step 3: Evaluate the importance of each channel based on the attention map difference;

[0014] Step 4: According to the different sensitivity of each module in the network to complexity and segmentation performance, the optimal pruning rate of each module is adaptively calculated;

[0015] Step 5: Obtain the channel importance threshold through the pruning rate, and prune the channels below the threshold; finally, combine the offline knowledge distillation technology to restore the lost segmentation performance and obtain the pruned model for 3D medical image segmentation.

[0016] Specifically, in step 1, image preprocessing includes operations such as rotation, flipping and cropping, which are used to perform a series of enhancement operations on the loaded image to enhance the network learning effect. Evaluation indicators can be divided into two categories: complexity and performance. The complexity indicators include the number of parameters, the amount of calculation and the model size, and the performance indicators include the Dice coefficient.

[0017] Specifically, in step 3, the importance of the channel is evaluated by comparing the difference in attention maps before and after channel pruning. Assume that F i ∈R C×H×W is the three-dimensional tensor output by the i-th convolutional layer, and its size is C×H×W. Define a spatial mapping function M as the square sum of the feature maps along the channel direction, then transform the tensor F i ∈R C×H×W The feature map A can be obtained through the spatial mapping function M i ∈R H×W , through A i It can represent the importance of the channel, A i It can be expressed by the following formula 2.

[0018]

[0019] Among them, F i,j Represents the feature map of the jth channel in the i-th convolutional layer.

[0020] Then, after pruning a channel j in convolution layer i, the new attention map It can be expressed as:

[0021]

[0022] Finally, the importance of channel j can be evaluated by calculating the difference between the two attention maps. As shown in Equation 4:

[0023]

[0024] Among them, γ i,j That is, the difference between the attention maps before and after pruning of channel j in convolutional layer i, γ i,j The smaller it is, the smaller the difference between the two attention maps is, which means that the importance of channel j is lower.

[0025] Specifically, in step 4, considering the different redundancies in each layer or each module in the neural network, the same pruning rate cannot achieve the best pruning effect. The present invention proposes a method for adaptively assigning the best pruning rate to each module in view of the different sensitivities of different modules in the three-dimensional medical image segmentation model to the segmentation accuracy and 95% HausdorffDistance (95HD) as well as the parameter amount and the calculation amount, so as to achieve the best compression effect. HausdorffDistance is used to measure the distance between two subsets A and B in space. In the medical image segmentation task, especially for the segmentation of organs such as the left atrium or pancreas, the boundary segmentation of the organ needs to be more careful. Then HausdorffDistance can represent the maximum distance between the predicted segmentation region boundary and the real region boundary, and then measure the segmentation quality, as shown in the following formula 5.

[0026] H(A,B)=max(h(A,B),h(B,A)) (5)

[0027]

[0028] In Equations 6 and 7, ||·|| represents the distance form between point set A and point set B.

[0029] The function used to calculate the pruning rate is shown in Equation 8 below.

[0030] F=Dice loss +HD+log e Params+log e MACs (8)

[0031] Among them, Dice lossIndicates the accuracy loss value of the model before and after pruning. HD is 95% Hausdorff Distance. Params and MACs represent the number of parameters and the amount of calculation after pruning, respectively. Taking the logarithm of the two can ensure that their sizes are within a certain range. The size of the F value can determine the impact of the pruned module on the network performance. The smaller the F value, the smaller the impact of the module on the network. By calculating the F value of each module and using the proportion of these values ​​to calculate the pruning rate of each module, the performance of the pruned network can be guaranteed and its complexity can be limited.

[0032] Specifically, in step 5, the large model before pruning is regarded as the teacher model, and the small model after pruning is regarded as the student model, and fine-tuning is performed using offline knowledge distillation to restore accuracy. In knowledge distillation, the teacher model transfers knowledge to the student model through a loss function, and the distillation loss function is designed to be the KL divergence between the predicted values ​​of the teacher model and the predicted values ​​of the student model, as shown in the following formula 9.

[0033] L kd (p s ,p t )=KL(σ(p s / T)||σ(p t / T)) (9)

[0034] Where σ represents the softmax operation and T is the temperature hyperparameter. In the present invention, T=5, p s and p t Represent the predicted values ​​of the student model and the teacher model respectively. Since pruning leads to a reduction in model capacity and performance degradation, the use of knowledge distillation helps to compensate for the lost segmentation accuracy, so that the pruned model reaches or even exceeds the model before pruning.

[0035] The loss function of the final pruned model training is composed of the model segmentation loss function and the distillation loss function, as shown in Equation 10.

[0036] L s =λL t +(1-λ)L kd (10)

[0037] Among them, L t is the model segmentation loss function, including dice loss and cross entropy loss, as shown in Formula 11. kd is the distillation loss function, λ∈[0,1] is used to adjust L t and L kd For the contribution weight between , the present invention sets λ=0.3.

[0038] L t =L dice +L CE (11)

[0039] Compared with the prior art, the present invention has at least the following beneficial effects:

[0040] (1) The present invention is applicable to a three-stage adaptive channel pruning method for three-dimensional medical image segmentation, which is called the ACPrune method. It can automatically analyze the network structure and group channels with dependencies to facilitate the use of the same pruning strategy.

[0041] (2) The present invention proposes to adaptively calculate the pruning rate on a module-by-module basis, which can effectively avoid the subjective influence caused by the manually set uniform pruning rate and achieve a balance between the complexity and performance of the model.

[0042] (3) The present invention combines offline knowledge distillation to restore the accuracy of the network after pruning, improves the learning efficiency of the model, and saves time cost and resource consumption. BRIEF DESCRIPTION OF THE DRAWINGS

[0043] Figure 1 It is a flow chart for realizing the present invention;

[0044] Figure 2 This is an example of pruning using the present invention on V-Net;

[0045] Figure 3 The difference between the attention map before and after pruning;

[0046] Figure 4 The change in the number of convolutional layer output channels before and after V-Net pruning. DETAILED DESCRIPTION

[0047] The present invention will be further described below in conjunction with the accompanying drawings and specific implementations.

[0048] It should be understood that the embodiments described in the present invention are exemplary, and the specific parameters used in the description of the embodiments are only for the convenience of describing the present invention and are not used to limit the present invention.

[0049] The present invention is an adaptive channel pruning method for a three-dimensional medical image segmentation network, and the specific steps are as follows:

[0050] Step 1: Read 3D medical images from the dataset and input them into the segmentation network for training and testing after preprocessing.

[0051] This example uses the left atrium and pancreas datasets. The left atrium dataset consists of 100 MR images with a resolution of 0.625×0.625×0.625mm. 80 images are used for training, 20 for validation, and randomly cropped to a size of 112×112×80 as network input. The pancreas dataset includes 82 abdominal CT images, which are randomly divided into 62 for training and 20 for testing. The data is enhanced by rotation, scaling, and flipping, and randomly cropped to a size of 96×96×96 as network input.

[0052] Five networks commonly used in 3D medical image segmentation tasks were selected for pruning to prove the effectiveness of the present invention. The five networks are 3D U-Net, V-Net, DeepLabv3, DenseVoxNet and Attention U-Net. All networks were trained using the SGD optimizer before pruning, with a learning rate of 0.01 and a weight decay factor of 10. -4 , the loss functions are Dice loss and cross entropy loss. The batch size of left atrial data is 4, the batch size of pancreatic data is 2, and the training cycle is 160. The accuracy after pruning is recovered by fine-tuning a small number of iterations through offline knowledge distillation. All experiments were completed using PyTorch 2.0.1 and CUDA 11.8 on NVIDIA 4060 GPU.

[0053] Step 2: If Figure 2 As shown, taking V-Net as an example, if the bottleneck module is to be pruned, the encoder4 and decoder4 modules that are interdependent with it also need to be pruned, otherwise the network structure will be destroyed and cannot be learned. Therefore, these three modules can be grouped into the same group and the same pruning strategy can be adopted to solve the above problems. According to the characteristics of different layers in the neural network, the above grouping method can also be adopted. And in order to reduce the redundancy during grouping, it is more finely divided into intra-layer dependency and inter-layer dependency. When there is dependency between the input and output within the layer, the input and output need to be pruned at the same time, such as the batch normalization layer. When there is no dependency between the input and output, different pruning schemes can be selected, such as the convolutional layer. In this case, the intra-layer dependency does not exist but the inter-layer dependency exists, and the output of the previous layer needs to maintain the same pruning scheme as the input of the layer. By checking the relationship between all inputs and outputs in the network, the present invention can complete the analysis and grouping of channel dependencies.

[0054] Step 3: If Figure 3 As shown, it shows the difference in the attention map before and after pruning channel j. Calculate A i and The difference between them is used to evaluate the importance of each channel and arrange them in a certain order.

[0055] Step 4: Adaptively calculate the pruning rate of each module and assign the best pruning rate to each module. Figure 4 As shown in Figure 2, taking V-Net as an example, the figure shows the change in the number of output channels of the convolutional layer before and after pruning. The pruning rate varies depending on the network's ability to learn images from different data sets.

[0056] Step 5: Fine-tune the recovery accuracy. Let the large model before pruning guide the small model after pruning to train a small number of iterations to improve learning efficiency and restore or even exceed the performance before pruning. Test and record the model complexity index and performance index after pruning, and compare them with those before pruning. The experimental results are shown in Tables 1 and 2 below.

[0057] Among them, Params represents the parameter quantity. For a standard three-dimensional convolution layer, let the input channel be C in , the output channel is C out , the depth, height, and width of the convolution kernel are K d , K h , K w , then the parameter calculation method of the three-dimensional convolutional layer is as shown in the following formula 12:

[0058] Params=C in ×C out ×K d ×K h ×K w +C out (12)

[0059] Fewer parameters means the model is lighter and easier to deploy.

[0060] The number of multiplication and addition operations MACs represents the amount of model calculation. In the three-dimensional convolution layer, the depth, height, and width of the input feature map are D in , H in , W in , the step size is S, the padding is P, then the output feature map size after the convolution operation is D out , H out , W out It can be calculated by the following formula:

[0061]

[0062] Then, MACs are calculated as:

[0063] MACs = D out ×H out ×W out ×C out ×C in ×K d ×Kh ×K w (16)

[0064] Lower MACs indicate that the model operates more efficiently.

[0065] Size is the model size. After saving the model, you can get this value by checking the memory space occupied by the model.

[0066] Dice is used to evaluate the segmentation accuracy of the model, and its value range is 0 to 1. The Dice calculation formula is as follows:

[0067]

[0068] Among them, A represents the predicted value and B represents the true value. The higher the Dice value, the more overlap the predicted result has with the true value, and the better the segmentation effect.

[0069] Table 1 Left atrium dataset results

[0070]

[0071] Table 2 Results of pancreas dataset

[0072]

[0073] It can be observed from Table 1 that after pruning, the number of parameters of the five models is reduced by at least 80%, the amount of calculation is reduced by at least 5 times, and the model size is reduced by at least 80%. In addition, the segmentation accuracy of the pruned models is almost intact after fine-tuning through knowledge distillation, and some models even exceed the segmentation level before pruning, indicating that there is a certain degree of overfitting or redundancy in the original model design. The channels removed during pruning can be regarded as unimportant. They do not have much positive impact on the performance of the model, and may even have a negative impact, because the segmentation performance of some models is more advantageous after pruning. Through fine-tuning of knowledge distillation, the model can use a more compact structure and more critical parameters to learn the features of the image and achieve better performance. Although the fine-tuning process consumes some time, a lighter model is obtained, which is convenient for deployment in more restricted environments, so these costs are worthwhile.

[0074] Since the number of parameters is only related to the network structure and has nothing to do with the size of the input dataset, the number of parameters of each network before pruning in Table 2 is the same as that of the left atrium experiment. It can be observed that the number of model parameters after pruning is reduced by at least 85%. Since the size of the pancreatic image input to the network is smaller than that of the left atrium, the MACs value of the model before pruning in Table 6 is also slightly lower than that of the left atrium experiment. It can be observed that the DeepLabv3 model has the highest redundancy in terms of computational complexity, and the MACs value is reduced by 16 times after pruning. The model size in Table 7 also reflects the effectiveness of pruning, which is reduced by at least 85%. Since the pancreas has irregular imaging features and fuzzy boundaries, and has a lower contrast than the surrounding fat, it can be seen from the segmentation accuracy shown in Table 2 that the segmentation difficulty is higher than that of the left atrium dataset on the same neural network. After pruning, 3 / 5 of the models surpassed the performance before pruning in segmentation accuracy, and the remaining 2 / 5 models were only reduced by 0.31% at most compared with before pruning. It can be concluded that pruning has achieved a very lightweight effect, especially for the DeepLabv3 network, where the pruning effect is the most obvious and the accuracy of the network is improved by 1.48%, indicating that the network has a lot of redundancy or overfitting problems. After pruning, the network will pay more attention to the features or connections that are most important for model prediction. Therefore, the present invention helps the network focus more on learning key features and improves the model's discrimination ability.

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

Claims

1. An adaptive channel pruning method for three-dimensional medical image segmentation network, characterized in that: The segmentation method includes the following contents: Step 1: Obtain medical images and perform image preprocessing; Step 2: Build a segmentation network model and automatically analyze the dependencies between adjacent layers in the network, which are divided into inter-layer dependencies and intra-layer dependencies: The input channel in the i-th layer is represented as The output channel is represented as Inter-layer dependency is the relationship between layers directly connected, that is, the output of the upper layer i With the input of the lower layer j With the same feature graph, this relationship is described as at this time, and will be grouped together, that is, if you want to Pruning, then the channel will also be pruned in the same way; the intra-layer dependency depends on the properties of the layer itself. The input and output of this type of layer are not independent of each other, but affect each other. At this time, the input and output channels of the layer will be grouped together and share the same pruning scheme. This dependency relationship is described as Then, grouping is performed according to the dependency relationship, ensuring that there is a dependency relationship between channels in the same group and that each channel is not grouped repeatedly, and the channel grouping is completed through a matrix D; Among them, ∨ and ∧ represent "OR" and "AND" operations respectively. represents the indicator function, and returns "True" to indicate that the condition is met; sch represents the scheme; the channels are grouped through the matrix D, and the corresponding pruning scheme is selected according to the dependency relationship; Step 3: Evaluate the importance of each channel based on the attention map difference; Step 4: According to the different sensitivity of each module in the network to complexity and segmentation performance, the optimal pruning rate of each module is adaptively calculated; Step 5: Obtain the channel importance threshold through the pruning rate, and prune the channels below the threshold; finally, combine the offline knowledge distillation technology to restore the lost segmentation performance and obtain the pruned model for 3D medical image segmentation.

2. The method according to claim 1, characterized in that The image preprocessing includes rotation, flipping and cropping operations.

3. The method according to claim 1, characterized in that In step 3, the importance of the channel is evaluated by comparing the difference in the attention map before and after channel pruning. Assuming F i ∈R C×H×W is the three-dimensional tensor output by the i-th convolutional layer, and its size is C×H×W; define a spatial mapping function M as the square sum of the feature maps along the channel direction, then transform the tensor F i ∈R C×H×W The feature map A can be obtained through the spatial mapping function M i ∈R H×W , through A i Indicates the importance of the channel, A i The following formula 2 represents: Among them, F i,j Represents the feature map of the jth channel in the i-th convolutional layer; Then, after pruning a channel j in convolution layer i, the new attention map It is expressed as: Finally, the importance of channel j can be evaluated by calculating the difference between the two attention maps, as shown in Equation 4: Among them,γij is the difference between the attention maps of the channel,before pruning and after pruning in the convolutional layer,i. The smaller,γi,j,is, the smaller the difference between the two attention maps is, which means that the channel,is less important.

4. The method according to claim 1, characterized in that In step 4, according to the different sensitivities of different modules in the three-dimensional medical image segmentation network to segmentation accuracy and 95% Hausdorff Distance (95HD) as well as parameter quantity and calculation quantity, each module is adaptively assigned its own optimal pruning rate mode to achieve the best compression effect; HausdorffDistance is used to measure the distance between two subsets A and B in space. In the medical image segmentation task, the boundary segmentation of organs is more detailed. Then HausdorffDistance represents the maximum distance between the predicted segmentation region boundary and the true region boundary, and then measures the segmentation quality, as shown in Formula 5: H(A,B)=max(h(A,B),h(B,A)) (5) In Equations 6 and 7, ||·|| represents the distance form between point set A and point set B; The function used to calculate the pruning rate is shown in Equation 8; F=Dice loss +HD+log e Params+log e MACs (8) Among them, Dice loss It represents the accuracy loss value of the model before and after pruning. HD is 95% Hausdorff Distance. Params and MACs represent the number of parameters and the amount of calculation after pruning, respectively. The logarithm of the two is taken to ensure that their sizes are within a certain range. The size of the F value is used to judge the impact of the pruned module on the network performance. The smaller the F value, the smaller the impact of the module on the network. By calculating the F value of each module and using the proportion of these values ​​to calculate the pruning rate of each module, the performance of the pruned network is guaranteed and its complexity is limited.

5. The method according to claim 1, characterized in that In step 5, the large model before pruning is regarded as the teacher model, and the small model after pruning is regarded as the student model, and fine-tuning is performed by offline knowledge distillation to restore accuracy; the distillation loss function is the KL divergence between the predicted value of the teacher model and the predicted value of the student model, as shown in Formula 9: L kd (p s ,p t )=KL(σ(p s / T)||σ(p t / T)) (9) Where σ represents the softmax operation and T is the temperature hyperparameter; The loss function of the final pruned model training is composed of the model segmentation loss function and the distillation loss function, as shown in Equation 10: THE s =λL t +(1-λ)L kd (10) Among them, L t is the model segmentation loss function, including dice loss L dice and the cross entropy loss L CE , as shown in formula 11; L kd is the distillation loss function, λ∈[0, 1] is used to adjust L t and L kd The contribution weight between L t =L dice +L CE (11)