Towards Optimal Structured CNN Pruning via Generative Adversarial Learning

Shaohui Lin, Rongrong Ji, Chenqian Yan, Baochang Zhang, Liujuan Cao, Qixiang Ye, Feiyue Huang, David Doermann

Introduction

Convolutional neural networks (CNNs) have achieved state-of-the-art accuracy in computer vision tasks such as image recognition and object detection . However, the success of CNNs is often accompanied by significant computation and memory consumption that restricts their usage on resource-limited devices, such as mobile or embedded devices. To address these issues, techniques have been proposed for CNN compression such as low-rank decomposition , parameter quantization , knowledge distillation and network pruning . Network pruning has received a great deal of research focus demonstrating significant compression and acceleration of CNNs in practice.

Our main contributions are summarized as follows:

We propose a generative adversarial learning (GAL) to effectively conduct structured pruning of CNNs. It is able to jointly prune redundant structures, including filters, branches and blocks to improve the compression and speedup rates.

Adversarial regularization is introduced to prevent a trivially-strong discriminator, soft mask is used to solve the slackness of hard filter pruning, and FISTA is employed to fast and reliably remove the redundant structures.

Extensive experiments demonstrate the superior performance of our approach. On ImageNet ILSVRC 2012 , the pruned ResNet-50 achieves 10.88% Top-5 error with a factor of 3.7×3.7\times speedup outperforming state-of-the-art methods.

Related Work

Recently, binary masks have been proposed to guide filter pruning. Yu et al. proposed a Neuron Importance Score Propagation (NISP) to optimize the reconstruction error of the “final response layer” and propagate an “importance score” to each node, i.e., 1 for important nodes, and 0 otherwise. Lin et al. directly learned a global mask with binary values, and pruned the filters whose mask values are 0. However, such a hard filter pruning lacks effectiveness and slackness, due to the NP-hard optimization caused by using the binary mask. Our method slacks the binary mask to the soft one, which largely improves the flexibility and accuracy.

In line with our work, sparse scaling parameters in batch normalization (BN) or in the specific structures were obtained by supervised training with a class-labelled dataset. In contrast, our approach obtains the sparse soft mask with label-free data and can transfer to other scenarios with unseen labels.

Neural Architecture Search: While state-of-the-art CNNs with compact architectures have been explored with hand-crafted design , automatic search of neural architectures is also becoming popular. Recent work on searching models with reinforcement learning or genetic algorithms greatly improve the performance of neural networks. However, the search space of these methods is extremely large, which requires significant computational overhead to search and select the best model from hundreds of models. In contrast, our method learns a compact neural architecture by a single training, which is more efficient. Group sparsity regularization on filters or multiple structures including filter shapes and layers has been proposed to sparsify them during training. This is also less efficient and cannot reliably remove the sparse structures since only stochastic gradient descent is used.

Knowledge Distillation: The proposed generative adversarial learning for structured pruning is also related to knowledge distillation (KD) to a certain extent. KD transfers knowledge from the teacher to the student using different kinds of knowledge (e.g., dark knowledge and attention ). Hinton et al. introduced dark knowledge for model compression, which uses the softened final output of a complicated teacher network to teach a small student network. Romero et al. proposed FitNets to train the student network by combining dark knowledge and the knowledge from the teacher’s hint layer. Zagoruyko et al. transferred the knowledge from attention maps from the teacher’s hidden layer to improve the performance of a student network. Unlike other methods, we do not require labels to train the pruned network. Furthermore, we directly copy the architecture of the student network from the teacher without being designed by experts, and then automatically learn how to prune the student network.

Note that our approach is orthogonal to other compression approaches, such as low-rank decomposition , or parameter quantization . We can integrate our approach into the above methods to achieve higher compression and speedup rates.

Our Method

2 Formulation

where LAdv(WG,m,WD)\mathcal{L}_{Adv}(\mathcal{W}_{G},\mathbf{m},\mathcal{W}_{D}) is the adversarial loss to train the two-player game between the baseline and the pruned network that compete with each other. This is defined as:

where pb(x)p_{b}(x) and pg(x)p_{g}(x) represent the feature distributions of the baseline and the pruned network, respectively. pz(x)p_{z}(x) corresponds to the prior distribution of noise input zz. Inspired by , we use the dropout as the noise input zz in the pruned network. This dropout is active only while updating the pruned network. For notation simplicity, we omit zz in fg(x,z)f_{g}(x,z).

In addition, Ldata(WG,m)\mathcal{L}_{data}(\mathcal{W}_{G},\mathbf{m}) is the data loss between output features from both the baseline and the pruned network, which is used to align the outputs of these two networks. Therefore, the data loss can be expressed by MSE loss:

where nn is the number of the mini-batch size.

Finally, Lreg(WG,m,WD)\mathcal{L}_{reg}(\mathcal{W}_{G},\mathbf{m},\mathcal{W}_{D}) is a regularizer on WG,m\mathcal{W}_{G},\mathbf{m} and WD\mathcal{W}_{D}, which can be split into three parts as follows:

We found the discriminator DD is updated only with correct prediction by using Eq. (2), which leads to a less valuable gradient updating that the pruned network receives. Therefore, adversarial regularization is introduced to also update the discriminator DD with the features of pruned network produced by the baseline, and to extend the time of the two-player game to achieve more valuable gradients.

3 Optimization

Following , Stochastic Gradient Descent (SGD) can be directly introduced to alternately update the discriminator DD and generator GG to solve the optimization problem in Eq. (1). However, SGD is less efficient in convergence, and by using SGD we have observed non-exact zero scaling factors in the soft mask m\mathbf{m}. We therefore need a threshold to remove the corresponding structures, whose scaling factors are lower than the threshold. By doing so, the accuracy of the pruned network is significantly lower than the baseline. To solve this problem, we introduce FISTA into the GAN to effectively solve the optimization problem of Eq. (1) via two alternating steps. Algorithm 1 presents the optimization process.

First, we use SGD to optimize the weights WD\mathcal{W}_{D} of the discriminator DD by ascending its stochastic gradient to solve Eq. (6). The entire procedure mainly relies on the standard forward-backward pass. Second, for better illustration, we shorten the first two terms of Eq. (7) as H(WG,m)\mathcal{H}(\mathcal{W}_{G},\mathbf{m}), and we have:

We solve the optimization problem of Eq. (8) by alternately updating WG\mathcal{W}_{G} and m\mathbf{m}. (1) Fixing m\mathbf{m}, we use SGD with momentum to update WG\mathcal{W}_{G} by descending its gradient. (2) Fixing WG\mathcal{W}_{G}, the optimization of m\mathbf{m} is reformulated as:

Then m\mathbf{m} is updated by FISTA with the initialization of α(1)=1\alpha_{(1)}=1:

where η(k+1)\eta_{(k+1)} is the learning rate at the iteration k+1k+1 and proxη(k+1)λ∥⋅∥1(zi)=sign(zi)∘(∣zi∣−η(k+1)λ)+\textbf{prox}_{\eta_{(k+1)}\lambda\|\cdot\|_{1}}(\mathbf{z}_{i})=\text{sign}(\mathbf{z}_{i})\circ(|\mathbf{z}_{i}|-\eta_{(k+1)}\lambda)_{+}.

We solve these two steps by following stochastic methods with the mini-batches and set the learning rate η\eta with fixed-step updating. Moreover, we update WG,m\small{\mathcal{W}_{G},\mathbf{m}} and WD\small{\mathcal{W}_{D}} at each iteration (i=j=1i=j=1 in Algorithm 1).

4 Structure Selection

To achieve flexible structure selection, we add a soft mask after the three different kinds of structures from coarse to fine-grained, including blocks, branches and channels, to remove the redundancy of different networks ResNets , GoogLeNet and DenseNets as shown in Fig. 1. Furthermore, these structures can be integrated into each other for jointly learning.

Block Selection: For ResNets, the residual block contains the residual mapping with a large number of parameters and the shortcut connections with few parameters. This achieves high performance by skipping the computation of specific layers to overcome the degradation problem. The block is removed by setting the residual mapping to zero, but cannot cut off the information flow in ResNets. Therefore, block selection is significantly effective when applied in ResNets. The new residual block by adding the soft mask is formulated as:

where zi\mathbf{z}^{i} and zi+1\mathbf{z}^{i+1} are the input and output of the ii-th block, respectively. F\mathcal{F} is a residual mapping and {WGi}\{\mathcal{W}_{G}^{i}\} are weights of the ii-th block. After optimization, we obtain a sparse soft mask m\mathbf{m}, in which the ii-th residual block can be pruned if mi=0m_{i}=0.

Branch Selection. Multi-branch networks such as GoogLeNet and ResNeXts have been proposed to enhance the information flow to achieve high performance. Similar to ResNets, there is redundancy in the branch that can be removed entirely by setting the corresponding soft mask to 0. Likewise, this does not cut off the information flow in multi-branch networks. Taking GoogLeNet for instance, we can formulate the new inception module by adding the soft mask as follows:

where [⋅][\cdot] represents concatenation operator. τi(z,{WGi})\tau^{i}(\mathbf{z},\{\mathcal{W}_{G}^{i}\}) is a transformation with all weights {WGi}\{\mathcal{W}_{G}^{i}\} at the ii-th branch and cc is the number of branch in one inception module. We can reliably remove the ii-th branch, which satisfies mi=0m_{i}=0 after optimization.

Channel Selection: The channel is a basic element in all CNNs and has large amounts of redundancy. In our framework, we add the soft mask after input at the current layer (the output feature maps at the upper layer) to guide the input channel pruning at the current layer and the output channel pruning at the upper layer. Therefore, the formulation at the ll-th layer is as follows:

where zil\mathbf{z}_{i}^{l} and zjl+1\mathbf{z}_{j}^{l+1} are the ii-th input feature map and the jj-th output feature map at the ll-the layer, respectively. WGi,jl\mathbf{W}_{G_{i,j}}^{l} represents the 2D kernel of ii-th input channel in the jj-th filter at the ll-th layer. ∗* and f(⋅)f(\cdot) refer to convolutional operator and non-linearity (ReLU), respectively. After training, we remove the feature maps with a zero soft mask that are associated with the corresponding channels at the current layer and the filters at the upper layer.

Experiments

We evaluate the proposed GAL approach on three widely-used datasets, MNIST , CIFAR-10 and ImageNet ILSVRC 2012 . We use channel selection to prune plain networks (LeNet and VGGNet ) and DenseNets , branch selection for GoogLeNet , and block selection for ResNets . For ResNets, we also leverage channel selection to block selection that jointly prunes these heterogeneous structures to largely improve the performance of the pruned network.

Implementations: We use PyTorch to implement GAL. We solve the optimization problem of Eq. (1) by running on two NVIDIA GTX 1080Ti GPUs with 128GB of RAM. The weight decay is set to 0.0002 and the momentum is set to 0.9. The hyper-parameter λ\lambda is selected by cross-validation in the range [0.01, 0.1] for channel pruning on LeNet, VGGNet and DenseNets, and the range [0.1, 1] for branch and block pruning on GoogLeNet and ResNets. The drop rate in dropout is set to 0.1. The other training parameters are discussed in different datasets in Section 4.2.

Discriminator Architecture: The discriminator DD plays a very important role in striking a balance between simplicity and network capacity to avoid being trivially fooled. In this paper, we select a unified and relative simple architecture, which is composed of three fully-connected (FC) layers and non-linearity (ReLU) with the neurons of 128-256-128. The input is the features from the baseline fb(x)f_{b}(x) and the pruned network fg(x)f_{g}(x), while the output is the binary prediction to predict the input from baseline or pruned network.

2 Comparison with the State-of-the-art

We evaluate the effectiveness of GAL on MNIST in LeNet. For training parameters, we apply GAL with three groups of hyper-parameter λ\lambda (0.01, 0.05 and 0.1) with the mini-batch size of 128 for 100 epochs. The initial learning rate is set to 0.001 and is scaled by 0.1 over 40 epochs. As shown in Table 1, compared to SSL and NISP , GAL achieves the best trade-off between FLOPs/parameter pruned rate and the classification error. For example, by setting λ\lambda to 0.05, the error of GAL only increases by 0.1% with 92.6% and 93% pruned rate in FLOPs and parameter, respectively. In addition, we found that fine-tuning the pruned LeNet with GAL only achieves a limited decrease in error. Fine-tuning instead increases the error when λ\lambda is set to 0.1. This is due to the fact that the output features learned by GAL have already had a strong discriminability, which may be reduced by fine-tuning.

2.2 CIFAR-10

We further evaluate the performance of the proposed GAL on CIFAR-10 in five popular networks, VGGNet, DenseNet-40, GoogLeNet, ResNet-56 and ResNet-110. For VGGNet, we take a variation of the original VGG-16 for CIFAR-10 from . DenseNet-40 has 40 layers with growth rate 12. For GoogLeNet, we also take a variation of the original GoogLeNet by changing the final output class number for CIFAR-10.

VGGNet: The baseline achieves the classification error 6.04%. GAL is applied to prune it with the mini-batch size of 128 for 100 epochs. The initial learning rate is set to 0.01, and is scaled by 0.1 over 30 epochs. As shown in Table 2, compared to L1 and SSS , our GAL achieves a lowest error and highest pruned rate in both FLOPs and parameters. For example, GAL with setting λ\lambda to 0.05 achieves the lowest error (6.23% vs. 6.60% by L1 and 6.37% by SSS) by the highest pruned rate of FLOPs (39.6% vs. 34.3% by L1 and 36.3% by SSS) and parameters (77.6% vs. 64.0% by L1 and 73.8% by SSS).

DenseNet-40: According to the principle of channel selection in Section 3.4, we should prune the input channels at the current layer and the corresponding output feature maps and the filters at the upper layer in DenseNets. But this leads to a mismatch of the dimension in the following layers. This is due to the complex dense connectivity of each layer in DenseNets. We therefore only prune the input channels in DenseNet-40, as suggested in . The training setup is the same to VGGNet, except the mini-batch size is 64. The pruning results of DenseNet-40 are summarized in Table 3. GAL achieves a comparable result with Liu et al. . For example, when λ\lambda is set to 0.01, 3362 out of 8904 channels are pruned by GAL with a higher computational saving of (35.3% vs. 32.8%), but only with a slightly higher error (5.39% vs. 5.19%), compared to Liu et al.-40%.

GoogLeNet: For better comparison, we re-implemented L1 and APoZ on GoogLeNet and also introduce random pruning, because of lack of pruning results on GoogLeNet in CIFAR-10. For Random, L1 and APoZ, we simply prune the same number of branches in each inception module based on their pruning criteria as GAL-0.5 for a fair comparison. The training parameters of GAL are the same to prune DenseNet-40 (not including λ\lambda) and the first convolutional layer is skipped to add the soft mask. As presented in Table 4, GAL achieves the best trade-off by removing 14 of 36 branches with a rate of FLOPs saving of 38.2%, parameter saving of 49.3% and only an increase of 0.49% classification error, compared to all methods. This is because GAL employs the more flexible branch selection by learning the soft mask than L1 and APoZ based on the statistical property. Note that the simplest random approach works reasonably well, which is possibly due to the self-recovery ability of the distributed representations. In addition, the branches of 3×33\times 3 convolutional filters with a large number of parameters are more removed by APoZ, which leads to significant FLOPs and parameters reduction and also significant error increase.

ResNets: To evaluate the effectiveness of block selection in GAL, we use ResNet-56 and ResNet-110 as our baseline models. The training parameters of GAL on both ResNet-56 and ResNet-110 are the same to prune VGGNet (not including λ\lambda) and the first convolutional layer is also skipped to add the soft mask. The pruning results of both ResNet-56 and ResNet-110 are summarized in Table 5. For ResNet-56, when λ\lambda is set to 0.6, 10 out of 27 residual blocks are removed by GAL, which achieves a 37.6% pruned rate in FLOPs while with a decrease of 0.12% error. This indicates that there are redundant residual blocks in ResNet-56. Moreover, compared to L1 and NISP , GAL-0.6 also achieves the best performance. When more residual blocks are pruned (16 when λ\lambda is set to 0.8), GAL-0.8 still achieves the higher pruned rate in FLOPs (60.2% vs. 50.6%), with a slightly higher classification error (8.42% vs. 8.20%) compared to He et al. . For ResNet-110, compared to L1, GAL achieves better results by pruning 10 out of 54 residual blocks, when λ\lambda is set to 0.1. After optimization for ResNet-56 and ResNet-110, the bottom residual blocks are easier to prune. To explain, top blocks often have high-level semantic information that is necessary for maintaining the classification accuracy.

2.3 ImageNet ILSVRC 2012

GAL was also evaluated on ImageNet using ResNet-50. We train the pruned network with the mini-batch size of 32 for 30 epochs. The initial learning rate is set to 0.01 and is scaled by 0.1 over 10 epochs. As shown in Table 6, GAL without jointly pruning blocks and channels is able to obtain 1.76×1.76\times and 2.59×2.59\times speedup (FLOPs rate) (2.33B and 1.58B vs. 4.09B in ResNet-50) by setting λ\lambda to 0.5 and 1, with an increase of 1.93% and 3.12% in Top-5 error, respectively. However, GAL-0.5 and GAL-1 only achieve a 1.2×1.2\times and 1.74×1.74\times parameter compression rate, which is due to the fact that most of the pruned blocks comes from the bottom layers with a small number of parameters. By jointly pruning blocks and channels, we achieve a higher speedup and compression. For example, compared to GAL-0.5, GAL-0.5-joint achieves the higher speedup and compression by a factor of 2.22×2.22\times and 1.32×1.32\times (vs. 1.75×1.75\times and 1.2×1.2\times), respectively. Furthermore, compared to SSS-26 , He et al and GDP-0.6 , GAL-0.5-joint also achieves the best trade-off between Top-5 error and speedup. With almost the same speedup, our GAL-1-joint outperforms ThiNet-30 by 0.89% and 0.82% in Top-1 and Top-5 error, respectively.

3 Ablation Study

To evaluate the effectiveness of GAL, which lies in adversarial regularization, FISTA and GANs, we select ResNet-56 and DenseNet-40 for an ablation study.

We train our GAL approach with three types of discriminator regularizers, L1-norm, L2-norm and adversarial regularization (AR). For a fair comparison, all the training parameters are the same. As shown in Fig. 2, adversarial regularization achieves the best performance, compared to the L1-norm and L2-norm. This is because AR prolongs the competition between generator and discriminator to achieve the better output features of generator, which are close to baseline and fool the discriminator.

3.2 Effect on the Optimizers

We compare our FISTA with SGD optimizer. For SGD, we cannot obtain the soft mask with an exact scaling factor of 0. Therefore, a hard threshold is required in the pruning stage. We set the threshold to 0.0001 in our experiments. As presented in Table 7, compared to the random method, SGD achieves a lower error with the same architecture. It indicates that SGD provides better initial values for the pruned network (PN). After pruning with thresholding, the accuracy drops significantly (See the columns of Error and PN in Table 7), as the pruned small near-zero weights might have large impact on the final network output. Advantageously, GAL with FISTA can safely remove the redundant structures in the training process, and achieves better performance compared to SGD.

3.3 Effect of the GANs

We train the pruned network with and without the GAN, and also make a comparison with CGAN by using the FISTA. For training CGAN, we only need to modify the adversarial loss function in Eq. (2) by the loss of CGAN, and the optimization with related training parameters are same as GAL. The results are summarized in Fig. 3. First, the lack of GANs leads to significant error increase. Second, the GAN achieves a better result than CGAN. For example, with the same regularization and optimizer on ResNet-56, label-free GAL achieves a 8.42% error with a 65.9% parameter pruned rate vs. 9.56% error with 50.5% parameter pruned rate in label-dependent CGAN. We conjecture this is due to the class label that is added to the discriminator in CGAN, which instead affects the output features of generator to approximate baseline during training.

Conclusion

References