Group Fisher Pruning for Practical Network Compression
Liyang Liu, Shilong Zhang, Zhanghui Kuang, Aojun Zhou, Jing-Hao Xue, Xinjiang Wang, Yimin Chen, Wenming Yang, Qingmin Liao, Wayne Zhang
Introduction
Modern computer vision models equipped with deep networks exhibit excellent performances in many tasks. However, they consume a great amount of memory and computation during inference. It can hinder the model deployment on edge devices where high-end hardwares are not available. It can also limit the throughput of services on clouds, resulting from considerable energy cost and inference latency. Network pruning aims at increasing the inference efficiency with negligible accuracy drop. It takes the trained dense model as input and prunes weights or channels with little importances. Through fine-tuning the pruned model can usually regain the lost performance caused by pruning.
Although numerous channel pruning methods have been proposed in literature (Molchanov et al., 2017; Luo et al., 2017; He et al., 2017; Liu et al., 2017), most of them study sequential networks such as AlexNet (Krizhevsky et al., 2012) and VGGNet (Simonyan & Zisserman, 2015) where pruning the input channel of a layer only affects the output channel of its single preceding layer. However, recently developed networks are designed with complicated structures such as residual connections in ResNet (He et al., 2016), group convolution (GConv) in ResNeXt (Xie et al., 2017) and RegNet (Radosavovic et al., 2020), depth-wise convolution (DWConv) in MobileNet (Sandler et al., 2018), and feature pyramid networks (FPN) (Lin et al., 2017a) in object detection frameworks. These structures have coupled channels distributed in multiple layers, which must be pruned or preserved simultaneously. Ignoring the coupled channels and pruning them independently will definitely hurt the efficiency in terms of both FLOPs (floating-point operations), memory access and actual speedup during inference.
In this paper we propose a general framework named Group Fisher Pruning that can be applied to various complicated structures. Particularly, we first introduce a binary mask initialized as for each input channel. Then we propose a layer grouping algorithm to automatically find the coupled channels given computation graph of the network, and we make the coupled channels share the same mask. The importance of a single channel is estimated by the loss change if it is discarded, which is approximated by Fisher information and is proportional to the squared mask gradient. Based on the single-channel importance, we obtain the overall importance of coupled channels by the principled chain rule of gradient computation. Pruning is done by iteratively setting the mask of the least important channel to , where the coupled channels are pruned together. Finally the network is fine-tuned to regain the lost accuracy. During fine-tuning and inference, the channels with masks are explicitly excluded from the network, and thus computation and memory cost can be practically reduced for acceleration.
Moreover, we propose to normalize importances of channels by their reductions of computation costs as we would like to prune the least important channels with the most computation overheads to achieve the best trade-off between accuracy and efficiency. However, we find the commonly-used reduction of FLOPs is a rather biased estimator for the actual inference speedup. In contrast, we propose to measure the computational complexity of a channel by its reduction of memory during pruning. Through experiments we find normalizing the channel importance by the reduction of memory is more correlated with the speedup than FLOPs in terms of the inference time on GPUs.
Our proposed Group Fisher Pruning has the following advantages. Firstly, it can prune any layers including those with coupled channels, and thus achieves better trade-off between accuracy drop and actual acceleration. Secondly, it prunes globally rather than locally (He et al., 2017; Luo et al., 2017), i.e., it obtains the pruning ratio for each layer automatically without the cumbersome sensitivity analysis of layer-wise pruning ratio (Yu et al., 2018), and thus leads to higher accuracy. Thirdly, it estimates importances of all channels in one pass via the principled Fisher information instead of multiple forward passes for individual channels (Luo & Wu, 2020), and thus is more efficient. Lastly, in contrast with (Liu et al., 2017), it does not depend on specific layers like batch normalization (BN) and thus is more general so that we can prune more sophisticated structures such as object detection networks, where such layers may not be naïvely adopted due to the larger input size.
To demonstrate the generalization ability and effectiveness of the proposed method to deal with complicated network structures, we conduct extensive experiments on various backbones, including classic ResNet (He et al., 2016) and ResNeXt (Xie et al., 2017), mobile-friendly MobileNetV2 (Sandler et al., 2018), and recent NAS-based RegNet (Radosavovic et al., 2020) on image classification (See Fig. 1). We also evaluate our method on object detection, which is more computation-intensive due to larger input image size and more complicated network structure than image classification, but rather under-explored (See Fig. 2).
Our main contributions are: we introduce the concept of coupled channels, find them by the proposed layer grouping algorithm, derive a unified metric to evaluate both coupled-channel and single-channel importances (based on Fisher information), and normalize the importance by memory reduction to realize higher speedup on GPUs without sacrificing the accuracy.
Related Work
Network pruning can be generally categorized into unstructured and structured methods. Unstructured methods (Han et al., 2015, 2016; Guo et al., 2016) prune unimportant weights in the model, but efficiency of the pruned sparse network can only be shown with the help of specialized libraries or hardwares. Recently, there are also efforts (Zhou et al., 2021; Mishra et al., 2021) to develop N:M fine-grained sparse models, leveraging the innovations in general-purpose GPUs (e.g., NVIDIA Ampere architecture). On the contrary, structured methods prune the whole channels or filters with little importances, and thus actual speedup can be easily achieved without requiring sparse accelerators.
For structured pruning methods (Wen et al., 2016; Lebedev & Lempitsky, 2016), different importance metrics have been proposed. PFEC (Li et al., 2017) employs norm of the channel weights, while SFP (He et al., 2018a) uses norm of each filter. These methods rely on the “smaller-norm-less-informative” assumption (Ye et al., 2018) which may not be true especially for structured pruning. CP (He et al., 2017) and ThiNet (Luo et al., 2017) cast channel selection as reconstruction error minimization of feature maps, where LASSO regression and greedy strategy are used to select the pruned channels, respectively. However, they can only prune networks in a layer-wise manner, as the least-square reconstruction happens locally. NISP (Yu et al., 2018) instead minimizes the reconstruction error of the final response layer and propagates importance scores through the entire network. The above methods also need sensitivity analysis to decide the pruning ratio for each layer, which may be time-consuming and sub-optimal. Network Slimming (Liu et al., 2017) reuses BN layer scaling factors as importance scores so that channels can be pruned globally. Although BN is prevalently used in image classification, many applications such as object detection can not trivially adopt BN because of the large input image size, which limits the application scenarios of pruning methods based on BN scaling factors. SSS (Huang & Wang, 2018) introduces extra scaling factors to scale the outputs of various micro-structures and solves the sparsity regularized optimization of scaling factors by the accelerated proximal gradient method.
Apart from the heuristic-based importance evaluation methods, one may use the exact loss change induced by removing a specific parameter (Luo & Wu, 2020) to measure its importance, but it is prohibitively expensive due to the large parameter number. Others try to approximate the importance score via Taylor expansion on the loss. The seminal work of OBD (LeCun et al., 1990) and OBS (Hassibi & Stork, 1993) exploit the second-order derivative information to estimate the importances of weights, but they may need to obtain the heavy-weight Hessian matrix which is too large to compute for modern large-scale networks. L-OBS (Dong et al., 2017) layer-wisely computes the Hessian matrix to achieve tractable approximation. WoodFisher (Singh & Alistarh, 2020) approximates the inverse of Hessian matrix by the Woodbury matrix identity and improves unstructured pruning based on OBD/OBS. PCNN (Molchanov et al., 2017) extends Taylor expansion to channel pruning and uses the first-order information instead. It takes a greedy strategy to prune the least important channels, interleaving pruning and fine-tuning. In place of estimating importances of feature maps, IE (Molchanov et al., 2019) applies Taylor expansion to the weights of a filter. These importance estimation methods are more principled than magnitude-based ones, but they seldom deal with structure constraints, for example, the residual connections.
Besides the importances, another factor needed to be concerned is computation, as we wish to prune the least important channels with the most computation costs. Current methods (Molchanov et al., 2017; Theis et al., 2018) typically add a regularization term to constrain FLOPs of the pruned network. However, the same amount of FLOPs reduction may lead to different actual speedups. Through experiments we empirically find that reduction of memory access can act as a more accurate estimator for efficiency gain, which is not explored in previous pruning methods.
Methodology
We first introduce Fisher information (Theis et al., 2018) as single-channel importance estimation, which can be used to prune channels in sequential networks but can not deal with complicated structures. Then we propose our layer grouping algorithm to find coupled channels in different layers, and make the coupled channels share the mask so as to prune them simultaneously. Finally we propose to use memory reduction as importance normalization to achieve better trade-off between efficiency and accuracy.
To evaluate the importance of a channel , we apply Taylor expansion to the network loss and approximate the loss change when discarding it (setting its mask to ):
which is proportional to the squared gradient of the mask. The above derivation is based on the model convergence, to satisfy it, a greedy pruning strategy is employed. Starting from a dense model, we first accumulate the importance scores by passing a few batches, then the least important channel is pruned. Next we fine-tune the pruned model and meanwhile re-accumulate the importance scores of the remained channels, following which the remained least important one is pruned, and the procedure recurs.
2 Prune Coupled Channels
Till now we can prune early-stage networks like AlexNet (Krizhevsky et al., 2012) and VGGNet (Simonyan & Zisserman, 2015) which involve normal Conv layers and sequential structures where pruning only affects a layer and its single preceding one. However, recent networks contain complicated structures such as residual connections (He et al., 2016), group convolutions (GConv) (Xie et al., 2017), depth-wise convolutions (DWConv) (Sandler et al., 2018) and feature pyramid networks (FPN) (Lin et al., 2017a) in object detection. There emerge coupled channels which should be pruned simultaneously to achieve higher speedup than pruning only the isolated channels. We propose mask sharing in coupled channels. Given the network computation graph containing nodes like convolution (Conv), batch normalization (BN), ReLU and pooling (Pool) layers, we adopt the proposed layer grouping algorithm to find the coupled channels as Alg. 1. Firstly, we use depth-first search (DFS) as Fig. 4 to find parents of each Conv/FC layer . Since channel pruning only affects the channel dimension, we ignore all layers except Conv/FC layers in during layer grouping. Then given parents of each layer, we can assign layers to different groups where layers in one group have coupled channels to be pruned simultaneously. It contains the following situations: (1) layers which have the same parents should be assigned to one group because their input channels (or equivalently, output channels of their parents) are coupled and should be pruned together as Fig. 3 (b); (2) layers whose parents contain GConv should be in the same group with their parents because the input and output channels of GConv are coupled as Fig. 3 (c). For the isolated channels, there is only one layer in the group such as the Conv of a residual bottleneck as Fig. 3 (a).
After obtaining the coupled channels via layer grouping, we make them share the same mask. Then the overall contribution of coupled channels can be computed by:
For pruning GConv with input channels divided into groups as Fig. 3 (c), each time we prune one group of channels which share the same mask, since generally they represent related features. We first compute the -dim gradient corresponding to individual channels and reshape it to (see Fig. 5 (a)), and do in-layer sum over the last dimension to obtain the -dim gradient corresponding to groups of channels. Next we compute the cross-layer summation across layers: the GConv layer itself and layers which are in the same group with the GConv including its children layers. Then we obtain the overall importance with the squared gradients as Eq. (4).
3 Importance Normalization
The raw importance scores do not take into consideration the computation costs of different channels, however it is more effective to prune the least important channel with the highest cost. Otherwise we may prune too many channels to achieve the desired speedup, but it may lead to degraded accuracy resulting from less parameters retained. We propose to normalize the importance scores by the computation reduction of each channel. We first try the widely-used FLOPs proxy and normalize the importance by the reduction of FLOPs / for pruning an input channel of normal Conv/GConv as and , where denotes the kernel height/width and is the group number in GConv. Different from previous methods which compute channel FLOPs in advance and fix them during pruning, we dynamically update the FLOPs by remained channels of the network. We use and to represent the unpruned input and output channel number of each layer. As the layers are connected internally, pruning a channel not only brings FLOPs reduction in the current layer, but also in its parent layers across the computation graph, which can be computed as and
However, we find that reduction of FLOPs is not directly correlated with the inference speedup. In contrast, reduction of memory increases linearly with speedup (Fig. 6 (a)), which motivates us to employ the memory reduction as the importance normalization. The memory reduction of pruning one channel can be obtained for normal Conv and GConv as and . Note that we discard one group at a time when pruning GConv. Similar to FLOPs reduction, pruning an input channel in one layer brings memory reduction from all of its parents (which can be found by DFS as Fig. 4), and we obtain the overall reduction by summing separate values. Finally we adopt the memory-normalized importance of each channel to evaluate its significance and prune the least important one every few iterations. There exist methods (Wang et al., 2020; Li et al., 2020) that directly profile the running time without resorting to proxies. However, it is not applicable here since we prune in a fine-grained manner. The running time difference of discarding one or few channels is too subtle to measure. Through memory-normalization (Fig. 6 (b)), we notice that in the pruned model, the reduction of memory is still linearly correlated with speedup. Besides, the reduction of FLOPs is more correlated with speedup than the FLOPs-normalization (Fig. 6 (a)) variant. In the appendix we demonstrate that the memory is a good proxy generally applicable to various networks.
Experiments
In this section we first conduct ablation studies to verify the effectiveness of our layer grouping and mask sharing strategy to prune coupled channels, and that of the proposed memory normalized importance scores. We measure the batch inference time on NVIDIA 2080 Ti GPU to prove our pruned models can significantly accelerate inference with little accuracy drop. Next we show that our method can outperform previous methods to prune various networks under different FLOPs constraints including the rather compact ResNet (He et al., 2016) and ResNeXt (Xie et al., 2017) where residual connection and GConv is adopted. Our method can be applied to prune MobileNetV2 (Sandler et al., 2018) where DWConv is presented, and it outperforms the uniform-scaled baselines remarkably. It can also be used to prune RegNet (Radosavovic et al., 2020) which is neural architecture search based and highly efficient and accurate, surprisingly we achieve higher accuracy and speed than the searched counterpart under the same FLOPs. Finally, we prune object detection networks with sophisticated structures and show significant speedup with negligible mAP drop. We conduct all experiments for the task of image classification and object detection on the ImageNet (Deng et al., 2009) and COCO (Lin et al., 2014) datasets, respectively. We prune a channel every iterations when pruning classification/detection networks. After the whole pruning process we fine-tune the pruned model for the same number of epochs that is used to train the unpruned model, which is trained following standard practices. The complete pruning pipeline of our proposed method is in Alg. 2. All experiments are done using PyTorch (Paszke et al., 2019) and more details can be found in the appendix.
As shown in Tab. 1, pruning coupled layers simultaneously (“-M”) as Fig. 3 (b) rather than the isolated channels only (“-I”) as Fig. 3 (a), for residual networks with both 50 and 101 layers, we achieve higher or comparable top-1 accuracy but with much higher inference speed, which verifies we can obtain better actual speedup under the same FLOPs. We also adopt the absolute value of the first-order gradient as the importance metric, but for pruning ResNet-50 it only obtains 75.8% top-1 accuracy, which lags behind our method based on Fisher information 76.4%.
Besides, we explore different importance normalization strategies: unnormalized (“-U”), normalized by FLOPs reduction (“-F”) and normalized by memory reduction (“-M”) in Fig. 2 (c) and Tab. 1. We find that normalizing importance scores by memory reduction can achieve the best accuracy-efficiency trade-off compared with the other two variants. The unnormalized importance score brings the worst efficiency gain and the largest accuracy drop, which results from the least parameter remained. For ResNet, ResNeXt and MobileNet, normalization by memory reduction is more efficiency-friendly and with higher accuracy.
Except for the rather compact residual networks, we also prune the light-weight networks MobileNetV2 (MBv2). In Fig. 2 (b) and Tab. 2 we compare the accuracy and speed of our pruned networks and the uniform scaled ones under different FLOPs budgets. To obtain a network with similar FLOPs as MBv2 (e.g., MBv2-0.7), we prune a uniform-scaled double-FLOPs MBv2 (e.g., MBv2-1.0) to 50% FLOPs remained. It can be seen that our pruned networks significantly outperform the uniform-scaled baselines. Other than only pruning the human-designed networks, we prune the highly-efficient RegNet to show that we can prune it to further boost the efficiency and accuracy. We prune RegNet (e.g., RegX-1.6G) to 50% FLOPs remained to compare with the searched half-FLOPs RegNet (e.g., RegX-0.8G). As shown in Fig. 2 (a) and Tab. 3, our pruned networks outperform the searched counterparts in this extreme circumstance.
2 Compare with SoTAs
To compare with previous state-of-the-arts, we conduct extensive experiments of image classification on ImageNet using different network structures and FLOPs constraints. From Fig. 1 and Tab. 4 we can see that our method performs best. Specifically, we outperform layer-wise pruning methods such as CP (He et al., 2017) and ThiNet (Luo et al., 2017) because we evaluate the importance scores globally throughout the network. Moreover, we do not need sensitivity analysis which is required by NISP (Yu et al., 2018) to decide the pruning ratio for each layer, as our method can automatically learn to prune the least important channels considering the current state of the network. We also achieve better accuracy than the methods C-SGD (Ding et al., 2019a), GBN (You et al., 2019) and IE (Molchanov et al., 2019) which compute the overall importance of coupled channels via heuristics, validating the benefits of our importance metric grounded on gradients obtained by the principled chain rule.
3 Prune for Detection
Pruning object detection is more challenging than image classification due to its larger input size and more complicated networks, which demands model pruning more than image classification. Besides, many pruning methods based on BN scaling parameters can not be directly applied. However, our method can not only be applied to image classification, but also object detection, thanks to its general importance estimation and the proposed layer grouping for pruning coupled channels. We prune one-stage methods including RetinaNet (Lin et al., 2017b), FSAF (Zhu et al., 2019), ATSS (Zhang et al., 2020) and PAA (Kim & Lee, 2020), and two-stage method Faster R-CNN (Ren et al., 2015) to extensively validate the effectiveness of our method. We present the pruning results of our method for various detection frameworks in Fig. 2 (d) and Tab. 5. Similar to pruning classification models, normalizing importances with memory reduction and pruning coupled channels together (“-M”) leads to the highest efficiency. Our method effectively prunes detection networks without losing average precision, in some cases our pruned model even receives higher mAP than the unpruned baseline. More importantly, our method delivers practical inference speedup, e.g., we achieve a speedup by pruning Faster R-CNN with only 0.8% mAP drop. We also compare the pruned networks with state-of-the-art method Slimmable Networks (Yu et al., 2019) in Tab. 6, as shown we can achieve higher mAP, lower mAP drop under comparable or less FLOPs.
Considering the intrinsic differences between image classification and object detection, we compare the pruned network structures between them. As in Fig. 7, we find that the pruned classification network keeps more capacity in later stages where the spatial resolution is rather small, as classification needs more global features. However, for detection the early stages also remain a large portion of channels, as detection should extract features at different scales to detect objects with various sizes. This validates that our method can be adaptively applied to different tasks.
Conclusion
We present a general channel pruning framework for complicated structures. We propose the layer grouping algorithm to find coupled channels and make them share the binary mask. Based on the single-channel importance approximated by Fisher information, we compute the overall importance of coupled channels by the chain rule of gradient computation. We prune the coupled channels simultaneously for better accuracy-efficiency trade-off. Moreover, normalizing channel importances by memory reduction rather than FLOPs is proposed to deliver more speedup. Extensive experiments on pruning various network structures with residual connections, GConv/DWConv and FPN in detection are explored and verify the effectiveness.
Inspired by the memory-bound nature of GPUs, we propose to normalize channel importance by memory reduction, which can bring a better trade-off between accuracy and speedup. In future work we will theoretically model the relationships between core factors (e.g., FLOPs, memory) and inference speed on various platforms (e.g., GPU/CPU/TPU).
Acknowledgements
The work described in this paper was partially supported by the Natural Science Foundation of Guangdong Province (No. 2020A1515010711), the Special Foundation for the Development of Strategic Emerging Industries of Shenzhen (No. JCYJ20200109143010272 and No. JCYJ20200109143035495), the Innovation and Technology Commission of the Hong Kong Special Administrative Region, China (Enterprise Support Scheme under the Innovation and Technology Fund B/E030/18) and the Shanghai Committee of Science and Technology, China (Grant No. 20DZ1100800).